diff --git a/ruoyi-fastapi-backend/.env.dev b/ruoyi-fastapi-backend/.env.dev index 78ec319dc..3b7a0d537 100644 --- a/ruoyi-fastapi-backend/.env.dev +++ b/ruoyi-fastapi-backend/.env.dev @@ -155,6 +155,77 @@ LOG_SERVICE_NAME = 'ruoyi-fastapi-backend' # Worker 标识(auto 自动生成) LOG_WORKER_ID = 'auto' +# -------- 统一认证中心 OIDC 配置 -------- +# 默认关闭;启用前必须配置生产密钥、Pepper 和已注册 Origin。 +# 是否启用统一认证中心 OIDC 功能 +OIDC_ENABLED = false +# OIDC 签发者地址(必须为公开的绝对 URL) +OIDC_ISSUER = 'https://auth.example.com' +# OIDC 对外公开基址,须与 OIDC_ISSUER 一致 +OIDC_PUBLIC_BASE_URL = 'https://auth.example.com' +# 是否强制所有授权请求使用 PKCE +OIDC_REQUIRE_PKCE = true +# 允许的 PKCE 方法,多个值使用逗号分隔(当前仅支持 S256) +OIDC_PKCE_METHODS = 'S256' +# 校验令牌时间戳时允许的时钟偏差(秒) +OIDC_ALLOWED_CLOCK_SKEW_SECONDS = 60 +# 授权码有效期(秒) +OIDC_AUTHORIZATION_CODE_TTL_SECONDS = 90 +# 登录、同意等交互状态有效期(秒) +OIDC_INTERACTION_TTL_SECONDS = 300 +# ID Token 有效期(秒) +OIDC_ID_TOKEN_TTL_SECONDS = 300 +# Access Token 默认有效期(秒) +OIDC_ACCESS_TOKEN_TTL_SECONDS = 600 +# Access Token 允许的最大有效期(秒) +OIDC_MAX_ACCESS_TOKEN_TTL_SECONDS = 1800 +# Refresh Token 闲置有效期(秒) +OIDC_REFRESH_TOKEN_IDLE_SECONDS = 604800 +# Refresh Token 绝对有效期(秒) +OIDC_REFRESH_TOKEN_ABSOLUTE_SECONDS = 2592000 +# SSO 会话闲置超时时间(秒) +OIDC_SSO_IDLE_SECONDS = 1800 +# SSO 会话绝对有效期(秒) +OIDC_SSO_ABSOLUTE_SECONDS = 28800 +# SSO“记住我”模式的绝对有效期(秒) +OIDC_SSO_REMEMBER_ABSOLUTE_SECONDS = 604800 +# OIDC 签名算法 +OIDC_SIGNING_ALGORITHM = 'RS256' +# 签名密钥来源,可选 file、kms、hsm +OIDC_SIGNING_KEY_SOURCE = 'file' +# 签名私钥文件路径(file 模式使用) +OIDC_SIGNING_PRIVATE_KEY_PATH = '' +# 签名私钥加密密钥 +OIDC_SIGNING_KEY_ENCRYPTION_KEY = '' +# 当前启用的签名密钥 KID,留空由数据库 active 密钥决定 +OIDC_ACTIVE_KID = '' +# 密钥轮换期间旧密钥保留时间(秒) +OIDC_KEY_ROTATION_OVERLAP_SECONDS = 86400 +# 不透明 Token 哈希 Pepper(至少 32 bytes,生产环境必须配置) +OIDC_TOKEN_HASH_PEPPER = '' +# 允许的 CORS Origin 列表,多个值使用逗号分隔 +OIDC_CORS_ALLOWED_ORIGINS = '' +# 登录交互页面 URL +OIDC_INTERACTION_LOGIN_URL = 'https://auth.example.com/auth-center/login' +# 授权同意页面 URL +OIDC_INTERACTION_CONSENT_URL = 'https://auth.example.com/auth-center/consent' +# 交互错误页面 URL +OIDC_INTERACTION_ERROR_URL = 'https://auth.example.com/auth-center/error' +# SSO Cookie 名称(必须使用 __Host- 前缀) +OIDC_SSO_COOKIE_NAME = '__Host-ruoyi-sso' +# SSO Cookie 是否启用 Secure 属性 +OIDC_SSO_COOKIE_SECURE = true +# SSO Cookie SameSite 策略(当前必须为 lax) +OIDC_SSO_COOKIE_SAMESITE = 'lax' +# SSO Cookie Domain(__Host- Cookie 必须留空) +OIDC_SSO_COOKIE_DOMAIN = '' +# 是否启用 Legacy 认证隔离 +OIDC_LEGACY_AUTH_ISOLATION_ENABLED = true +# 审计日志保留天数 +OIDC_AUDIT_RETENTION_DAYS = 180 +# Back-channel Logout 请求超时时间(秒) +OIDC_BACKCHANNEL_LOGOUT_TIMEOUT_SECONDS = 5 + # -------- 传输层加解密配置 -------- # 是否启用传输层加解密 TRANSPORT_CRYPTO_ENABLED = false @@ -187,7 +258,7 @@ TRANSPORT_CRYPTO_ENABLED_PATHS = '' # 强制要求传输层加密的路径列表,多个值使用逗号分隔 TRANSPORT_CRYPTO_REQUIRED_PATHS = '' # 排除传输层加密的路径列表,多个值使用逗号分隔 -TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download' +TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download,/.well-known/openid-configuration,/.well-known/oauth-authorization-server,/oauth2/authorize,/oauth2/token,/oauth2/userinfo,/oauth2/jwks,/oauth2/revoke,/oauth2/introspect,/oauth2/logout' # -------- 插件依赖安装策略配置 -------- # 策略模式:可填单一模式,或按环境映射。可选 disabled、plan_only、explicit、locked、offline diff --git a/ruoyi-fastapi-backend/.env.dockermy b/ruoyi-fastapi-backend/.env.dockermy index a14db6a00..a3bb6bcf7 100644 --- a/ruoyi-fastapi-backend/.env.dockermy +++ b/ruoyi-fastapi-backend/.env.dockermy @@ -155,6 +155,77 @@ LOG_SERVICE_NAME = 'ruoyi-fastapi-backend' # Worker 标识(auto 自动生成) LOG_WORKER_ID = 'auto' +# -------- 统一认证中心 OIDC 配置 -------- +# 默认关闭;启用前必须配置生产密钥、Pepper 和已注册 Origin。 +# 是否启用统一认证中心 OIDC 功能 +OIDC_ENABLED = false +# OIDC 签发者地址(必须为公开的绝对 URL) +OIDC_ISSUER = 'https://auth.example.com' +# OIDC 对外公开基址,须与 OIDC_ISSUER 一致 +OIDC_PUBLIC_BASE_URL = 'https://auth.example.com' +# 是否强制所有授权请求使用 PKCE +OIDC_REQUIRE_PKCE = true +# 允许的 PKCE 方法,多个值使用逗号分隔(当前仅支持 S256) +OIDC_PKCE_METHODS = 'S256' +# 校验令牌时间戳时允许的时钟偏差(秒) +OIDC_ALLOWED_CLOCK_SKEW_SECONDS = 60 +# 授权码有效期(秒) +OIDC_AUTHORIZATION_CODE_TTL_SECONDS = 90 +# 登录、同意等交互状态有效期(秒) +OIDC_INTERACTION_TTL_SECONDS = 300 +# ID Token 有效期(秒) +OIDC_ID_TOKEN_TTL_SECONDS = 300 +# Access Token 默认有效期(秒) +OIDC_ACCESS_TOKEN_TTL_SECONDS = 600 +# Access Token 允许的最大有效期(秒) +OIDC_MAX_ACCESS_TOKEN_TTL_SECONDS = 1800 +# Refresh Token 闲置有效期(秒) +OIDC_REFRESH_TOKEN_IDLE_SECONDS = 604800 +# Refresh Token 绝对有效期(秒) +OIDC_REFRESH_TOKEN_ABSOLUTE_SECONDS = 2592000 +# SSO 会话闲置超时时间(秒) +OIDC_SSO_IDLE_SECONDS = 1800 +# SSO 会话绝对有效期(秒) +OIDC_SSO_ABSOLUTE_SECONDS = 28800 +# SSO“记住我”模式的绝对有效期(秒) +OIDC_SSO_REMEMBER_ABSOLUTE_SECONDS = 604800 +# OIDC 签名算法 +OIDC_SIGNING_ALGORITHM = 'RS256' +# 签名密钥来源,可选 file、kms、hsm +OIDC_SIGNING_KEY_SOURCE = 'file' +# 签名私钥文件路径(file 模式使用) +OIDC_SIGNING_PRIVATE_KEY_PATH = '' +# 签名私钥加密密钥 +OIDC_SIGNING_KEY_ENCRYPTION_KEY = '' +# 当前启用的签名密钥 KID,留空由数据库 active 密钥决定 +OIDC_ACTIVE_KID = '' +# 密钥轮换期间旧密钥保留时间(秒) +OIDC_KEY_ROTATION_OVERLAP_SECONDS = 86400 +# 不透明 Token 哈希 Pepper(至少 32 bytes,生产环境必须配置) +OIDC_TOKEN_HASH_PEPPER = '' +# 允许的 CORS Origin 列表,多个值使用逗号分隔 +OIDC_CORS_ALLOWED_ORIGINS = '' +# 登录交互页面 URL +OIDC_INTERACTION_LOGIN_URL = 'https://auth.example.com/auth-center/login' +# 授权同意页面 URL +OIDC_INTERACTION_CONSENT_URL = 'https://auth.example.com/auth-center/consent' +# 交互错误页面 URL +OIDC_INTERACTION_ERROR_URL = 'https://auth.example.com/auth-center/error' +# SSO Cookie 名称(必须使用 __Host- 前缀) +OIDC_SSO_COOKIE_NAME = '__Host-ruoyi-sso' +# SSO Cookie 是否启用 Secure 属性 +OIDC_SSO_COOKIE_SECURE = true +# SSO Cookie SameSite 策略(当前必须为 lax) +OIDC_SSO_COOKIE_SAMESITE = 'lax' +# SSO Cookie Domain(__Host- Cookie 必须留空) +OIDC_SSO_COOKIE_DOMAIN = '' +# 是否启用 Legacy 认证隔离 +OIDC_LEGACY_AUTH_ISOLATION_ENABLED = true +# 审计日志保留天数 +OIDC_AUDIT_RETENTION_DAYS = 180 +# Back-channel Logout 请求超时时间(秒) +OIDC_BACKCHANNEL_LOGOUT_TIMEOUT_SECONDS = 5 + # -------- 传输层加解密配置 -------- # 是否启用传输层加解密 TRANSPORT_CRYPTO_ENABLED = true @@ -187,7 +258,7 @@ TRANSPORT_CRYPTO_ENABLED_PATHS = '' # 强制要求传输层加密的路径列表,多个值使用逗号分隔 TRANSPORT_CRYPTO_REQUIRED_PATHS = '' # 排除传输层加密的路径列表,多个值使用逗号分隔 -TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download' +TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download,/.well-known/openid-configuration,/.well-known/oauth-authorization-server,/oauth2/authorize,/oauth2/token,/oauth2/userinfo,/oauth2/jwks,/oauth2/revoke,/oauth2/introspect,/oauth2/logout' # -------- 插件依赖安装策略配置 -------- # 策略模式:可填单一模式,或按环境映射。可选 disabled、plan_only、explicit、locked、offline diff --git a/ruoyi-fastapi-backend/.env.dockerpg b/ruoyi-fastapi-backend/.env.dockerpg index dcb6d5ec3..81aad9c7a 100644 --- a/ruoyi-fastapi-backend/.env.dockerpg +++ b/ruoyi-fastapi-backend/.env.dockerpg @@ -155,6 +155,77 @@ LOG_SERVICE_NAME = 'ruoyi-fastapi-backend' # Worker 标识(auto 自动生成) LOG_WORKER_ID = 'auto' +# -------- 统一认证中心 OIDC 配置 -------- +# 默认关闭;启用前必须配置生产密钥、Pepper 和已注册 Origin。 +# 是否启用统一认证中心 OIDC 功能 +OIDC_ENABLED = false +# OIDC 签发者地址(必须为公开的绝对 URL) +OIDC_ISSUER = 'https://auth.example.com' +# OIDC 对外公开基址,须与 OIDC_ISSUER 一致 +OIDC_PUBLIC_BASE_URL = 'https://auth.example.com' +# 是否强制所有授权请求使用 PKCE +OIDC_REQUIRE_PKCE = true +# 允许的 PKCE 方法,多个值使用逗号分隔(当前仅支持 S256) +OIDC_PKCE_METHODS = 'S256' +# 校验令牌时间戳时允许的时钟偏差(秒) +OIDC_ALLOWED_CLOCK_SKEW_SECONDS = 60 +# 授权码有效期(秒) +OIDC_AUTHORIZATION_CODE_TTL_SECONDS = 90 +# 登录、同意等交互状态有效期(秒) +OIDC_INTERACTION_TTL_SECONDS = 300 +# ID Token 有效期(秒) +OIDC_ID_TOKEN_TTL_SECONDS = 300 +# Access Token 默认有效期(秒) +OIDC_ACCESS_TOKEN_TTL_SECONDS = 600 +# Access Token 允许的最大有效期(秒) +OIDC_MAX_ACCESS_TOKEN_TTL_SECONDS = 1800 +# Refresh Token 闲置有效期(秒) +OIDC_REFRESH_TOKEN_IDLE_SECONDS = 604800 +# Refresh Token 绝对有效期(秒) +OIDC_REFRESH_TOKEN_ABSOLUTE_SECONDS = 2592000 +# SSO 会话闲置超时时间(秒) +OIDC_SSO_IDLE_SECONDS = 1800 +# SSO 会话绝对有效期(秒) +OIDC_SSO_ABSOLUTE_SECONDS = 28800 +# SSO“记住我”模式的绝对有效期(秒) +OIDC_SSO_REMEMBER_ABSOLUTE_SECONDS = 604800 +# OIDC 签名算法 +OIDC_SIGNING_ALGORITHM = 'RS256' +# 签名密钥来源,可选 file、kms、hsm +OIDC_SIGNING_KEY_SOURCE = 'file' +# 签名私钥文件路径(file 模式使用) +OIDC_SIGNING_PRIVATE_KEY_PATH = '' +# 签名私钥加密密钥 +OIDC_SIGNING_KEY_ENCRYPTION_KEY = '' +# 当前启用的签名密钥 KID,留空由数据库 active 密钥决定 +OIDC_ACTIVE_KID = '' +# 密钥轮换期间旧密钥保留时间(秒) +OIDC_KEY_ROTATION_OVERLAP_SECONDS = 86400 +# 不透明 Token 哈希 Pepper(至少 32 bytes,生产环境必须配置) +OIDC_TOKEN_HASH_PEPPER = '' +# 允许的 CORS Origin 列表,多个值使用逗号分隔 +OIDC_CORS_ALLOWED_ORIGINS = '' +# 登录交互页面 URL +OIDC_INTERACTION_LOGIN_URL = 'https://auth.example.com/auth-center/login' +# 授权同意页面 URL +OIDC_INTERACTION_CONSENT_URL = 'https://auth.example.com/auth-center/consent' +# 交互错误页面 URL +OIDC_INTERACTION_ERROR_URL = 'https://auth.example.com/auth-center/error' +# SSO Cookie 名称(必须使用 __Host- 前缀) +OIDC_SSO_COOKIE_NAME = '__Host-ruoyi-sso' +# SSO Cookie 是否启用 Secure 属性 +OIDC_SSO_COOKIE_SECURE = true +# SSO Cookie SameSite 策略(当前必须为 lax) +OIDC_SSO_COOKIE_SAMESITE = 'lax' +# SSO Cookie Domain(__Host- Cookie 必须留空) +OIDC_SSO_COOKIE_DOMAIN = '' +# 是否启用 Legacy 认证隔离 +OIDC_LEGACY_AUTH_ISOLATION_ENABLED = true +# 审计日志保留天数 +OIDC_AUDIT_RETENTION_DAYS = 180 +# Back-channel Logout 请求超时时间(秒) +OIDC_BACKCHANNEL_LOGOUT_TIMEOUT_SECONDS = 5 + # -------- 传输层加解密配置 -------- # 是否启用传输层加解密 TRANSPORT_CRYPTO_ENABLED = true @@ -187,7 +258,7 @@ TRANSPORT_CRYPTO_ENABLED_PATHS = '' # 强制要求传输层加密的路径列表,多个值使用逗号分隔 TRANSPORT_CRYPTO_REQUIRED_PATHS = '' # 排除传输层加密的路径列表,多个值使用逗号分隔 -TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download' +TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download,/.well-known/openid-configuration,/.well-known/oauth-authorization-server,/oauth2/authorize,/oauth2/token,/oauth2/userinfo,/oauth2/jwks,/oauth2/revoke,/oauth2/introspect,/oauth2/logout' # -------- 插件依赖安装策略配置 -------- # 策略模式:可填单一模式,或按环境映射。可选 disabled、plan_only、explicit、locked、offline diff --git a/ruoyi-fastapi-backend/.env.prod b/ruoyi-fastapi-backend/.env.prod index f90d51ed0..a9bef3d60 100644 --- a/ruoyi-fastapi-backend/.env.prod +++ b/ruoyi-fastapi-backend/.env.prod @@ -155,6 +155,77 @@ LOG_SERVICE_NAME = 'ruoyi-fastapi-backend' # Worker 标识(auto 自动生成) LOG_WORKER_ID = 'auto' +# -------- 统一认证中心 OIDC 配置 -------- +# 默认关闭;启用前必须配置生产密钥、Pepper 和已注册 Origin。 +# 是否启用统一认证中心 OIDC 功能 +OIDC_ENABLED = false +# OIDC 签发者地址(必须为公开的绝对 URL) +OIDC_ISSUER = 'https://auth.example.com' +# OIDC 对外公开基址,须与 OIDC_ISSUER 一致 +OIDC_PUBLIC_BASE_URL = 'https://auth.example.com' +# 是否强制所有授权请求使用 PKCE +OIDC_REQUIRE_PKCE = true +# 允许的 PKCE 方法,多个值使用逗号分隔(当前仅支持 S256) +OIDC_PKCE_METHODS = 'S256' +# 校验令牌时间戳时允许的时钟偏差(秒) +OIDC_ALLOWED_CLOCK_SKEW_SECONDS = 60 +# 授权码有效期(秒) +OIDC_AUTHORIZATION_CODE_TTL_SECONDS = 90 +# 登录、同意等交互状态有效期(秒) +OIDC_INTERACTION_TTL_SECONDS = 300 +# ID Token 有效期(秒) +OIDC_ID_TOKEN_TTL_SECONDS = 300 +# Access Token 默认有效期(秒) +OIDC_ACCESS_TOKEN_TTL_SECONDS = 600 +# Access Token 允许的最大有效期(秒) +OIDC_MAX_ACCESS_TOKEN_TTL_SECONDS = 1800 +# Refresh Token 闲置有效期(秒) +OIDC_REFRESH_TOKEN_IDLE_SECONDS = 604800 +# Refresh Token 绝对有效期(秒) +OIDC_REFRESH_TOKEN_ABSOLUTE_SECONDS = 2592000 +# SSO 会话闲置超时时间(秒) +OIDC_SSO_IDLE_SECONDS = 1800 +# SSO 会话绝对有效期(秒) +OIDC_SSO_ABSOLUTE_SECONDS = 28800 +# SSO“记住我”模式的绝对有效期(秒) +OIDC_SSO_REMEMBER_ABSOLUTE_SECONDS = 604800 +# OIDC 签名算法 +OIDC_SIGNING_ALGORITHM = 'RS256' +# 签名密钥来源,可选 file、kms、hsm +OIDC_SIGNING_KEY_SOURCE = 'file' +# 签名私钥文件路径(file 模式使用) +OIDC_SIGNING_PRIVATE_KEY_PATH = '' +# 签名私钥加密密钥 +OIDC_SIGNING_KEY_ENCRYPTION_KEY = '' +# 当前启用的签名密钥 KID,留空由数据库 active 密钥决定 +OIDC_ACTIVE_KID = '' +# 密钥轮换期间旧密钥保留时间(秒) +OIDC_KEY_ROTATION_OVERLAP_SECONDS = 86400 +# 不透明 Token 哈希 Pepper(至少 32 bytes,生产环境必须配置) +OIDC_TOKEN_HASH_PEPPER = '' +# 允许的 CORS Origin 列表,多个值使用逗号分隔 +OIDC_CORS_ALLOWED_ORIGINS = '' +# 登录交互页面 URL +OIDC_INTERACTION_LOGIN_URL = 'https://auth.example.com/auth-center/login' +# 授权同意页面 URL +OIDC_INTERACTION_CONSENT_URL = 'https://auth.example.com/auth-center/consent' +# 交互错误页面 URL +OIDC_INTERACTION_ERROR_URL = 'https://auth.example.com/auth-center/error' +# SSO Cookie 名称(必须使用 __Host- 前缀) +OIDC_SSO_COOKIE_NAME = '__Host-ruoyi-sso' +# SSO Cookie 是否启用 Secure 属性 +OIDC_SSO_COOKIE_SECURE = true +# SSO Cookie SameSite 策略(当前必须为 lax) +OIDC_SSO_COOKIE_SAMESITE = 'lax' +# SSO Cookie Domain(__Host- Cookie 必须留空) +OIDC_SSO_COOKIE_DOMAIN = '' +# 是否启用 Legacy 认证隔离 +OIDC_LEGACY_AUTH_ISOLATION_ENABLED = true +# 审计日志保留天数 +OIDC_AUDIT_RETENTION_DAYS = 180 +# Back-channel Logout 请求超时时间(秒) +OIDC_BACKCHANNEL_LOGOUT_TIMEOUT_SECONDS = 5 + # -------- 传输层加解密配置 -------- # 是否启用传输层加解密 TRANSPORT_CRYPTO_ENABLED = true @@ -187,7 +258,7 @@ TRANSPORT_CRYPTO_ENABLED_PATHS = '' # 强制要求传输层加密的路径列表,多个值使用逗号分隔 TRANSPORT_CRYPTO_REQUIRED_PATHS = '' # 排除传输层加密的路径列表,多个值使用逗号分隔 -TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download' +TRANSPORT_CRYPTO_EXCLUDE_PATHS = '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,/common/files,/system/file/download,/.well-known/openid-configuration,/.well-known/oauth-authorization-server,/oauth2/authorize,/oauth2/token,/oauth2/userinfo,/oauth2/jwks,/oauth2/revoke,/oauth2/introspect,/oauth2/logout' # -------- 插件依赖安装策略配置 -------- # 策略模式:可填单一模式,或按环境映射。可选 disabled、plan_only、explicit、locked、offline diff --git a/ruoyi-fastapi-backend/cli/core/app_builder.py b/ruoyi-fastapi-backend/cli/core/app_builder.py index a8d366d00..d03c47794 100644 --- a/ruoyi-fastapi-backend/cli/core/app_builder.py +++ b/ruoyi-fastapi-backend/cli/core/app_builder.py @@ -349,6 +349,7 @@ def build(self) -> typer.Typer: 'job': 'cli.groups.job', 'config': 'cli.groups.config', 'crypto': 'cli.groups.crypto', + 'oidc': 'cli.groups.oidc', 'gen': 'cli.groups.gen', 'dev': 'cli.groups.dev', 'plugin': 'cli.groups.plugin', diff --git a/ruoyi-fastapi-backend/cli/groups/app/controller.py b/ruoyi-fastapi-backend/cli/groups/app/controller.py index 17f658259..7e3df5ae6 100644 --- a/ruoyi-fastapi-backend/cli/groups/app/controller.py +++ b/ruoyi-fastapi-backend/cli/groups/app/controller.py @@ -13,6 +13,7 @@ from cli.runtime.app import APP_RUNTIME, AppRuntimeService from cli.runtime.crypto import CRYPTO_RUNTIME, CryptoRuntimeService from cli.runtime.db import DATABASE_RUNTIME, DatabaseRuntimeService +from cli.runtime.oidc import OIDC_RUNTIME, OidcRuntimeCliService from cli.runtime.ops import OPERATIONS_RUNTIME, OperationsRuntimeService from .presenter import AppCommandPresenter @@ -36,6 +37,7 @@ def __init__( database_runtime: DatabaseRuntimeService | None = None, operations_runtime: OperationsRuntimeService | None = None, crypto_runtime: CryptoRuntimeService | None = None, + oidc_runtime: OidcRuntimeCliService | None = None, bootstrap_service: AppBootstrapService | None = None, ) -> None: """ @@ -48,6 +50,7 @@ def __init__( :param database_runtime: 数据库运行时服务 :param operations_runtime: 运维运行时服务 :param crypto_runtime: 传输加密运行时服务 + :param oidc_runtime: OIDC 就绪检查运行时服务 :param bootstrap_service: 应用引导服务 :return: None """ @@ -58,6 +61,7 @@ def __init__( self.database_runtime = database_runtime or DATABASE_RUNTIME self.operations_runtime = operations_runtime or OPERATIONS_RUNTIME self.crypto_runtime = crypto_runtime or CRYPTO_RUNTIME + self.oidc_runtime = oidc_runtime or OIDC_RUNTIME self.bootstrap_service = bootstrap_service or APP_BOOTSTRAP def run_app(self, env: str) -> None: @@ -81,13 +85,15 @@ def doctor(self, env: str, output: str) -> None: db_status = self.execution_service.run_async(self.database_runtime.ping_database()) redis_status = self.execution_service.run_async(self.operations_runtime.ping_redis()) crypto_status = self.crypto_runtime.validate_crypto_config() + oidc_status = self.execution_service.run_async(self.oidc_runtime.check_readiness()) payload = { 'env': ctx.env, 'database': db_status, 'redis': redis_status, 'crypto': crypto_status, + 'oidc': oidc_status, } - payload['ok'] = all(item.get('ok', False) for item in (db_status, redis_status, crypto_status)) + payload['ok'] = all(item.get('ok', False) for item in (db_status, redis_status, crypto_status, oidc_status)) exit_code = SUCCESS if payload['ok'] else DEPENDENCY_ERROR self.execution_service.complete_payload_with_text( ctx, diff --git a/ruoyi-fastapi-backend/cli/groups/app/presenter.py b/ruoyi-fastapi-backend/cli/groups/app/presenter.py index 7a6cb7999..8e50e02cf 100644 --- a/ruoyi-fastapi-backend/cli/groups/app/presenter.py +++ b/ruoyi-fastapi-backend/cli/groups/app/presenter.py @@ -132,6 +132,7 @@ def build_doctor_text(self, payload: dict[str, Any]) -> str: self._build_check_status_line('database', payload.get('database')), self._build_check_status_line('redis', payload.get('redis')), self._build_check_status_line('crypto', payload.get('crypto')), + self._build_check_status_line('oidc', payload.get('oidc')), ] ) diff --git a/ruoyi-fastapi-backend/cli/groups/oidc/__init__.py b/ruoyi-fastapi-backend/cli/groups/oidc/__init__.py new file mode 100644 index 000000000..0a2282e59 --- /dev/null +++ b/ruoyi-fastapi-backend/cli/groups/oidc/__init__.py @@ -0,0 +1,3 @@ +from .command import app + +__all__ = ['app'] diff --git a/ruoyi-fastapi-backend/cli/groups/oidc/command.py b/ruoyi-fastapi-backend/cli/groups/oidc/command.py new file mode 100644 index 000000000..99c94fbf7 --- /dev/null +++ b/ruoyi-fastapi-backend/cli/groups/oidc/command.py @@ -0,0 +1,53 @@ +from typing import Annotated + +import typer + +from cli.context import AllowProdOption, DryRunOption, EnvOption, OutputOption, YesOption + +from .controller import OidcCommandController + +app = typer.Typer( + help='OIDC 认证中心部署与运维命令', + no_args_is_help=True, + context_settings={'help_option_names': ['-h', '--help']}, +) +key_app = typer.Typer( + help='OIDC 签名密钥命令', + no_args_is_help=True, + context_settings={'help_option_names': ['-h', '--help']}, +) +app.add_typer(key_app, name='key') +_OIDC_COMMAND_CONTROLLER = OidcCommandController() + + +@key_app.command('bootstrap', help='幂等创建并激活首把 OIDC 签名密钥') +def bootstrap_key( + env: EnvOption = 'dev', + output: OutputOption = 'text', + allow_prod: AllowProdOption = False, + yes: YesOption = False, + dry_run: DryRunOption = False, + kid: Annotated[str, typer.Option('--kid', help='签名密钥公开编号')] = 'bootstrap-primary', + actor: Annotated[str, typer.Option('--actor', help='写入审计记录的部署操作者')] = 'system:deployment', +) -> None: + """ + 初始化首把 OIDC 签名密钥 + + :param env: 当前命令运行环境 + :param output: 输出格式 + :param allow_prod: 是否允许生产环境危险命令 + :param yes: 是否跳过确认 + :param dry_run: 是否只执行预演 + :param kid: 首把签名密钥标识 + :param actor: 部署操作人标识 + :return: None + """ + _OIDC_COMMAND_CONTROLLER.bootstrap_key( + env, + output, + allow_prod, + yes, + dry_run, + kid=kid, + actor=actor, + ) diff --git a/ruoyi-fastapi-backend/cli/groups/oidc/controller.py b/ruoyi-fastapi-backend/cli/groups/oidc/controller.py new file mode 100644 index 000000000..4cd554767 --- /dev/null +++ b/ruoyi-fastapi-backend/cli/groups/oidc/controller.py @@ -0,0 +1,67 @@ +from cli.core import DEFAULT_CORE_SERVICES, CliContextFactory, CliExecutionService +from cli.runtime.oidc import OIDC_RUNTIME, OidcRuntimeCliService + + +class OidcCommandController: + """ + OIDC 命令控制器 + + 该控制器负责组织 OIDC 命令组的上下文准备、危险操作保护、 + runtime 调用和结果输出。 + """ + + def __init__( + self, + *, + context_factory: CliContextFactory | None = None, + execution_service: CliExecutionService | None = None, + runtime_service: OidcRuntimeCliService | None = None, + ) -> None: + """ + 初始化 OIDC 命令控制器 + + :param context_factory: CLI 上下文工厂 + :param execution_service: CLI 执行服务 + :param runtime_service: OIDC CLI 运行时服务 + :return: None + """ + self.context_factory = context_factory or DEFAULT_CORE_SERVICES.context_factory + self.execution_service = execution_service or DEFAULT_CORE_SERVICES.execution_service + self.runtime_service = runtime_service or OIDC_RUNTIME + + def bootstrap_key( + self, + env: str, + output: str, + allow_prod: bool, + yes: bool, + dry_run: bool, + *, + kid: str, + actor: str, + ) -> None: + """ + 执行幂等首把 OIDC 签名密钥初始化 + + :param env: 当前命令运行环境 + :param output: 输出格式 + :param allow_prod: 是否允许生产环境危险命令 + :param yes: 是否跳过确认 + :param dry_run: 是否只执行预演 + :param kid: 首把签名密钥标识 + :param actor: 部署操作人标识 + :return: None + """ + ctx = self.context_factory.build_dangerous( + env, + output, + allow_prod, + yes, + dry_run, + command_name='oidc key bootstrap', + ) + payload = self.execution_service.run_async( + self.runtime_service.bootstrap_signing_key(kid, actor=actor, dry_run=dry_run) + ) + payload['env'] = ctx.env + self.execution_service.complete_payload(ctx, payload) diff --git a/ruoyi-fastapi-backend/cli/guards.py b/ruoyi-fastapi-backend/cli/guards.py index a32c5317f..cdc3976ac 100644 --- a/ruoyi-fastapi-backend/cli/guards.py +++ b/ruoyi-fastapi-backend/cli/guards.py @@ -196,6 +196,9 @@ def guard(self, ctx: CliContext, *, rule: DangerousCommandRule) -> CommandResult supports_dry_run=False, ), 'crypto rotate': DangerousCommandRule(command_name='crypto rotate', risk_level='high', supports_dry_run=True), + 'oidc key bootstrap': DangerousCommandRule( + command_name='oidc key bootstrap', risk_level='high', supports_dry_run=True + ), 'job run-once': DangerousCommandRule(command_name='job run-once', risk_level='normal', supports_dry_run=False), 'job pause': DangerousCommandRule(command_name='job pause', risk_level='normal', supports_dry_run=False), 'job resume': DangerousCommandRule(command_name='job resume', risk_level='normal', supports_dry_run=False), diff --git a/ruoyi-fastapi-backend/cli/runtime/oidc/__init__.py b/ruoyi-fastapi-backend/cli/runtime/oidc/__init__.py new file mode 100644 index 000000000..222d3edf1 --- /dev/null +++ b/ruoyi-fastapi-backend/cli/runtime/oidc/__init__.py @@ -0,0 +1,3 @@ +from .service import OIDC_RUNTIME, OidcRuntimeCliService + +__all__ = ['OIDC_RUNTIME', 'OidcRuntimeCliService'] diff --git a/ruoyi-fastapi-backend/cli/runtime/oidc/gateway.py b/ruoyi-fastapi-backend/cli/runtime/oidc/gateway.py new file mode 100644 index 000000000..da82c3473 --- /dev/null +++ b/ruoyi-fastapi-backend/cli/runtime/oidc/gateway.py @@ -0,0 +1,56 @@ +from importlib import import_module +from typing import Any + + +class OidcInfrastructureGateway: + """ + OIDC 基础设施网关 + + 该对象负责延迟加载 OIDC CLI 所需的数据库、Redis、运行时检查 + 与签名密钥管理依赖。 + """ + + @staticmethod + def get_data_source_registry() -> Any: + """ + 获取数据源注册表 + + :return: 数据源注册表 + """ + return import_module('config.database').DataSourceRegistry + + @staticmethod + def get_oidc_config() -> Any: + """ + 获取 OIDC 配置对象 + + :return: OIDC 配置对象 + """ + return import_module('config.env').OidcConfig + + @staticmethod + def get_redis_util() -> Any: + """ + 获取 Redis 工具类 + + :return: Redis 工具类 + """ + return import_module('config.get_redis').RedisUtil + + @staticmethod + def get_runtime_service() -> Any: + """ + 获取 OIDC 运行时服务 + + :return: OIDC 运行时服务类 + """ + return import_module('module_identity.service.runtime_service').OidcRuntimeService + + @staticmethod + def get_key_management_service() -> Any: + """ + 获取 OIDC 签名密钥管理服务 + + :return: OIDC 签名密钥管理服务类 + """ + return import_module('module_identity.service.key_service').OidcKeyManagementService diff --git a/ruoyi-fastapi-backend/cli/runtime/oidc/service.py b/ruoyi-fastapi-backend/cli/runtime/oidc/service.py new file mode 100644 index 000000000..09bc2cc9e --- /dev/null +++ b/ruoyi-fastapi-backend/cli/runtime/oidc/service.py @@ -0,0 +1,133 @@ +from datetime import datetime, timezone +from typing import Any + +from cli.exit_codes import RUNTIME_ERROR + +from .gateway import OidcInfrastructureGateway + + +class OidcRuntimeCliService: + """ + OIDC CLI 运行时服务 + + 该服务对外统一暴露 OIDC 就绪检查和首把签名密钥初始化入口。 + + :param infrastructure_gateway: OIDC 基础设施网关 + """ + + def __init__(self, infrastructure_gateway: OidcInfrastructureGateway | None = None) -> None: + """ + 初始化 OIDC CLI 运行时服务 + + :param infrastructure_gateway: OIDC 基础设施网关 + :return: None + """ + self.infrastructure_gateway = infrastructure_gateway or OidcInfrastructureGateway() + + async def check_readiness(self) -> dict[str, Any]: + """ + 检查部署配置中的 OIDC 是否已具备签名能力 + + :return: OIDC 就绪检查结果 + """ + data_source_registry = self.infrastructure_gateway.get_data_source_registry() + oidc_config = self.infrastructure_gateway.get_oidc_config() + runtime_service = self.infrastructure_gateway.get_runtime_service() + + if not oidc_config.oidc_enabled: + return { + 'ok': True, + 'message': 'OIDC 未启用', + 'enabled': False, + 'ready': False, + 'reason': 'disabled', + } + try: + await data_source_registry.initialize(log_enabled=False) + async with data_source_registry.session() as db: + readiness = await runtime_service.inspect_readiness(db) + result = { + 'ok': readiness.ready, + 'message': 'OIDC 已就绪' if readiness.ready else 'OIDC 尚未初始化可用签名密钥', + 'enabled': readiness.enabled, + 'ready': readiness.ready, + 'reason': readiness.reason, + } + if not readiness.ready: + result['error'] = readiness.reason + return result + except Exception as exc: + return { + 'ok': False, + 'message': 'OIDC 就绪检查失败', + 'enabled': True, + 'ready': False, + 'reason': 'check_failed', + 'error': str(exc), + } + finally: + await data_source_registry.dispose_all() + + async def bootstrap_signing_key( + self, + kid: str, + *, + actor: str, + dry_run: bool = False, + ) -> dict[str, Any]: + """ + 幂等创建并激活首把 OIDC 签名密钥 + + :param kid: 首把签名密钥标识 + :param actor: 部署操作人标识 + :param dry_run: 是否只生成预演结果 + :return: OIDC 签名密钥初始化结果 + """ + if dry_run: + return { + 'ok': True, + 'message': 'OIDC 签名密钥初始化预演完成,未写入数据库', + 'kid': kid, + 'created': False, + 'active': False, + 'dryRun': True, + } + + data_source_registry = self.infrastructure_gateway.get_data_source_registry() + redis_util = self.infrastructure_gateway.get_redis_util() + key_management_service = self.infrastructure_gateway.get_key_management_service() + + redis = None + try: + await data_source_registry.initialize(log_enabled=False) + redis = await redis_util.create_redis_pool(log_enabled=False) + async with data_source_registry.session() as db: + row, created = await key_management_service.bootstrap( + db, + kid=kid, + actor=actor, + redis=redis, + now=datetime.now(timezone.utc), + ) + return { + 'ok': True, + 'message': 'OIDC 首把签名密钥已初始化' if created else 'OIDC 已存在可用签名密钥', + 'kid': row['kid'], + 'created': created, + 'active': row['status'] == 'active', + 'dryRun': False, + } + except Exception as exc: + return { + 'ok': False, + 'message': 'OIDC 签名密钥初始化失败', + 'error': str(exc), + 'exit_code': RUNTIME_ERROR, + } + finally: + if redis is not None: + await redis.close() + await data_source_registry.dispose_all() + + +OIDC_RUNTIME = OidcRuntimeCliService() diff --git a/ruoyi-fastapi-backend/cli/tui/actions/builders.py b/ruoyi-fastapi-backend/cli/tui/actions/builders.py index 708e80bd8..c56f6f24b 100644 --- a/ruoyi-fastapi-backend/cli/tui/actions/builders.py +++ b/ruoyi-fastapi-backend/cli/tui/actions/builders.py @@ -1,7 +1,4 @@ -from __future__ import annotations - from dataclasses import dataclass -from typing import TYPE_CHECKING from cli.tui.actions.models import ( ActionParameterBuilder, @@ -9,11 +6,9 @@ ActionTextBuilder, TuiActionSpec, ) +from cli.tui.adapters.models import BrowserRecordSnapshot from cli.tui.copy import TUI_COPY -if TYPE_CHECKING: - from cli.tui.adapters.models import BrowserRecordSnapshot - @dataclass(frozen=True) class TuiActionSpecFactory: diff --git a/ruoyi-fastapi-backend/cli/tui/actions/registry.py b/ruoyi-fastapi-backend/cli/tui/actions/registry.py index 90534e2cf..2f2f396e4 100644 --- a/ruoyi-fastapi-backend/cli/tui/actions/registry.py +++ b/ruoyi-fastapi-backend/cli/tui/actions/registry.py @@ -1,12 +1,44 @@ -from __future__ import annotations - from dataclasses import dataclass -from typing import TYPE_CHECKING -if TYPE_CHECKING: - from cli.tui.actions.builders import TuiActionSpecFactory, TuiActionTemplate - from cli.tui.actions.models import ActionSlot, TuiActionSpec - from cli.tui.adapters.models import BrowserRecordSnapshot +from cli.tui.actions.builders import TuiActionSpecFactory, TuiActionTemplate +from cli.tui.actions.models import ActionSlot, TuiActionSpec +from cli.tui.adapters.models import BrowserRecordSnapshot + + +@dataclass(frozen=True) +class TuiActionSlotResolver: + """ + TUI 动作槽位解析器。 + + 该对象将一个页面支持的槽位动作集中定义为映射关系,避免继续在 + `resolve_*_action` 中堆叠条件分支。 + + :param slot_templates: 槽位到动作模板的映射 + :param spec_factory: 动作规格构建器 + """ + + slot_templates: dict[ActionSlot, TuiActionTemplate] + spec_factory: TuiActionSpecFactory + + def resolve( + self, + *, + slot: ActionSlot, + record: BrowserRecordSnapshot | None, + env: str, + ) -> TuiActionSpec | None: + """ + 解析指定槽位的动作定义。 + + :param slot: 动作槽位 + :param record: 当前选中记录 + :param env: 当前运行环境 + :return: 动作定义 + """ + template = self.slot_templates.get(slot) + if template is None: + return None + return template.build(record=record, env=env, spec_factory=self.spec_factory) @dataclass(frozen=True) @@ -65,39 +97,3 @@ def resolve_detail_action( if resolver is None: return None return resolver.resolve(slot=slot, record=None, env=env) - - -@dataclass(frozen=True) -class TuiActionSlotResolver: - """ - TUI 动作槽位解析器。 - - 该对象将一个页面支持的槽位动作集中定义为映射关系,避免继续在 - `resolve_*_action` 中堆叠条件分支。 - - :param slot_templates: 槽位到动作模板的映射 - :param spec_factory: 动作规格构建器 - """ - - slot_templates: dict[ActionSlot, TuiActionTemplate] - spec_factory: TuiActionSpecFactory - - def resolve( - self, - *, - slot: ActionSlot, - record: BrowserRecordSnapshot | None, - env: str, - ) -> TuiActionSpec | None: - """ - 解析指定槽位的动作定义。 - - :param slot: 动作槽位 - :param record: 当前选中记录 - :param env: 当前运行环境 - :return: 动作定义 - """ - template = self.slot_templates.get(slot) - if template is None: - return None - return template.build(record=record, env=env, spec_factory=self.spec_factory) diff --git a/ruoyi-fastapi-backend/common/constant.py b/ruoyi-fastapi-backend/common/constant.py index aee6467f8..1c6c8bf56 100644 --- a/ruoyi-fastapi-backend/common/constant.py +++ b/ruoyi-fastapi-backend/common/constant.py @@ -135,6 +135,17 @@ class JobConstant: class LockConstant: """ 分布式锁常量 + + APP_STARTUP_LOCK_KEY: 应用启动分布式锁键 + LOCK_EXPIRE_SECONDS: 通用分布式锁过期时间(秒) + LOCK_RENEWAL_INTERVAL: 分布式锁续期间隔(秒) + PLUGIN_STARTUP_READY_KEY: 插件启动就绪状态键 + PLUGIN_STARTUP_READY_EXPIRE_SECONDS: 插件启动就绪状态过期时间(秒) + PLUGIN_STARTUP_FAILED_EXPIRE_SECONDS: 插件启动失败状态过期时间(秒) + PLUGIN_STARTUP_READY_WAIT_TIMEOUT_SECONDS: 等待插件启动就绪的超时时间(秒) + PLUGIN_STARTUP_READY_WAIT_INTERVAL_SECONDS: 等待插件启动就绪的轮询间隔(秒) + PLUGIN_LIFECYCLE_LOCK_PREFIX: 插件生命周期锁键前缀 + PLUGIN_LIFECYCLE_LOCK_EXPIRE_SECONDS: 插件生命周期锁过期时间(秒) """ APP_STARTUP_LOCK_KEY = 'app:startup:lock' @@ -149,9 +160,118 @@ class LockConstant: PLUGIN_LIFECYCLE_LOCK_EXPIRE_SECONDS = 1800 +class OidcConstant: + """ + 统一认证中心命名空间常量。 + + OIDC 状态与现有 ``access_token:*`` Legacy 会话严格隔离;Redis 具体键由 + ``module_identity.redis_keys.OidcRedisKey`` 生成,这里只保存稳定的命名空间。 + + REDIS_PREFIX: OIDC Redis 键前缀 + INTERACTION_NAMESPACE: 认证交互状态命名空间 + AUTHORIZATION_CODE_NAMESPACE: 授权码命名空间 + SSO_SESSION_NAMESPACE: 单点登录会话命名空间 + USER_SESSIONS_NAMESPACE: 用户会话集合命名空间 + SSO_COOKIE_NAMESPACE: 单点登录 Cookie 命名空间 + REVOKED_JTI_NAMESPACE: 已撤销令牌 JTI 命名空间 + RATE_LIMIT_NAMESPACE: OIDC 限流命名空间 + EVENT_NAMESPACE: OIDC 审计事件命名空间 + CACHE_NAMESPACE: OIDC 缓存命名空间 + AUTHORIZE_RATE_LIMIT: 授权端点限流命名空间 + LOGIN_RATE_LIMIT: 登录端点限流命名空间 + TOKEN_RATE_LIMIT: 令牌端点限流命名空间 + INTROSPECT_RATE_LIMIT: 令牌自省端点限流命名空间 + """ + + REDIS_PREFIX = 'oidc' + INTERACTION_NAMESPACE = 'oidc:interaction' + AUTHORIZATION_CODE_NAMESPACE = 'oidc:authorization_code' + SSO_SESSION_NAMESPACE = 'oidc:sso_session' + USER_SESSIONS_NAMESPACE = 'oidc:user_sessions' + SSO_COOKIE_NAMESPACE = 'oidc:sso_cookie' + REVOKED_JTI_NAMESPACE = 'oidc:revoked_jti' + RATE_LIMIT_NAMESPACE = 'oidc:rate_limit' + EVENT_NAMESPACE = 'oidc:event' + CACHE_NAMESPACE = 'oidc:cache' + + AUTHORIZE_RATE_LIMIT = 'oidc:rate_limit:authorize' + LOGIN_RATE_LIMIT = 'oidc:rate_limit:login' + TOKEN_RATE_LIMIT = 'oidc:rate_limit:token' + INTROSPECT_RATE_LIMIT = 'oidc:rate_limit:introspect' + + +class OidcAuditEvent: + """ + 统一认证中心首批审计事件名称。 + + 事件值用于审计持久化和风险检索,禁止写入 Token、Secret 或 Cookie 明文。 + + AUTHORIZE_REQUESTED: 收到授权请求 + AUTHORIZE_SUCCEEDED: 授权请求成功 + AUTHORIZE_DENIED: 授权请求被拒绝 + LOGIN_SUCCEEDED: 登录成功 + LOGIN_FAILED: 登录失败 + CONSENT_GRANTED: 用户授予同意 + CONSENT_REVOKED: 用户撤销同意 + AUTHORIZATION_CODE_REUSED: 授权码重复使用 + TOKEN_ISSUED: 令牌签发成功 + TOKEN_FAILED: 令牌处理失败 + REFRESH_ROTATED: 刷新令牌轮换成功 + REFRESH_REUSE_DETECTED: 检测到刷新令牌重复使用 + TOKEN_REVOKED: 令牌已撤销 + SESSION_REVOKED: 会话已撤销 + CLIENT_ACCESS_BLOCKED: 用户应用访问已禁止 + CLIENT_ACCESS_ALLOWED: 用户应用访问已允许 + GRANT_REVOKED: 授权授予已撤销 + RESOURCE_POLICY_CHANGED: 资源策略已变更 + SCOPE_POLICY_CHANGED: Scope 策略已变更 + CLIENT_CREATED: OAuth 客户端已创建 + CLIENT_DISABLED: OAuth 客户端已停用 + CLIENT_SECRET_ROTATED: OAuth 客户端密钥已轮换 + SIGNING_KEY_ROTATED: OIDC 签名密钥已轮换 + BACKCHANNEL_LOGOUT_SUCCEEDED: Back-Channel Logout 成功 + BACKCHANNEL_LOGOUT_FAILED: Back-Channel Logout 失败 + IDENTITY_SUBJECT_MISSING: 身份 Subject 缺失 + SECURITY_VERSION_CHANGED: 安全版本已变更 + INVALID_CLIENT: 无效客户端请求 + """ + + AUTHORIZE_REQUESTED = 'authorize_requested' + AUTHORIZE_SUCCEEDED = 'authorize_succeeded' + AUTHORIZE_DENIED = 'authorize_denied' + LOGIN_SUCCEEDED = 'login_succeeded' + LOGIN_FAILED = 'login_failed' + CONSENT_GRANTED = 'consent_granted' + CONSENT_REVOKED = 'consent_revoked' + AUTHORIZATION_CODE_REUSED = 'authorization_code_reused' + TOKEN_ISSUED = 'token_issued' + TOKEN_FAILED = 'token_failed' + REFRESH_ROTATED = 'refresh_rotated' + REFRESH_REUSE_DETECTED = 'refresh_reuse_detected' + TOKEN_REVOKED = 'token_revoked' + SESSION_REVOKED = 'session_revoked' + GRANT_REVOKED = 'grant_revoked' + CLIENT_ACCESS_BLOCKED = 'client_access_blocked' + CLIENT_ACCESS_ALLOWED = 'client_access_allowed' + RESOURCE_POLICY_CHANGED = 'resource_policy_changed' + SCOPE_POLICY_CHANGED = 'scope_policy_changed' + CLIENT_CREATED = 'client_created' + CLIENT_DISABLED = 'client_disabled' + CLIENT_SECRET_ROTATED = 'client_secret_rotated' + SIGNING_KEY_ROTATED = 'signing_key_rotated' + BACKCHANNEL_LOGOUT_SUCCEEDED = 'backchannel_logout_succeeded' + BACKCHANNEL_LOGOUT_FAILED = 'backchannel_logout_failed' + IDENTITY_SUBJECT_MISSING = 'identity_subject_missing' + SECURITY_VERSION_CHANGED = 'security_version_changed' + INVALID_CLIENT = 'invalid_client' + + class PluginRuntimeConstant: """ 插件运行时常量。 + + PLUGIN_HOOK_TIMEOUT_SECONDS: 插件钩子执行超时时间(秒) + PLUGIN_HEALTH_TIMEOUT_SECONDS: 插件健康检查超时时间(秒) """ PLUGIN_HOOK_TIMEOUT_SECONDS = 30 @@ -172,9 +292,49 @@ class ApiNamespace: LOGIN_USER_ROUTERS: 登录用户路由接口命名空间 CAPTCHA_IMAGE: 图片验证码接口命名空间 COMMON_UPLOAD: 通用上传接口命名空间 + COMMON_PRIVATE_UPLOAD: 通用私有文件上传接口命名空间 + COMMON_FILE_DOWNLOAD: 通用文件下载接口命名空间 TRANSPORT_CRYPTO_PUBLIC_KEY: 传输层加密公钥接口命名空间 TRANSPORT_CRYPTO_FRONTEND_CONFIG: 传输层加密前端配置接口命名空间 + OIDC_AUTHORIZE: OIDC 授权接口命名空间 + OIDC_LOGIN: OIDC 登录接口命名空间 + OIDC_TOKEN: OIDC 令牌接口命名空间 + OIDC_USERINFO: OIDC 用户信息接口命名空间 + OIDC_INTROSPECT: OIDC 令牌自省接口命名空间 + OIDC_REVOKE: OIDC 令牌撤销接口命名空间 + OIDC_LOGOUT: OIDC 登出接口命名空间 + OIDC_INTERACTION: OIDC 认证交互接口命名空间 + OIDC_INTERACTION_LOGIN: OIDC 认证交互登录接口命名空间 + OIDC_INTERACTION_CONSENT: OIDC 认证交互同意接口命名空间 + + SYSTEM_OAUTH_CLIENT_CREATE: OAuth 客户端创建接口命名空间 + SYSTEM_OAUTH_CLIENT_DISABLE: OAuth 客户端停用接口命名空间 + SYSTEM_OAUTH_CLIENT_STATUS: OAuth 客户端状态接口命名空间 + SYSTEM_OAUTH_CLIENT_UPDATE: OAuth 客户端更新接口命名空间 + SYSTEM_OAUTH_CLIENT_SECRET_ROTATE: OAuth 客户端密钥轮换接口命名空间 + SYSTEM_OAUTH_CLIENT_SECRET_REVOKE: OAuth 客户端密钥撤销接口命名空间 + SYSTEM_OAUTH_CLIENT_URI_ADD: OAuth 客户端回调地址新增接口命名空间 + SYSTEM_OAUTH_CLIENT_URI_REMOVE: OAuth 客户端回调地址移除接口命名空间 + SYSTEM_OAUTH_RESOURCE_CREATE: OAuth 资源创建接口命名空间 + SYSTEM_OAUTH_RESOURCE_UPDATE: OAuth 资源更新接口命名空间 + SYSTEM_OAUTH_RESOURCE_DISABLE: OAuth 资源停用接口命名空间 + SYSTEM_OAUTH_RESOURCE_STATUS: OAuth 资源状态接口命名空间 + SYSTEM_OAUTH_SCOPE_CREATE: OAuth Scope 创建接口命名空间 + SYSTEM_OAUTH_SCOPE_UPDATE: OAuth Scope 更新接口命名空间 + SYSTEM_OAUTH_SCOPE_DISABLE: OAuth Scope 停用接口命名空间 + SYSTEM_OAUTH_SCOPE_STATUS: OAuth Scope 状态接口命名空间 + SYSTEM_OAUTH_SESSION_REVOKE: OAuth 会话撤销接口命名空间 + SYSTEM_OAUTH_SESSION_USER_REVOKE: OAuth 用户会话批量撤销接口命名空间 + SYSTEM_OAUTH_GRANT_ACCESS: 用户应用访问控制 + SYSTEM_OAUTH_GRANT_REVOKE: OAuth 授权授予撤销接口命名空间 + SYSTEM_OAUTH_KEY_ROTATE: OAuth 签名密钥轮换接口命名空间 + SYSTEM_OAUTH_KEY_ACTIVATE: OAuth 签名密钥激活接口命名空间 + SYSTEM_OAUTH_KEY_RETIRE: OAuth 签名密钥退役接口命名空间 + SYSTEM_OAUTH_KEY_DELETE: OAuth 签名密钥删除接口命名空间 + MONITOR_OAUTH_AUDIT_EXPORT: OAuth 审计日志导出接口命名空间 + MONITOR_OAUTH_AUDIT_CLEANUP: OAuth 审计日志清理接口命名空间 + MONITOR_SERVER_INFO: 服务监控信息接口命名空间 MONITOR_CACHE_CLEAR_NAME: 缓存名称清理接口命名空间 MONITOR_CACHE_CLEAR_KEY: 缓存键清理接口命名空间 @@ -219,6 +379,7 @@ class ApiNamespace: SYSTEM_NOTICE_DETAIL: 通知公告详情接口命名空间 SYSTEM_FILE_DOWNLOAD: 文件管理下载接口命名空间 SYSTEM_FILE_DELETE: 文件管理删除接口命名空间 + SYSTEM_FILE_ACL: 文件访问控制接口命名空间 SYSTEM_FILE_RETENTION_POLICY: 文件业务保留策略接口命名空间 SYSTEM_FILE_RECONCILE: 文件存储对账接口命名空间 SYSTEM_FILE_TRANSFER: 文件管理转移接口命名空间 @@ -273,6 +434,44 @@ class ApiNamespace: TRANSPORT_CRYPTO_PUBLIC_KEY = 'transport-crypto:public-key' TRANSPORT_CRYPTO_FRONTEND_CONFIG = 'transport-crypto:frontend-config' + OIDC_AUTHORIZE = 'oidc:authorize' + OIDC_LOGIN = 'oidc:login' + OIDC_TOKEN = 'oidc:token' + OIDC_USERINFO = 'oidc:userinfo' + OIDC_INTROSPECT = 'oidc:introspect' + OIDC_REVOKE = 'oidc:revoke' + OIDC_LOGOUT = 'oidc:logout' + OIDC_INTERACTION = 'oidc:interaction' + OIDC_INTERACTION_LOGIN = 'oidc:interaction:login' + OIDC_INTERACTION_CONSENT = 'oidc:interaction:consent' + + SYSTEM_OAUTH_CLIENT_CREATE = 'system:oauth-client:create' + SYSTEM_OAUTH_CLIENT_DISABLE = 'system:oauth-client:disable' + SYSTEM_OAUTH_CLIENT_STATUS = 'system:oauth-client:status' + SYSTEM_OAUTH_CLIENT_UPDATE = 'system:oauth-client:update' + SYSTEM_OAUTH_CLIENT_SECRET_ROTATE = 'system:oauth-client:secret-rotate' + SYSTEM_OAUTH_CLIENT_SECRET_REVOKE = 'system:oauth-client:secret-revoke' + SYSTEM_OAUTH_CLIENT_URI_ADD = 'system:oauth-client:uri-add' + SYSTEM_OAUTH_CLIENT_URI_REMOVE = 'system:oauth-client:uri-remove' + SYSTEM_OAUTH_RESOURCE_CREATE = 'system:oauth-resource:create' + SYSTEM_OAUTH_RESOURCE_UPDATE = 'system:oauth-resource:update' + SYSTEM_OAUTH_RESOURCE_DISABLE = 'system:oauth-resource:disable' + SYSTEM_OAUTH_RESOURCE_STATUS = 'system:oauth-resource:status' + SYSTEM_OAUTH_SCOPE_CREATE = 'system:oauth-scope:create' + SYSTEM_OAUTH_SCOPE_UPDATE = 'system:oauth-scope:update' + SYSTEM_OAUTH_SCOPE_DISABLE = 'system:oauth-scope:disable' + SYSTEM_OAUTH_SCOPE_STATUS = 'system:oauth-scope:status' + SYSTEM_OAUTH_SESSION_REVOKE = 'system:oauth-session:revoke' + SYSTEM_OAUTH_SESSION_USER_REVOKE = 'system:oauth-session:user-revoke' + SYSTEM_OAUTH_GRANT_REVOKE = 'system:oauth-grant:revoke' + SYSTEM_OAUTH_GRANT_ACCESS = 'system:oauth:grant:access' + SYSTEM_OAUTH_KEY_ROTATE = 'system:oauth-key:rotate' + SYSTEM_OAUTH_KEY_ACTIVATE = 'system:oauth-key:activate' + SYSTEM_OAUTH_KEY_RETIRE = 'system:oauth-key:retire' + SYSTEM_OAUTH_KEY_DELETE = 'system:oauth-key:delete' + MONITOR_OAUTH_AUDIT_EXPORT = 'monitor:oauth-audit:export' + MONITOR_OAUTH_AUDIT_CLEANUP = 'monitor:oauth-audit:cleanup' + MONITOR_SERVER_INFO = 'monitor:server:info' MONITOR_CACHE_CLEAR_NAME = 'monitor:cache:clear-name' MONITOR_CACHE_CLEAR_KEY = 'monitor:cache:clear-key' @@ -598,6 +797,8 @@ class GenConstant: QUERY_LIKE: 模糊查询 QUERY_EQ: 相等查询 REQUIRE: 需要 + COLUMNNAME_NOT_ADD_SHOW: 页面新增时不显示字段 + COLUMNNAME_NOT_EDIT_SHOW: 页面编辑时不显示字段 DB_TO_SQLALCHEMY_TYPE_MAPPING: 数据库类型与sqlalchemy类型映射 DB_TO_PYTHON_TYPE_MAPPING: 数据库类型与python类型映射 """ diff --git a/ruoyi-fastapi-backend/config/env.py b/ruoyi-fastapi-backend/config/env.py index 48a28f4b2..d515a2cda 100644 --- a/ruoyi-fastapi-backend/config/env.py +++ b/ruoyi-fastapi-backend/config/env.py @@ -5,7 +5,8 @@ import re import secrets import sys -from typing import Annotated, Literal +from typing import Annotated, ClassVar, Literal +from urllib.parse import urlsplit, urlunsplit from dotenv import load_dotenv from pydantic import BaseModel, ConfigDict, Field, SecretStr, computed_field, field_validator, model_validator @@ -255,10 +256,265 @@ class TransportCryptoSettings(BaseSettings): transport_crypto_exclude_paths: str = ( '/openapi.json,/docs,/docs/oauth2-redirect,/redoc,' '/transport/crypto/frontend-config,/transport/crypto/public-key,/common/download,/common/download/resource,' - '/common/files,/system/file/download' + '/common/files,/system/file/download,' + '/.well-known/openid-configuration,/.well-known/oauth-authorization-server,' + '/oauth2/authorize,/oauth2/token,/oauth2/userinfo,/oauth2/jwks,' + '/oauth2/revoke,/oauth2/introspect,/oauth2/logout' ) +class OidcSettings(BaseSettings): + """ + OIDC/OAuth2 认证中心配置 + + OIDC 默认关闭,关闭时只保留可安全解析的默认值,不要求发行者、签名密钥 + 或不透明令牌 Pepper,从而保证现有 Legacy JWT 启动路径完全不变。 + """ + + oidc_enabled: bool = False + oidc_issuer: str = 'https://auth.example.com' + oidc_public_base_url: str = 'https://auth.example.com' + + oidc_require_pkce: bool = True + oidc_pkce_methods: str = 'S256' + oidc_allowed_clock_skew_seconds: int = Field(default=60, ge=0) + oidc_authorization_code_ttl_seconds: int = Field(default=90, gt=0) + oidc_interaction_ttl_seconds: int = Field(default=300, gt=0) + oidc_id_token_ttl_seconds: int = Field(default=300, gt=0) + oidc_access_token_ttl_seconds: int = Field(default=600, gt=0) + oidc_max_access_token_ttl_seconds: int = Field(default=1800, gt=0) + oidc_refresh_token_idle_seconds: int = Field(default=604800, gt=0) + oidc_refresh_token_absolute_seconds: int = Field(default=2592000, gt=0) + + oidc_sso_idle_seconds: int = Field(default=1800, gt=0) + oidc_sso_absolute_seconds: int = Field(default=28800, gt=0) + oidc_sso_remember_absolute_seconds: int = Field(default=604800, gt=0) + oidc_sso_cookie_name: str = '__Host-ruoyi-sso' + oidc_sso_cookie_secure: bool = True + oidc_sso_cookie_samesite: Literal['lax', 'strict', 'none'] = 'lax' + oidc_sso_cookie_domain: str | None = None + + oidc_signing_algorithm: Literal['RS256'] = 'RS256' + oidc_signing_key_source: Literal['file', 'kms', 'hsm'] = 'file' + oidc_signing_private_key_path: str = '' + oidc_signing_key_encryption_key: str = '' + oidc_active_kid: str = '' + oidc_key_rotation_overlap_seconds: int = Field(default=86400, gt=0) + + oidc_token_hash_pepper: str = '' + + oidc_cors_allowed_origins: str = '' + oidc_interaction_login_url: str = 'https://auth.example.com/auth-center/login' + oidc_interaction_consent_url: str = 'https://auth.example.com/auth-center/consent' + oidc_interaction_error_url: str = 'https://auth.example.com/auth-center/error' + + oidc_audit_retention_days: int = Field(default=180, gt=0) + oidc_backchannel_logout_timeout_seconds: int = Field(default=5, gt=0) + oidc_legacy_auth_isolation_enabled: bool = True + OIDC_PEPPER_MIN_BYTES: ClassVar[int] = 32 + + @staticmethod + def _normalise_url(value: str, field_name: str, *, allow_empty: bool = False) -> str: + """ + 规范化并校验认证中心地址。 + + :param value: 待校验的 URL 文本 + :param field_name: 配置字段名称 + :param allow_empty: 是否允许空值 + :return: 去除尾部斜杠且不含 query/fragment 的 URL + :raises ValueError: URL 不是绝对 HTTP(S) 地址或含有不安全部分 + """ + value = value.strip().strip('\'"') + if not value: + if allow_empty: + return '' + raise ValueError(f'{field_name} 不能为空') + parsed = urlsplit(value) + if parsed.scheme not in {'http', 'https'} or not parsed.netloc: + raise ValueError(f'{field_name} 必须是绝对 HTTP(S) URL') + if parsed.username or parsed.password: + raise ValueError(f'{field_name} 不得包含用户信息') + if parsed.query or parsed.fragment: + raise ValueError(f'{field_name} 不得包含查询参数或片段标识') + path = parsed.path.rstrip('/') + return urlunsplit((parsed.scheme.lower(), parsed.netloc.lower(), path, '', '')) + + @staticmethod + def _is_local_http(url: str) -> bool: + """ + 判断 URL 是否属于仅开发环境允许的本机 HTTP 地址。 + + :param url: 已解析的 URL + :return: 是否为 localhost、127.0.0.1 或 ::1 的 HTTP 地址 + """ + parsed = urlsplit(url) + host = (parsed.hostname or '').lower() + return parsed.scheme == 'http' and host in {'localhost', '127.0.0.1', '::1'} + + @property + def pkce_method_list(self) -> tuple[str, ...]: + """ + 返回规范化后的 PKCE 方法集合。 + + :return: 以逗号分隔配置解析出的 PKCE 方法 + """ + return tuple(item.strip() for item in self.oidc_pkce_methods.split(',') if item.strip()) + + @property + def cors_origin_list(self) -> tuple[str, ...]: + """ + 返回规范化后的、用于精确匹配的注册 Origin。 + + :return: 去除空项和尾部斜杠后的 Origin 集合 + """ + return tuple(item.strip().rstrip('/') for item in self.oidc_cors_allowed_origins.split(',') if item.strip()) + + def _validate_issuer_urls(self, app_env: str) -> None: + """ + 校验 issuer、公开基址及环境协议要求。 + + :param app_env: 当前应用环境 + :raises ValueError: issuer 不安全、含路径前缀或公开地址不一致 + """ + if self.oidc_issuer != self.oidc_public_base_url: + raise ValueError('OIDC_PUBLIC_BASE_URL 必须与 OIDC_ISSUER 一致') + parsed = urlsplit(self.oidc_issuer) + if parsed.path.rstrip('/'): + raise ValueError('OIDC_ISSUER 不得包含路径前缀') + if parsed.scheme != 'https' and not ( + app_env in {'dev', 'test', 'local'} and self._is_local_http(self.oidc_issuer) + ): + raise ValueError('OIDC_ISSUER 生产环境必须使用 HTTPS') + if urlsplit(self.oidc_public_base_url).scheme != parsed.scheme: + raise ValueError('OIDC_PUBLIC_BASE_URL 必须与签发者地址使用相同协议') + + def _validate_protocol_options(self) -> None: + """ + 校验 PKCE 和各类协议 TTL。 + + :raises ValueError: 协议选项不符合统一认证安全边界 + """ + if not self.oidc_require_pkce: + raise ValueError('OIDC_REQUIRE_PKCE 启用认证中心时必须为 true') + if self.pkce_method_list != ('S256',): + raise ValueError('OIDC_PKCE_METHODS 目前只能配置为 S256') + if self.oidc_access_token_ttl_seconds > self.oidc_max_access_token_ttl_seconds: + raise ValueError('OIDC_MAX_ACCESS_TOKEN_TTL_SECONDS 不能小于默认访问令牌有效期') + if self.oidc_refresh_token_idle_seconds > self.oidc_refresh_token_absolute_seconds: + raise ValueError('刷新令牌闲置有效期不能大于绝对有效期') + if self.oidc_sso_absolute_seconds > self.oidc_sso_remember_absolute_seconds: + raise ValueError('单点登录保持登录期限不能小于普通会话的绝对有效期') + + def _validate_secret_material(self) -> None: + """ + 校验 OIDC Pepper、签名密钥定位信息和 Legacy 隔离开关。 + + :raises ValueError: Pepper、密钥或隔离策略不符合要求 + """ + if not self.oidc_legacy_auth_isolation_enabled: + raise ValueError('OIDC_LEGACY_AUTH_ISOLATION_ENABLED 启用认证中心时必须为 true') + pepper = self.oidc_token_hash_pepper.strip() + if len(pepper.encode('utf-8')) < self.OIDC_PEPPER_MIN_BYTES: + raise ValueError('OIDC_TOKEN_HASH_PEPPER 至少需要 32 字节') + secret_values = { + os.getenv('JWT_SECRET_KEY', '').strip(), + os.getenv('TRANSPORT_CRYPTO_PRIVATE_KEY', '').strip(), + os.getenv('TRANSPORT_CRYPTO_PUBLIC_KEY', '').strip(), + } + legacy_config = globals().get('JwtConfig') + transport_config = globals().get('TransportCryptoConfig') + if legacy_config is not None: + secret_values.add(str(getattr(legacy_config, 'jwt_secret_key', '')).strip()) + if transport_config is not None: + secret_values.add(str(getattr(transport_config, 'transport_crypto_private_key', '')).strip()) + secret_values.add(str(getattr(transport_config, 'transport_crypto_public_key', '')).strip()) + secret_values.add(self.oidc_signing_key_encryption_key.strip()) + if pepper in secret_values: + raise ValueError('OIDC_TOKEN_HASH_PEPPER 不得复用原有 JWT 密钥或传输加密密钥') + # 运行时以数据库 active 密钥为唯一事实;密钥可来自数据库加密密文, + # 因此不能在配置层强制 OIDC_ACTIVE_KID 或文件路径。 + + def _validate_cookie(self) -> None: + """ + 校验认证中心 SSO Cookie 的固定安全属性。 + + :raises ValueError: Cookie 名称、Secure、SameSite 或 Domain 不符合要求 + """ + if not self.oidc_sso_cookie_name.startswith('__Host-'): + raise ValueError('OIDC_SSO_COOKIE_NAME 必须使用 __Host- 前缀') + if not self.oidc_sso_cookie_secure: + raise ValueError('OIDC_SSO_COOKIE_SECURE 启用认证中心时必须为 true') + if self.oidc_sso_cookie_domain: + raise ValueError('__Host- SSO Cookie 不得设置 Domain') + if self.oidc_sso_cookie_samesite != 'lax': + raise ValueError('OIDC_SSO_COOKIE_SAMESITE 必须为 lax') + + def _validate_interaction_urls(self) -> None: + """ + 校验认证交互 URL 与 issuer 的同源约束。 + + :raises ValueError: 交互地址不是 HTTPS 同源地址 + """ + issuer = urlsplit(self.oidc_issuer) + for field_name, value in ( + ('OIDC_INTERACTION_LOGIN_URL', self.oidc_interaction_login_url), + ('OIDC_INTERACTION_CONSENT_URL', self.oidc_interaction_consent_url), + ('OIDC_INTERACTION_ERROR_URL', self.oidc_interaction_error_url), + ): + normalised = self._normalise_url(value, field_name) + parsed = urlsplit(normalised) + if (parsed.scheme, parsed.hostname, parsed.port) != (issuer.scheme, issuer.hostname, issuer.port): + raise ValueError(f'{field_name} 的协议、主机名和端口必须与 OIDC_ISSUER 一致') + setattr(self, field_name.lower(), normalised) + + def _validate_cors_origins(self, app_env: str) -> None: + """ + 校验显式 CORS Origin 配置。 + + :param app_env: 当前应用环境 + :raises ValueError: Origin 含 userinfo、path、query、fragment 或协议不安全 + """ + for origin in self.cors_origin_list: + parsed = urlsplit(origin) + if ( + parsed.scheme not in {'http', 'https'} + or not parsed.netloc + or parsed.username + or parsed.password + or parsed.path + or parsed.query + or parsed.fragment + ): + raise ValueError('OIDC_CORS_ALLOWED_ORIGINS 只能包含协议、主机名和可选端口') + if parsed.scheme != 'https' and not (app_env in {'dev', 'test', 'local'} and self._is_local_http(origin)): + raise ValueError('生产环境 OIDC_CORS_ALLOWED_ORIGINS 必须使用 HTTPS') + + @model_validator(mode='after') + def validate_oidc_configuration(self) -> 'OidcSettings': + """ + 执行 OIDC 启动级安全校验。 + + :return: 已完成规范化和安全校验的配置对象 + """ + self.oidc_issuer = self._normalise_url(self.oidc_issuer, 'OIDC_ISSUER', allow_empty=not self.oidc_enabled) + self.oidc_public_base_url = self._normalise_url( + self.oidc_public_base_url, + 'OIDC_PUBLIC_BASE_URL', + allow_empty=not self.oidc_enabled, + ) + if not self.oidc_enabled: + return self + + app_env = os.getenv('APP_ENV', 'dev').strip().strip('\'"').lower() or 'dev' + self._validate_issuer_urls(app_env) + self._validate_protocol_options() + self._validate_secret_material() + self._validate_cookie() + self._validate_interaction_urls() + self._validate_cors_origins(app_env) + return self + + class PluginDependencyPolicySettings(BaseSettings): """ 插件依赖安装策略配置 @@ -408,6 +664,14 @@ def get_transport_crypto_config(self) -> TransportCryptoSettings: """ return TransportCryptoSettings() + def get_oidc_config(self) -> OidcSettings: + """ + 获取统一认证中心配置。 + + OIDC 配置由模型自身执行启动级校验;关闭时不会校验密钥和 Pepper。 + """ + return OidcSettings() + def get_plugin_dependency_policy_config(self) -> PluginDependencyPolicySettings: """ 获取插件依赖安装策略配置 @@ -478,6 +742,8 @@ def parse_cli_args() -> str: LogConfig = get_config.get_log_config() # 传输层加解密配置 TransportCryptoConfig = get_config.get_transport_crypto_config() +# 统一认证中心配置 +OidcConfig = get_config.get_oidc_config() # 插件依赖安装策略配置 PluginDependencyPolicyConfig = get_config.get_plugin_dependency_policy_config() # 代码生成配置 diff --git a/ruoyi-fastapi-backend/docs/unified_authentication_developer_guide.md b/ruoyi-fastapi-backend/docs/unified_authentication_developer_guide.md new file mode 100644 index 000000000..2cdcb3316 --- /dev/null +++ b/ruoyi-fastapi-backend/docs/unified_authentication_developer_guide.md @@ -0,0 +1,749 @@ +# 统一认证中心开发者使用手册 + +更新日期:2026-10-01。适用于当前 Vue3、Vue2 两个仓库的统一认证中心;两者使用相同的后端协议和接入配置。 + +本文面向平台管理员、业务应用开发者和资源服务开发者,按“启用服务 → 注册应用 → 登录 → 调用 API → 续期和退出”的顺序说明接入方法。本文的域名、客户端编号和凭据均为占位示例,使用时替换为实际值。 + +## 1. 选择接入方式 + +统一认证中心为外部应用提供 OpenID Connect(OIDC)登录和 OAuth 2.0 授权,支持以下方式: + +| 场景 | 客户端类型 | 授权方式 | 客户端认证 | +| ---------------------------------- | ---------------------------- | ---------------------------------------------------------- | -------------------------- | +| 有后端的 Web 应用、BFF | `confidential`(机密) | `authorization_code` + PKCE S256;按需增加 `refresh_token` | `client_secret_basic` | +| 纯浏览器 SPA、不能保管密钥的客户端 | `public`(公开) | `authorization_code` + PKCE S256;按需增加 `refresh_token` | `none`,表单传 `client_id` | +| 定时任务、服务间调用 | `confidential` | `client_credentials` | `client_secret_basic` | +| 资源 API 的在线令牌校验 | 独立的 `confidential` 客户端 | 调用 Introspection,并绑定到对应资源 | `client_secret_basic` | + +当前不支持密码授权、Implicit、`client_secret_post`、PKCE `plain`,也不提供动态客户端注册。应用由平台管理员预先登记。原生应用若使用公开客户端,仍须满足当前 HTTP(S) 回调注册规则,不支持自定义 URI scheme。 + +原有管理后台/App 的登录 Token、菜单权限与 OIDC 相互隔离。OAuth Access Token 不能直接替代管理后台 Token 调用 `/system/**`;业务应用也不应把认证中心登录页当成管理后台登录接口。 + +### 1.1 各方负责什么 + +| 角色 | 需要完成的工作 | +| -------------------------------- | ------------------------------------------------------------------- | +| 平台管理员 | 启用认证中心、初始化签名密钥、登记客户端/资源/Scope、管理授权和会话 | +| 业务应用(Client / RP) | 发起授权、检查回调、兑换令牌、验证用户身份、维护自己的登录会话 | +| 资源服务(Resource Server / RS) | 验证 Access Token、受众和 Scope,并执行自己的业务数据权限判断 | + +### 1.2 先理解三个标识 + +| 标识 | 示例 | 用途 | +| ------------ | -------------------------------- | -------------------------------------------- | +| `client_id` | 管理端生成的 `cli_...` | 标识一个接入应用;客户端配置中的 `clientId` | +| `resourceId` | `orders-api` | 管理端绑定资源的内部业务标识 | +| `audience` | `https://api.example.com/orders` | 令牌受众;协议请求的 `resource` 参数填写此值 | + +管理 API 的 JSON 字段使用 `camelCase`,协议参数使用 `snake_case`。例如管理端的 `redirectUris` 是列表,授权请求中的 `redirect_uri` 是其中一个精确匹配的地址。 + +## 2. 准备地址和环境 + +后续示例统一使用: + +| 配置 | 示例值 | +| --------------- | ----------------------------------------- | +| 认证中心 Issuer | `https://auth.example.com` | +| 业务应用 | `https://app.example.com` | +| 登录回调 | `https://app.example.com/oidc/callback` | +| 退出回调 | `https://app.example.com/oidc/logged-out` | +| 订单资源受众 | `https://api.example.com/orders` | + +Issuer 必须是无路径前缀、无查询串的固定公开地址。`OIDC_PUBLIC_BASE_URL` 必须与它一致;不能把 `/prod-api`、`/dev-api` 或 `/oauth2` 写进 Issuer。 + +生产环境使用 HTTPS。仅在应用环境为 `dev`、`test`、`local` 时允许 `localhost`、`127.0.0.1`、`::1` 的本机 HTTP;本地开发建议统一使用 `http://localhost`,不要在 Issuer、回调和浏览器地址之间混用 `localhost` 与 `127.0.0.1`。SSO Cookie 始终要求 `Secure`,不能通过关闭 Cookie 安全属性解决代理或浏览器问题。 + +## 3. 平台启用与检查 + +所有后端命令在 `ruoyi-fastapi-backend` 目录执行,并使用已安装项目依赖的 Python 3 解释器。命令适用于 Windows PowerShell 和 macOS Terminal,不依赖特定的 Python 环境管理工具。下面以 `dev` 环境为例,其他环境替换为对应的 `--env` 和配置文件。 + +`ruoyi` 命令在 Windows 和 macOS 上的参数一致。直接调用 Python 时,下文分别给出 Windows 的 `python` 和 macOS 的 `python3` 写法;若本机解释器命令不同,替换为实际安装项目依赖的解释器。 + +### 3.1 配置功能开关和密钥材料 + +在对应 `.env.<环境>` 中配置: + +```dotenv +OIDC_ENABLED=true +OIDC_ISSUER=https://auth.example.com +OIDC_PUBLIC_BASE_URL=https://auth.example.com +OIDC_INTERACTION_LOGIN_URL=https://auth.example.com/auth-center/login +OIDC_INTERACTION_CONSENT_URL=https://auth.example.com/auth-center/consent +OIDC_INTERACTION_ERROR_URL=https://auth.example.com/auth-center/error + +OIDC_REQUIRE_PKCE=true +OIDC_PKCE_METHODS=S256 +OIDC_SIGNING_ALGORITHM=RS256 +OIDC_TOKEN_HASH_PEPPER=<独立生成并长期保存的随机字符串> +OIDC_SIGNING_KEY_ENCRYPTION_KEY=<另一个独立生成并长期保存的随机字符串> + +OIDC_SSO_COOKIE_NAME=__Host-ruoyi-sso +OIDC_SSO_COOKIE_SECURE=true +OIDC_SSO_COOKIE_SAMESITE=lax +OIDC_SSO_COOKIE_DOMAIN= +OIDC_LEGACY_AUTH_ISOLATION_ENABLED=true +``` + +三个交互页面地址必须与 Issuer 同源,即协议、主机和端口一致。修改密码页由认证中心流程进入,不需要另设一个交互 URL 配置。 + +Pepper 和签名密钥加密材料各至少 32 个 UTF-8 字节,彼此独立,也不要复用原有 JWT 或传输加密密钥。可执行两次以下命令生成两份不同的值,再通过部署配置或密钥管理系统保存: + +Windows PowerShell: + +```shell +python -c "import secrets; print(secrets.token_urlsafe(48))" +``` + +macOS Terminal: + +```shell +python3 -c "import secrets; print(secrets.token_urlsafe(48))" +``` + +多实例使用同一套持久化配置。每次重启重新生成 Pepper 会使已有不透明凭据无法匹配;更换加密材料会影响数据库中现有签名私钥的解密。配置文件不提交真实密钥。修改环境配置后重启后端实例。 + +### 3.2 数据库、Redis 和首把签名密钥 + +先完成平台本身的数据库和 Redis 配置。新空库使用对应数据库的初始化 SQL;存量库按实际结构升级。完整初始化 SQL 含 `DROP TABLE`,不能向存量业务库直接重新导入。 + +当前启动流程可通过 `Base.metadata.create_all` 补建缺失表,但不会修改已有列或回填历史数据;当前没有认证模块专用的 Alembic 升级命令。 + +在数据库结构准备完成后启动后端,再初始化首把签名密钥: + +```shell +ruoyi app run --env dev +``` + +在另一个终端执行: + +```shell +ruoyi oidc key bootstrap --env dev --kid auth-primary --dry-run +ruoyi oidc key bootstrap --env dev --kid auth-primary --yes +ruoyi app doctor --env dev +``` + +若终端找不到 `ruoyi` 命令,可直接使用 Python 模块入口,将上面命令的 `ruoyi` 前缀替换为下表中的调用方式,其余参数保持不变: + +| 系统 | 替代前缀 | 检查命令示例 | +| ------------------ | ----------------------------- | -------------------------------------------------- | +| Windows PowerShell | `python -X utf8 -m cli.main` | `python -X utf8 -m cli.main app doctor --env dev` | +| macOS Terminal | `python3 -X utf8 -m cli.main` | `python3 -X utf8 -m cli.main app doctor --env dev` | + +生产密钥初始化还需要显式传入 `--allow-prod`,并指定实际生产环境。 + +`bootstrap` 会创建并激活首把 RSA 签名密钥,已有有效 Active 密钥时执行校验,可重复执行。`--dry-run` 不创建密钥。此流程将私钥加密存入数据库;使用这条路径不需要自行填写私钥文件路径或手工设置 `OIDC_ACTIVE_KID`。不要把客户端 Secret 当作签名私钥。 + +### 3.3 代理路由和前端开关 + +前端和协议端点需要通过正确的公开地址访问。部署代理至少保持以下关系: + +| 浏览器请求路径 | 目标 | +| --------------------------------------- | ------------------------------------------- | +| `/auth-center/**` | Vue 前端,深层路径刷新回落到 `index.html` | +| `/.well-known/**`、`/oauth2/**` | 后端,保留原始路径 | +| `/auth/interaction/**` | 后端,保留原始路径、Cookie、Origin 等请求头 | +| `/prod-api/**` 或开发环境 `/dev-api/**` | 平台普通 API,移除该前缀后转发给后端 | + +平台前端通过普通 API 前缀请求 `GET /auth/status`,例如浏览器中的 `/prod-api/auth/status`。该匿名接口只返回开关状态,数据部分为 `{"enabled": true}`,使用平台响应包装并禁止缓存。它不代表数据库、Redis 或签名密钥已经就绪。 + +当 `OIDC_ENABLED=false` 时,Vue3/Vue2 都会阻止进入认证中心登录、同意和修改密码页,展示“统一认证服务未启用”;状态请求失败时展示“暂不可用”。错误页仍可访问,普通后台登录不受影响。协议端点关闭时返回 404。前后端需同步发布,否则旧后端缺少状态接口会使新前端拒绝进入认证页。 + +### 3.4 确认服务可用 + +使用 Python 标准库检查公开端点,无需额外安装 HTTP 命令行工具。 + +Windows PowerShell: + +```shell +python -c "import urllib.request; print(urllib.request.urlopen('https://auth.example.com/.well-known/openid-configuration', timeout=10).read().decode('utf-8'))" +python -c "import urllib.request; print(urllib.request.urlopen('https://auth.example.com/oauth2/jwks', timeout=10).read().decode('utf-8'))" +``` + +macOS Terminal: + +```shell +python3 -c "import urllib.request; print(urllib.request.urlopen('https://auth.example.com/.well-known/openid-configuration', timeout=10).read().decode('utf-8'))" +python3 -c "import urllib.request; print(urllib.request.urlopen('https://auth.example.com/oauth2/jwks', timeout=10).read().decode('utf-8'))" +``` + +检查 Discovery 中的 `issuer` 和各端点是否为预期公开 HTTPS 地址,JWKS 是否包含可用 RSA 公钥,并确认 `ruoyi app doctor` 的 OIDC 检查通过。开关为真但无有效签名密钥时,Discovery/JWKS 等就绪检查可能返回 503;先检查密钥状态和加密配置。 + +## 4. 注册资源、Scope 和客户端 + +使用具有 OAuth 管理权限的后台账号操作。以下 JSON 用于说明与管理页面对应的配置,也可由受信任的管理程序调用 API。管理 API 需要原有管理后台身份和权限,不能拿业务应用的 OAuth Token 调用。 + +下文的 `/system/**` 为后端原始路径;通过前端代理调用时加上部署中的普通 API 前缀。管理结果使用平台 `{code, msg, data}` 等响应结构;OAuth 协议响应不使用这个包装。 + +### 4.1 登记业务资源 + +只有登录需求时可以暂时跳过资源和业务 Scope,只申请身份范围。若应用要访问订单 API,先在资源管理中创建以下资源,或调用 `POST /system/oauth/resource`: + +```json +{ + "resourceId": "orders-api", + "resourceName": "订单服务", + "audience": "https://api.example.com/orders", + "tokenFormat": "jwt", + "signingAlg": "RS256", + "accessTokenTtlSeconds": 600, + "allowedClaims": [] +} +``` + +`audience` 是稳定的受众标识,不需要认证中心访问该 URL。以后协议参数 `resource` 使用这个值。需要在线内省时,按第 7 节登记校验客户端,并在资源中设置它的 `introspectionClientId`。 + +### 4.2 登记业务 Scope + +在权限范围管理中创建,或调用 `POST /system/oauth/scope`: + +```json +{ + "scopeCode": "orders.read", + "scopeName": "查看订单", + "scopeType": "resource", + "resourceId": "orders-api", + "claims": [], + "consentRequired": true, + "sensitive": false, + "status": "0" +} +``` + +资源 Scope 必须绑定资源;身份 Scope 不绑定资源。标识应使用稳定的英文编码,例如 `orders.read`、`orders.write`;显示名称可用中文。`status="0"` 为启用,`"1"` 为停用。 + +系统已有以下身份 Scope,无需重复创建: + +| Scope | 用途 | +| ---------------- | --------------------------------------------- | +| `openid` | OIDC 用户登录,当前所有用户授权请求均要求携带 | +| `profile` | 基础资料,例如显示名、用户名、头像 | +| `email`、`phone` | 邮箱、手机及相应验证状态 | +| `dept` | 部门信息 | +| `roles` | 经客户端允许列表过滤后的角色信息 | +| `offline_access` | 申请刷新令牌资格 | + +最终返回字段还受用户实际资料、获准 Scope 和 Claim 策略控制。允许申请 `roles` 不等于可读取全部后台角色;客户端的 `allowedRoleKeys` 为空时不发布角色。业务服务仍需自行判断数据归属和操作权限。 + +### 4.3 登记有后端的业务应用 + +在客户端管理页面(默认 `/oauth/client`)新增应用,或调用 `POST /system/oauth/client`: + +```json +{ + "clientName": "订单工作台", + "clientType": "confidential", + "tokenEndpointAuthMethod": "client_secret_basic", + "grantTypes": ["authorization_code", "refresh_token"], + "responseTypes": ["code"], + "requirePkce": true, + "requireConsent": true, + "trustedClient": false, + "scopeCodes": ["openid", "profile", "offline_access", "orders.read"], + "resourceIds": ["orders-api"], + "allowedRoleKeys": [], + "preAuthorizedScopeCodes": [], + "redirectUris": ["https://app.example.com/oidc/callback"], + "postLogoutRedirectUris": ["https://app.example.com/oidc/logged-out"], + "backchannelLogoutUris": [], + "corsOrigins": [] +} +``` + +保存响应中的 `data.clientId`。`client_id` 由平台生成,不使用应用名称或资源编号替代。 + +创建客户端不会自动返回 Secret。机密客户端保存后,还要在行操作中“生成/轮换密钥”,或调用 `POST /system/oauth/client/{client_id}/secret`,请求体可为 `{}`。响应中 `data.clientSecret` 只显示这一次,应立即保存到应用后端的安全配置;后续详情接口无法找回明文。公开客户端不生成 Secret。 + +URI 和授权策略需要满足以下规则: + +- 登录回调与退出回调分别注册、精确匹配;不能使用通配符、URL 用户信息或 fragment。路径、端口、查询串和尾部斜杠应与实际请求保持一致。 +- 生产回调使用 HTTPS;本机开发可登记本机 HTTP。Back-Channel 地址还受更严格的公网 HTTPS 和网络地址校验,见第 10 节。 +- 所有 `authorization_code` 客户端都必须开启 PKCE S256、配置 `responseTypes: ["code"]` 和至少一个登录回调。 +- `refresh_token` 必须与 `authorization_code` 同时启用。用户还需实际获准 `offline_access`,才会收到 Refresh Token。 +- `scopeCodes`、`resourceIds` 指向已存在且启用的配置;一个请求当前只能选择一个业务资源,相关业务 Scope 应属于该资源。 +- 客户端类型创建后不能直接修改;由公开类型切换为机密类型时应另建客户端。常规接入保持显式同意,可信应用及预授权范围由管理员按实际业务配置。 + +### 4.4 纯前端 SPA 的差异 + +将客户端设为 `public`,认证方式设为 `none`,其余授权码与 PKCE 配置保留;`corsOrigins` 登记应用 Origin,例如 `https://app.example.com`,不带路径或尾部斜杠。也可由管理员配置 `OIDC_CORS_ALLOWED_ORIGINS` 的逗号分隔 Origin 列表。 + +浏览器直接调用 Token/UserInfo 端点时才需要相应跨域配置。有后端代为兑换令牌的应用通常无需这项配置。发起授权和退出使用浏览器页面跳转,不通过跨域 AJAX 获取登录页面。SPA 不保管客户端 Secret,不能使用 `client_credentials`。 + +## 5. 接入用户登录 + +外部应用从 `/oauth2/authorize` 发起流程,由认证中心创建交互并导航到登录、同意或修改密码页。不要直接拼接 `/auth-center/login`,也不要调用 `/auth/interaction/**` 模拟外部应用登录。 + +```mermaid +sequenceDiagram + participant B as 浏览器 + participant A as 业务应用 + participant I as 认证中心 + participant R as 资源 API + B->>A: 点击登录 + A-->>B: 保存 state、nonce、verifier 后跳转 + B->>I: authorize + PKCE challenge + I-->>B: 登录与同意;有有效 SSO 时按策略跳过 + I-->>B: 回调 code、state、iss + B->>A: 登录回调 + A->>I: token + code + verifier + I-->>A: Access Token、ID Token、可选 Refresh Token + A->>A: 验证身份并建立应用会话 + A->>R: Bearer Access Token + R->>R: 校验令牌、受众、Scope 与业务权限 +``` + +### 5.1 生成授权请求 + +建议用支持 OIDC Authorization Code + PKCE 的库实现完整流程。下面展示 Python 服务端的核心参数生成,方便与现有框架集成;后续示例可放在同一模块中,依赖 `httpx` 和 `PyJWT[crypto]`。 + +```python +import base64 +import hashlib +import secrets +import time +from urllib.parse import urlencode + +ISSUER = 'https://auth.example.com' +CLIENT_ID = 'cli_REPLACE_WITH_REGISTERED_ID' +REDIRECT_URI = 'https://app.example.com/oidc/callback' +RESOURCE = 'https://api.example.com/orders' + + +def begin_login(): + verifier = secrets.token_urlsafe(48) + challenge = base64.urlsafe_b64encode(hashlib.sha256(verifier.encode('ascii')).digest()).rstrip(b'=').decode('ascii') + transaction = { + 'state': secrets.token_urlsafe(32), + 'nonce': secrets.token_urlsafe(32), + 'verifier': verifier, + 'expires_at': time.time() + 300, + } + parameters = { + 'client_id': CLIENT_ID, + 'response_type': 'code', + 'response_mode': 'query', + 'redirect_uri': REDIRECT_URI, + 'scope': 'openid profile offline_access orders.read', + 'resource': RESOURCE, + 'state': transaction['state'], + 'nonce': transaction['nonce'], + 'code_challenge': challenge, + 'code_challenge_method': 'S256', + } + return f'{ISSUER}/oauth2/authorize?{urlencode(parameters)}', transaction +``` + +把 `transaction` 保存在与当前浏览器会话绑定的短期服务端存储中,并按 `state` 区分并发登录。回调处理时原子取出并删除,防止重复使用;不要把 verifier 或 Secret 放到授权 URL、前端日志或可被其他浏览器重放的共享状态中。公开 SPA 可由 OIDC 库在浏览器中维护对应短期状态。 + +仅登录时从 `scope` 移除 `orders.read`,并完全省略 `resource`;不需要续期时移除 `offline_access`。当前用户授权请求即使只关心业务 API,也必须包含 `openid` 和 `nonce`。一次请求不能携带多个业务资源。 + +可选参数:`prompt=login` 强制重新登录,`prompt=consent` 要求重新同意,`prompt=none` 尝试无交互授权,`max_age` 要求认证新鲜度。`none` 不能与其他 prompt 值一起使用;静默授权返回 `login_required`/`consent_required` 时应重新发起允许交互的流程。 + +授权端点也接受 `application/x-www-form-urlencoded` 的 POST,此时所有授权参数放在表单体,不能同时在查询串传参,不能重复传同名参数。 + +### 5.2 接收回调并兑换授权码 + +成功回调形如: + +```text +https://app.example.com/oidc/callback?code=...&state=...&iss=https%3A%2F%2Fauth.example.com +``` + +回调路由应拒绝重复的 `state`、`iss`、`code`、`error` 参数,校验与本次浏览器事务绑定的 `state`、过期时间和固定 Issuer,再处理成功或错误结果。错误回调不能建立登录会话。 + +以下函数接收框架已提取、无重复参数的字典,以及从当前浏览器会话中一次性取出的事务。`client_secret` 由服务端安全配置提供;公开客户端传 `None`。 + +```python +import httpx + + +async def exchange_code(query, transaction, client_secret=None): + if not transaction or time.time() >= transaction['expires_at']: + raise ValueError('登录事务不存在或已过期,请重新登录') + if not secrets.compare_digest(query.get('state', ''), transaction['state']): + raise ValueError('state 不匹配') + if query.get('iss') != ISSUER: + raise ValueError('授权响应 issuer 不匹配') + if query.get('error'): + raise ValueError(f'授权未完成:{query["error"]}') + if not query.get('code'): + raise ValueError('回调缺少授权码') + + form = { + 'grant_type': 'authorization_code', + 'code': query['code'], + 'redirect_uri': REDIRECT_URI, + 'code_verifier': transaction['verifier'], + } + auth = None + if client_secret is None: + form['client_id'] = CLIENT_ID + else: + auth = httpx.BasicAuth(CLIENT_ID, client_secret) + async with httpx.AsyncClient(timeout=10, follow_redirects=False) as client: + response = await client.post(f'{ISSUER}/oauth2/token', data=form, auth=auth) + response.raise_for_status() + return response.json() +``` + +这里的 Basic 示例适用于平台生成的 URL 安全字符客户端编号和 Secret。自行实现客户端认证时,按 OAuth Basic 规则对客户端编号和 Secret 分别做表单编码后再组装 Basic Header;不要把 `client_secret` 放到 URL 或请求体。客户端认证和令牌传输都依赖 HTTPS。 + +授权码默认 90 秒有效、只能使用一次。兑换时 `redirect_uri` 必须与授权时一致,verifier 必须对应本次 challenge。该次兑换不要附加 `scope` 或 `resource`,授权内容已绑定在授权码上。兑换失败或超时后不要反复重放同一个授权码,应重新开始登录。 + +协议返回的是原始 OAuth JSON,例如: + +```json +{ + "access_token": "", + "token_type": "Bearer", + "expires_in": 600, + "scope": "openid profile offline_access orders.read", + "id_token": "", + "refresh_token": "" +} +``` + +以实际返回的 `scope`、`expires_in` 和可选字段为准;用户可能只同意部分范围。收到这些字符串并不等于已经验证用户身份。 + +### 5.3 验证 ID Token,再创建应用登录会话 + +ID Token 用于登录身份,Access Token 用于资源访问,Refresh Token 用于续期。三者不要混用。JWKS 地址从固定可信 Issuer 的 Discovery 获取,核对 Discovery 的 `issuer`;不得按未验证 Token 的 `iss`、`jku` 等字段动态选择信任来源。 + +下面是适配当前 RS256 Profile 的最小校验示例。固定 JWKS 来源,校验类型、算法、签名、受众、时效、nonce 和 `at_hash`;如请求使用了 `max_age`,还需由 OIDC 库按 `auth_time` 校验相应要求。 + +```python +import jwt + +jwks_client = jwt.PyJWKClient(f'{ISSUER}/oauth2/jwks', timeout=5) + + +def verify_identity(tokens, expected_nonce): + id_token = tokens['id_token'] + header = jwt.get_unverified_header(id_token) + if header.get('alg') != 'RS256' or header.get('typ') != 'JWT': + raise ValueError('ID Token 类型或算法不正确') + if not isinstance(header.get('kid'), str) or not header['kid'].strip(): + raise ValueError('ID Token 缺少 kid') + if any(name in header for name in ('crit', 'jku', 'jwk', 'x5u', 'x5c')): + raise ValueError('不支持的 JWT Header') + key = jwks_client.get_signing_key_from_jwt(id_token).key + claims = jwt.decode( + id_token, + key, + algorithms=['RS256'], + issuer=ISSUER, + audience=CLIENT_ID, + leeway=60, + options={'require': ['iss', 'sub', 'aud', 'exp', 'iat', 'nonce', 'sid']}, + ) + if not isinstance(claims.get('sub'), str) or not claims['sub']: + raise ValueError('ID Token 缺少有效 sub') + nonce = claims.get('nonce') + if not isinstance(nonce, str) or not secrets.compare_digest(nonce, expected_nonce): + raise ValueError('nonce 不匹配') + audiences = claims['aud'] if isinstance(claims['aud'], list) else [claims['aud']] + if len(audiences) > 1 or 'azp' in claims: + if claims.get('azp') != CLIENT_ID: + raise ValueError('azp 不匹配') + if 'at_hash' in claims: + digest = hashlib.sha256(tokens['access_token'].encode('ascii')).digest() + expected_hash = base64.urlsafe_b64encode(digest[:16]).rstrip(b'=').decode('ascii') + if not isinstance(claims['at_hash'], str) or not secrets.compare_digest(claims['at_hash'], expected_hash): + raise ValueError('at_hash 不匹配') + return claims +``` + +`get_signing_key_from_jwt` 可能访问 JWKS 网络端点,在异步 Web 路由中将同步验证放在线程池,或使用支持异步获取与缓存 JWKS 的 OIDC 库。遇到未知 `kid` 时刷新可信 JWKS 并重新校验,失败则拒绝;不能临时关闭签名验证。 + +完成 `verify_identity(tokens, transaction["nonce"])` 后,以 `(iss, sub)` 作为外部用户稳定标识,按业务需要关联本地用户,并保存 `sid` 以处理退出通知。用户名、手机号、邮箱都不应当作不变的外部主键。 + +有后端的应用把令牌保存在服务端,只给浏览器自己的随机会话 Cookie,并设置 `HttpOnly`、`Secure` 和适合业务的 `SameSite`。应用自己的写操作仍需要相应 CSRF 防护。认证中心不会替应用创建本地登录态,也不会自动配置应用的数据权限。 + +### 5.4 获取用户资料 + +用用户 Access Token 请求 `GET /oauth2/userinfo`,或带相同 Bearer Header 的 POST: + +```http +GET /oauth2/userinfo HTTP/1.1 +Host: auth.example.com +Authorization: Bearer +``` + +返回原始用户 Claims JSON;核对返回的 `sub` 与已验证 ID Token 相同。字段取决于获准身份 Scope 及 Claim 策略。机器令牌不用于 UserInfo。 + +## 6. 应用调用业务 API + +将 Access Token 放在 Header,不放到 URL: + +```http +GET /orders HTTP/1.1 +Host: api.example.com +Authorization: Bearer +``` + +订单服务必须验证令牌针对 `https://api.example.com/orders` 签发,并检查包含 `orders.read`。拿到 `openid profile` 不代表具有订单读取权限;用户允许登录也不等于应用可以访问所有资源。 + +用户 Access Token 的受众包括认证中心 UserInfo,申请业务资源时另含相应业务受众;机器 Access Token 仅面向选定业务资源。不能因为 JWT 签名正确或包含某个 Scope 就省略 `aud` 校验。 + +## 7. 资源服务验证令牌 + +### 7.1 在线内省:用于及时响应撤销 + +新建一个专供订单 API 校验令牌的机密客户端,例如: + +```json +{ + "clientName": "订单服务令牌校验", + "clientType": "confidential", + "tokenEndpointAuthMethod": "client_secret_basic", + "grantTypes": ["client_credentials"], + "responseTypes": [], + "requirePkce": true, + "scopeCodes": [], + "resourceIds": [], + "redirectUris": [] +} +``` + +生成并保存该客户端的 Secret,然后编辑订单资源,将 `introspectionClientId` 设置为这个新客户端的 `clientId`。这里的校验权限来自资源绑定,不是来自 `grantTypes` 或业务客户端的申请范围。 + +订单 API 使用该校验客户端的 Basic 凭据调用: + +```http +POST /oauth2/introspect HTTP/1.1 +Host: auth.example.com +Authorization: Basic <校验客户端编号与密钥按规则编码后的值> +Content-Type: application/x-www-form-urlencoded + +token=&token_type_hint=access_token +``` + +校验顺序: + +1. 检查 HTTP 调用成功,且响应 `active` 严格为 `true`。 +2. 检查 `iss`、`aud` 是否属于预期认证中心和订单资源,必要时限制 `client_id`。 +3. 将空格分隔的 `scope` 拆分,检查本次操作要求的范围。 +4. 根据 `gty` 区分用户访问和机器访问,再执行资源自己的用户/租户/记录级授权。 + +失效响应为 `{"active": false}`。未绑定到该资源的其他客户端,即使知道某个有效 Token,也不能据此内省成功。只面向 UserInfo、没有业务资源受众的 Token 不用于这条资源内省路径;获取用户资料应调用 UserInfo。 + +内省超时或 5xx 应作为认证依赖不可用处理,不能放行。缓存内省结果会引入撤销延迟;需要即时反映授权撤销或会话下线的接口,应控制或避免正向缓存。 + +### 7.2 本地 JWT 校验:接受失效传播延迟时使用 + +固定可信 Issuer/JWKS 和 `RS256`,校验 `typ=at+jwt`、签名、`kid`、`iss`、预期业务 `aud`、`exp`/`nbf`/`iat`、Scope,以及当前接口允许的用户/机器令牌类型。不要接受 `typ=JWT` 的 ID Token 或 `logout+jwt` 的退出通知令牌作为 API 凭据。 + +本地验签不能自动得知 Grant 撤销、会话下线或客户端停用;已有 JWT 在本地到期前可能仍通过纯密码学检查。需要及时失效时使用在线内省或可靠的状态同步。当前实现的完整 Token Profile 可参考 [JWT 校验代码](../module_identity/security/jwt_profile.py)。 + +## 8. 刷新令牌与会话续期 + +取得 Refresh Token 需要客户端允许 `refresh_token`,且用户本次实际获准 `offline_access`。“记住授权选择”和浏览器“保持登录”均不等于获得离线访问资格。 + +机密客户端使用 Basic,公开客户端不传 Basic、改在表单中加入 `client_id`: + +```http +POST /oauth2/token HTTP/1.1 +Host: auth.example.com +Authorization: Basic <业务应用的客户端认证值> +Content-Type: application/x-www-form-urlencoded + +grant_type=refresh_token&refresh_token= +``` + +每次成功续期都会轮换 Refresh Token。将新的 Access Token、Refresh Token 和到期信息原子替换到应用会话中;当前刷新响应不重新签发 ID Token,不要覆盖原有已验证身份记录为“空”。 + +同一个应用会话的刷新操作必须串行化,多个并发 API 请求不能各自使用同一 Refresh Token。复用已使用的 Refresh Token 会触发令牌家族失效;请求超时且无法确认是否已经轮换时,不要盲目重试旧凭据,应重新授权恢复。 + +刷新时可选 `scope` 仅能缩小原授权;不能扩大权限,不能切换业务资源。若发送 `resource`,必须与原资源一致。新增 Scope 或切换资源应重新走授权码流程。 + +默认时限如下;客户端、资源配置和全局上限可能进一步限制实际结果,应以响应为准: + +| 配置 | 默认值 | +| ------------------------------- | ------------------------------ | +| 授权码 | 90 秒 | +| 登录/同意交互 | 300 秒 | +| ID Token | 300 秒 | +| Access Token | 600 秒,全局最大值默认 1800 秒 | +| Refresh Token 闲置期限 | 7 天 | +| Refresh Token 绝对期限 | 30 天,轮换不重置原始绝对期限 | +| SSO 闲置期限 | 30 分钟 | +| SSO 普通绝对期限 / 保持登录期限 | 8 小时 / 7 天 | +| 时钟容差 | 60 秒 | + +SSO 会话自然到期不自动取消已允许的离线授权;显式退出、管理员下线、授权撤销、用户被禁止访问或相关安全状态变化会阻断对应访问和续期。 + +## 9. 服务间调用 + +为后台任务单独登记机密客户端,只开启 `client_credentials`,绑定需要的资源和资源 Scope。例如: + +```json +{ + "clientName": "订单同步任务", + "clientType": "confidential", + "tokenEndpointAuthMethod": "client_secret_basic", + "grantTypes": ["client_credentials"], + "responseTypes": [], + "requirePkce": true, + "scopeCodes": ["orders.read"], + "resourceIds": ["orders-api"], + "redirectUris": [] +} +``` + +该流程没有授权回调,也没有 PKCE 参数交换;管理模型保留 `requirePkce=true` 不影响机器授权。保存后生成客户端 Secret,再发送: + +```http +POST /oauth2/token HTTP/1.1 +Host: auth.example.com +Authorization: Basic <机器客户端的认证值> +Content-Type: application/x-www-form-urlencoded + +grant_type=client_credentials&scope=orders.read&resource=https%3A%2F%2Fapi.example.com%2Forders +``` + +`scope` 和 `resource` 都要明确填写,Scope 只能包含该资源的业务范围。不要申请 `openid`、`profile`、`offline_access`。响应只有 Access Token 等访问字段,没有 ID Token 或 Refresh Token,到期后重新获取。 + +机器令牌的 `gty=client_credentials`,`sub` 为 `client:`。资源 API 不得将它当作某个用户身份,用户专属操作应拒绝机器令牌或使用单独的机器授权策略。 + +## 10. 退出、撤销与 Back-Channel + +### 10.1 浏览器退出 + +应用发起退出时,为当前会话保存随机退出 `state`,将浏览器导航到 Discovery 中的 `end_session_endpoint`,参数示例: + +```text +https://auth.example.com/oauth2/logout + ?id_token_hint= + &post_logout_redirect_uri=https%3A%2F%2Fapp.example.com%2Foidc%2Flogged-out + &state= +``` + +实际 URL 不换行,各参数均 URL 编码。`post_logout_redirect_uri` 必须已经在该客户端登记。保留本次需要的 `id_token_hint` 后,应用应按自己的退出语义清理本地会话;回到退出回调时验证预先保存的退出 state。 + +当前实现会展示退出确认页,用户确认后才清理认证中心会话及相关凭据。仅 GET 退出地址不立即执行全局退出。确认提交由认证中心页面完成,应用不要自行调用 `/oauth2/logout/confirm`。取消认证中心退出时,不能声称其他应用也已退出。 + +### 10.2 令牌撤销与管理员撤销的区别 + +客户端可以向 `/oauth2/revoke` 发送表单 `token` 和可选 `token_type_hint`,采用该客户端原有认证方式。未知或已撤销凭据按幂等语义成功处理。 + +| 操作 | 当前影响 | +| --------------------- | ---------------------------------------------------------------------- | +| 撤销 Refresh Token | 终止相应刷新能力;不能视为立即撤销所有已经签发的 Access Token | +| 撤销指定 Access Token | 阻断该 Token 的在线校验;纯本地验签仍需状态同步 | +| 后台“外部授权 → 撤销” | 撤销该用户对该应用的现有 Grant 和关联刷新凭据;旧 Token 的在线校验失败 | +| 后台“访问策略 → 禁止” | 阻止该用户再向该应用授权,并撤销已有授权;可在首次登录前配置 | +| 解除访问禁止 | 允许重新授权,不恢复旧 Grant 或旧 Token | +| 后台“外部会话 → 下线” | 撤销对应 `sid` 的访问资格,包括自然过期会话遗留的离线资格 | +| 客户端停用 | 阻止新登录/签发,并撤销该客户端相关授权与凭据 | + +应用本地会话与资源 Token 分别管理:本地 Cookie 没有过期不代表资源授权仍然有效;收到资源拒绝或退出通知后,应更新应用的会话状态。 + +### 10.3 接收后端退出通知 + +需要跨应用及时退出时,登记客户端的 `backchannelLogoutUris`,例如 `https://app.example.com/oidc/backchannel-logout`。默认只允许通过安全校验的公网 HTTPS 地址,不允许查询串和私网/回环地址。 + +接收端支持表单 POST 参数 `logout_token`,并完成以下处理: + +1. 固定可信 Issuer 和 JWKS,验证 RS256 签名、`typ=logout+jwt`、`kid`、`iss` 和 `aud=本客户端编号`。 +2. 校验 `iat`、`exp`、`jti`,以及 `events` 中的 `http://schemas.openid.net/event/backchannel-logout: {}`。 +3. 要求包含 `sid` 或 `sub`,且不能包含 `nonce`;按照被验证的 `sid`/`sub` 查找本应用会话。存在 `sid` 时优先精确匹配,避免误下线同一用户的其他设备。 +4. 按 `(iss, client_id, jti)` 实现幂等处理,清理对应会话后返回 200;重复通知仍返回成功。临时处理失败返回 5xx,由投递端重试。 + +通知目标来自该 SSO 会话实际参与的客户端,不是所有使用过同一用户的应用。接收端不依赖用户浏览器 Cookie,也不要信任未验签的退出 Claims。 + +## 11. 密钥维护 + +客户端 Secret 和认证中心签名密钥是两套不同的凭据: + +| 项目 | 客户端 Secret | OIDC 签名密钥 | +| ------------ | ----------------------------------------------------------- | ------------------------------------------------------------- | +| 用途 | 业务应用/资源服务向认证中心认证 | 认证中心签发 JWT,接入方通过 JWKS 验签 | +| 操作入口 | 客户端行操作“生成/轮换密钥” | 签名密钥管理,默认 `/oauth/key` | +| 应用保存什么 | 机密客户端保存自己的 Secret | 接入方只保存可信 Issuer 配置和缓存的公钥 | +| 轮换重点 | 使用生效时间和旧 Secret 退役窗口,完成部署后再撤销旧 Secret | 先发布新公钥,再激活签名;保留旧公钥覆盖已有 Token 的验证窗口 | + +普通客户端 Secret 轮换不会自动撤销正常用户授权。泄漏处置若需要阻断已有访问,应另外停用客户端或撤销相关授权,不能只依赖换 Secret。 + +签名密钥可通过 `POST /system/oauth/key/rotate` 创建 Pending 密钥,再通过 `PUT /system/oauth/key/{kid}/activate` 激活;安排发布、生效和旧密钥退役时间时考虑 Token 有效期、时钟容差及接入方 JWKS 缓存。不要删除仍用于验证未过期 Token 的旧公钥。 + +## 12. 接口速查 + +下表协议路径相对于固定 Issuer,不加平台 API 前缀: + +| 方法 | 路径 | 用途与身份 | +| ---------- | ----------------------------------------- | -------------------------------------------- | +| GET | `/.well-known/openid-configuration` | OIDC Discovery | +| GET | `/.well-known/oauth-authorization-server` | OAuth 服务端元数据 | +| GET | `/oauth2/jwks` | 公开验签密钥 | +| GET / POST | `/oauth2/authorize` | 浏览器用户授权;POST 使用表单 | +| POST | `/oauth2/token` | 授权码兑换、刷新、机器授权;表单和客户端认证 | +| GET / POST | `/oauth2/userinfo` | 用户 Access Token 的 Bearer 认证 | +| POST | `/oauth2/introspect` | 绑定资源的机密客户端 Basic 认证 | +| POST | `/oauth2/revoke` | 令牌撤销;对应客户端认证 | +| GET / POST | `/oauth2/logout` | 浏览器发起确认退出 | + +平台管理及内部页面接口: + +| 路径 | 用途 | +| ----------------------------------------------------- | ---------------------------------------------------------------------------- | +| `/auth/status` | 匿名读取启用开关;平台响应包装,前端通过普通 API 前缀访问 | +| `/auth/interaction/**` | 认证中心自己的登录、验证码、同意、修改密码和完成交互;不作为外部应用接入 API | +| `/system/oauth/client/**` | 客户端与 Secret 管理 | +| `/system/oauth/resource/**`、`/system/oauth/scope/**` | 资源和权限范围管理 | +| `/system/oauth/session/**`、`/system/oauth/grant/**` | 会话、授权和访问策略管理 | +| `/system/oauth/key/**` | 签名密钥管理 | +| `/monitor/oauth/audit/**` | OAuth 审计查询与导出 | + +管理端要求相应的 `system:oauthClient:*`、`system:oauthResource:*`、`system:oauthScope:*`、`system:oauthSession:*`、`system:oauthGrant:*`、`system:oauthKey:*` 等权限。菜单不可见时检查角色分配和对应菜单数据,不要改成匿名开放管理接口。 + +## 13. 常见问题 + +| 现象 | 优先检查与处理 | +| ----------------------------------------- | ------------------------------------------------------------------------------------------------------------- | +| 认证页显示“未启用” | 当前后端实例的 `OIDC_ENABLED`、配置环境及是否重启;`/auth/status` 返回的开关值 | +| 认证页显示“暂不可用” | 状态接口代理、后端是否同步升级、网络/超时;正常后台登录可用不代表 OIDC 已就绪 | +| 协议端点 404 | 开关是否关闭,代理是否把协议误加 `/prod-api` 或转发到 SPA | +| Discovery/JWKS 503 | Active 签名密钥、加密材料、数据库与 Redis 就绪状态;运行 `ruoyi app doctor` | +| `invalid_client` | 客户端类型/认证方式、编号、Secret 生效与到期、是否已停用;机密客户端使用 Basic | +| 回调提示 `redirect_uri is not registered` | 精确比对已登记的协议、域名、端口、路径和查询串;不要依赖前缀匹配 | +| `invalid_request` | 缺少 nonce/PKCE、重复字段、JSON 代替表单、授权 POST 混用查询参数、错误响应模式 | +| `invalid_scope` / `invalid_target` | 是否包含 `openid`,范围是否启用并授权给客户端,资源是否允许,`resource` 是否填了 audience,是否混入第二个资源 | +| `invalid_grant` | 授权码过期/复用、verifier 或回调不一致、刷新凭据已使用/撤销、会话或授权状态变化;重新授权 | +| 没收到 Refresh Token | 客户端是否启用刷新、请求是否含 `offline_access`、用户是否实际同意该范围 | +| 刷新后下次刷新失败 | 是否原子保存新 Refresh Token,是否多个请求并发使用同一个旧值 | +| 内省始终 `active=false` | RS 校验客户端是否绑定到对应资源、Token 是否含业务受众、授权/会话/客户端是否失效 | +| UserInfo 401 | 是否传用户 Access Token、是否过期/撤销;不要使用 ID Token、机器令牌或后台 Token | +| 角色/资料未返回 | 获准 Scope、客户端角色允许列表、Claim 策略和用户实际资料;不要假设请求即获准 | +| SPA 跨域失败 | 精确登记 Origin,检查预检代理;不要用 `*` 掩盖配置错误 | +| 已登录却反复显示登录页 | Issuer 与交互页是否同源,Cookie 是否因 HTTPS、主机、代理或浏览器策略未保存 | +| 登录/同意交互过期 | 默认 5 分钟,重新从业务应用开始授权,不复用旧 `interaction` 链接 | +| `login_required` / `consent_required` | 静默授权无法完成,重新发起允许页面交互的授权 | +| 429 / `temporarily_unavailable` | 遵守 `Retry-After`,检查依赖与限流;不要循环重放一次性授权码或刷新凭据 | +| 退出后另一应用仍显示登录 | 是否接收并验签 Back-Channel、是否映射并清理对应 sid、本地会话是否更新;仅清当前 Cookie 不会通知其他应用 | + +排查时结合 OAuth 审计记录和服务端日志中的错误代码、客户端编号、会话/授权标识。日志中不记录 Secret、密码、授权码、完整 Token 或 verifier。 + +## 14. 接入验收清单 + +- 配置启用,Discovery/JWKS 可访问,关闭开关时三个认证交互页面被阻止,错误页与普通后台登录可用。 +- 首次登录完成授权码 + PKCE;错误 state、nonce、Issuer、签名、受众或过期 Token 会被应用拒绝。 +- 同一浏览器接入第二个应用可以复用 SSO,各应用仍分别遵守自己的授权策略。 +- 仅身份登录、允许订单访问、拒绝订单范围三种场景的结果符合预期,API 不越权。 +- Refresh Token 正确轮换并串行使用;用户不同意离线访问时不假定存在刷新凭据。 +- 机器令牌可访问允许的资源,不能冒充用户或访问未授权资源。 +- 撤销 Grant、禁止访问、会话下线和客户端停用能使在线内省/续期失败;解除禁止后需要重新授权。 +- 全局退出经过用户确认;退出回调 state 正确;Back-Channel 验签、会话隔离、重复通知和失败重试有效。 +- 客户端 Secret、签名密钥轮换后新凭据可用,过渡期旧凭据按预期验证,日志无敏感凭据。 + +在目标环境完成真实 HTTPS、反向代理、数据库并发、Cookie 与公网 Back-Channel 联调;单元测试不能替代这些部署验证。 + +## 15. 代码参考 + +需要了解具体实现时,可查阅以下配置、协议模型和令牌校验源码。 + +| 资料 | 用途 | +| ---------------------------------------------------------------------- | --------------------------------------- | +| [OIDC 配置](../config/env.py) | 配置名、默认值和启动校验 | +| [协议数据模型](../module_identity/entity/vo/protocol_vo.py) | 授权、Token 等入参约束 | +| [客户端管理模型](../module_identity/entity/vo/oauth_client_vo.py) | 客户端类型、Grant、URI 与 Secret 配置 | +| [资源与 Scope 模型](../module_identity/entity/vo/oauth_resource_vo.py) | 资源受众、内省绑定与 Scope 配置 | +| [JWT Profile](../module_identity/security/jwt_profile.py) | Access/ID/Logout Token 的用途与完整校验 | diff --git a/ruoyi-fastapi-backend/exceptions/exception.py b/ruoyi-fastapi-backend/exceptions/exception.py index 5fe5617b7..a32ceeb0e 100644 --- a/ruoyi-fastapi-backend/exceptions/exception.py +++ b/ruoyi-fastapi-backend/exceptions/exception.py @@ -1,3 +1,6 @@ +from utils.oidc_util import OidcUtil + + class LoginException(Exception): """ 自定义登录异常LoginException @@ -18,6 +21,127 @@ def __init__(self, data: str | None = None, message: str | None = None) -> None: self.message = message +class OAuthProtocolException(Exception): + """ + OAuth/OIDC 协议异常。 + + 异常 message 和字符串表示使用中文,error_description 保持 OAuth 要求的 ASCII 格式。 + 统一异常处理器仅返回标准协议字段和经过验证的重定向状态,不回传内部诊断详情。 + """ + + _DEFAULT_BAD_REQUEST_STATUS = 400 + + def __init__( + self, + error: str, + error_description: str | None = None, + status_code: int = 400, + *, + redirect_uri: str | None = None, + state: str | None = None, + redirect_uri_verified: bool = False, + issuer: str | None = None, + headers: dict[str, str] | None = None, + ) -> None: + self.error = error + self.error_description = OidcUtil.protocol_error_description(error, error_description) + self.message = OidcUtil.localized_oauth_message(error, error_description) + self.status_code = ( + 401 if error == 'invalid_client' and status_code == self._DEFAULT_BAD_REQUEST_STATUS else status_code + ) + self.redirect_uri = redirect_uri + self.state = state + self.redirect_uri_verified = redirect_uri_verified + self.issuer = issuer + self.headers = dict(headers or {}) + super().__init__(self.message) + + @property + def can_redirect(self) -> bool: + """ + 判断是否允许协议错误重定向。 + + :return: Redirect URI 存在且已完成服务端精确校验时为 True + """ + return bool(self.redirect_uri and self.redirect_uri_verified) + + @property + def redirect_safe(self) -> bool: + """ + 返回安全重定向状态别名。 + + :return: 与 :attr:`can_redirect` 相同的安全状态 + """ + return self.can_redirect + + def as_dict(self, *, include_state: bool = False) -> dict[str, str]: + """ + 转换为 OAuth 标准错误 JSON 字段。 + + :param include_state: 是否在错误 JSON 中包含 state + :return: 标准 OAuth 错误字段 + """ + result: dict[str, str] = {'error': self.error} + if self.error_description: + result['error_description'] = self.error_description + if include_state and self.state: + result['state'] = self.state + return result + + +class OidcInteractionException(Exception): + """ + 认证交互状态异常。 + + 交互异常不默认跳转到外部地址;只有 Interaction 已绑定并验证了 Client + Redirect URI 时,调用方才可将其转换为 ``OAuthProtocolException``。 + """ + + def __init__( + self, + interaction_id: str | None = None, + message: str | None = None, + *, + error: str = 'interaction_required', + status_code: int = 400, + redirect_uri: str | None = None, + state: str | None = None, + redirect_uri_verified: bool = False, + ) -> None: + self.interaction_id = interaction_id + self.message = message or '认证交互无效或已过期' + self.error = error + self.status_code = status_code + self.redirect_uri = redirect_uri + self.state = state + self.redirect_uri_verified = redirect_uri_verified + super().__init__(self.message) + + @property + def can_redirect(self) -> bool: + """ + 判断交互是否已经具备安全重定向条件。 + + :return: 交互已绑定并验证 Client Redirect URI 时为 True + """ + return bool(self.redirect_uri and self.redirect_uri_verified) + + def as_protocol_exception(self) -> OAuthProtocolException: + """ + 显式转换为标准 OAuth 协议异常。 + + :return: 可由协议控制器处理的 OAuth 异常 + """ + return OAuthProtocolException( + self.error, + self.message, + self.status_code, + redirect_uri=self.redirect_uri, + state=self.state, + redirect_uri_verified=self.redirect_uri_verified, + ) + + class PermissionException(Exception): """ 自定义权限异常PermissionException diff --git a/ruoyi-fastapi-backend/exceptions/handle.py b/ruoyi-fastapi-backend/exceptions/handle.py index 6c463ef31..996fb0d36 100644 --- a/ruoyi-fastapi-backend/exceptions/handle.py +++ b/ruoyi-fastapi-backend/exceptions/handle.py @@ -1,5 +1,9 @@ -from fastapi import FastAPI, Request, Response +from urllib.parse import urlsplit + +from fastapi import FastAPI, Request, Response, status from fastapi.exceptions import HTTPException +from fastapi.responses import JSONResponse as FastAPIJSONResponse +from fastapi.responses import RedirectResponse from pydantic_validation_decorator import FieldValidationError from exceptions.exception import ( @@ -7,13 +11,18 @@ FileRangeNotSatisfiableException, LoginException, ModelValidatorException, + OAuthProtocolException, + OidcInteractionException, PermissionException, ServiceException, ServiceWarning, ) from utils.log_util import logger +from utils.oidc_util import OidcUtil from utils.response_util import JSONResponse, ResponseUtil, jsonable_encoder +_OAUTH_RESPONSE_PARAMETER_NAMES = frozenset({'code', 'error', 'error_description', 'error_uri', 'iss', 'state'}) + def handle_exception(app: FastAPI) -> None: """ @@ -25,6 +34,37 @@ def handle_exception(app: FastAPI) -> None: async def auth_exception_handler(request: Request, exc: AuthException) -> Response: return ResponseUtil.unauthorized(data=exc.data, msg=exc.message) + # 自定义OAuth协议异常 + @app.exception_handler(OAuthProtocolException) + async def oauth_protocol_exception_handler(request: Request, exc: OAuthProtocolException) -> Response: + if exc.can_redirect: + return _build_oauth_redirect(exc) + headers = {**exc.headers, 'Cache-Control': 'no-store', 'Pragma': 'no-cache'} + if exc.error == 'invalid_client' and exc.status_code == status.HTTP_401_UNAUTHORIZED: + headers.setdefault('WWW-Authenticate', 'Basic realm="oauth2/token"') + return FastAPIJSONResponse( + content=exc.as_dict(), + status_code=exc.status_code, + headers=headers, + ) + + # 自定义OIDC认证交互异常 + @app.exception_handler(OidcInteractionException) + async def oidc_interaction_exception_handler(request: Request, exc: OidcInteractionException) -> Response: + safe_message = { + 'invalid_request': '认证交互请求无效', + 'interaction_required': '认证交互已过期或不可用', + 'login_required': '需要登录', + 'consent_required': '需要授权确认', + 'invalid_scope': '请求权限无效', + 'server_error': '认证服务暂不可用', + }.get(exc.error, '认证交互无效或已过期') + return ResponseUtil.failure( + data=exc.interaction_id, + msg=safe_message, + headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'}, + ) + # 自定义登录检验异常 @app.exception_handler(LoginException) async def login_exception_handler(request: Request, exc: LoginException) -> Response: @@ -86,3 +126,34 @@ async def http_exception_handler(request: Request, exc: HTTPException) -> Respon async def exception_handler(request: Request, exc: Exception) -> Response: logger.exception(exc) return ResponseUtil.error(msg=str(exc)) + + +def _build_oauth_redirect(exc: OAuthProtocolException) -> Response: + """ + 构造安全的 OAuth 授权响应重定向。 + + :param exc: 已完成 Redirect URI 精确注册校验的协议异常 + :return: 303 重定向或本地 400 标准错误响应 + """ + parsed = urlsplit(exc.redirect_uri or '') + if parsed.fragment or not parsed.scheme or not parsed.netloc: + return FastAPIJSONResponse( + content={'error': 'server_error', 'error_description': 'Invalid validated redirect URI'}, + status_code=400, + headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'}, + ) + params = [('error', exc.error)] + if exc.error_description: + params.append(('error_description', exc.error_description)) + if exc.state: + params.append(('state', exc.state)) + if exc.issuer: + params.append(('iss', exc.issuer)) + location = OidcUtil.replace_query_parameters( + exc.redirect_uri or '', params, _OAUTH_RESPONSE_PARAMETER_NAMES, fragment='', doseq=True + ) + return RedirectResponse( + url=location, + status_code=303, + headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'}, + ) diff --git a/ruoyi-fastapi-backend/middlewares/cors_middleware.py b/ruoyi-fastapi-backend/middlewares/cors_middleware.py index 42206a6fc..435fdd8fe 100644 --- a/ruoyi-fastapi-backend/middlewares/cors_middleware.py +++ b/ruoyi-fastapi-backend/middlewares/cors_middleware.py @@ -7,10 +7,7 @@ def add_cors_middleware(app: FastAPI) -> None: 添加跨域中间件 :param app: FastAPI对象 - :return: """ - # 前端页面url - origins = ['*'] expose_headers = [ 'x-body-encrypted', 'x-key-id', @@ -21,11 +18,9 @@ def add_cors_middleware(app: FastAPI) -> None: 'content-range', 'content-length', ] - - # 后台api允许跨域 app.add_middleware( CORSMiddleware, - allow_origins=origins, + allow_origins=['*'], allow_credentials=True, allow_methods=['*'], allow_headers=['*'], diff --git a/ruoyi-fastapi-backend/middlewares/handle.py b/ruoyi-fastapi-backend/middlewares/handle.py index b56c0c0e7..890f37b8a 100644 --- a/ruoyi-fastapi-backend/middlewares/handle.py +++ b/ruoyi-fastapi-backend/middlewares/handle.py @@ -1,11 +1,12 @@ from fastapi import FastAPI -from config.env import AppConfig +from config.env import AppConfig, OidcConfig from middlewares.api_response_header_middleware import add_api_response_header_middleware from middlewares.context_middleware import add_context_cleanup_middleware from middlewares.cors_middleware import add_cors_middleware from middlewares.demo_mode_middleware import add_demo_mode_middleware from middlewares.gzip_middleware import add_gzip_middleware +from middlewares.oidc_cors_middleware import add_oidc_cors_middleware from middlewares.trace_middleware import add_trace_middleware from middlewares.transport_crypto_middleware import add_transport_crypto_middleware @@ -18,6 +19,9 @@ def handle_middleware(app: FastAPI) -> None: add_context_cleanup_middleware(app) # 加载跨域中间件 add_cors_middleware(app) + if OidcConfig.oidc_enabled: + # 加载 OIDC 专用 CORS 中间件 + add_oidc_cors_middleware(app) # 加载gzip压缩中间件 add_gzip_middleware(app) # 加载接口响应头追加中间件 diff --git a/ruoyi-fastapi-backend/middlewares/oidc_cors_middleware.py b/ruoyi-fastapi-backend/middlewares/oidc_cors_middleware.py new file mode 100644 index 000000000..0c0824c94 --- /dev/null +++ b/ruoyi-fastapi-backend/middlewares/oidc_cors_middleware.py @@ -0,0 +1,249 @@ +from urllib.parse import urlsplit + +from fastapi import FastAPI +from starlette.datastructures import Headers +from starlette.responses import PlainTextResponse +from starlette.types import ASGIApp, Message, Receive, Scope, Send + +from config.env import OidcConfig +from module_identity.service.runtime_service import OidcRuntimeService + + +class OidcCorsMiddleware: + """ + 为 OIDC 协议和认证交互接口提供独立的 CORS 边界。 + """ + + _OIDC_PREFIXES = ('/oauth2/', '/.well-known/') + _INTERACTION_PREFIX = '/auth/interaction/' + _NO_CORS_PATHS = ('/oauth2/authorize',) + _PUBLIC_METADATA_PATHS = ( + '/.well-known/openid-configuration', + '/.well-known/oauth-authorization-server', + '/oauth2/jwks', + ) + + def __init__(self, app: ASGIApp) -> None: + self.app = app + + @classmethod + def _is_oidc_path(cls, path: str) -> bool: + """ + 判断请求路径是否属于 OIDC 协议端点。 + + :param path: 请求路径 + :return: 是否为 OIDC 协议路径 + """ + normalized = path.rstrip('/') or '/' + return normalized.startswith(cls._OIDC_PREFIXES) + + @classmethod + def _is_interaction_path(cls, path: str) -> bool: + """ + 判断请求路径是否属于认证交互接口。 + + :param path: 请求路径 + :return: 是否为认证交互路径 + """ + normalized = path.rstrip('/') or '/' + return normalized.startswith(cls._INTERACTION_PREFIX) or normalized == '/oauth2/logout/confirm' + + @staticmethod + def _issuer_origin() -> str | None: + """ + 从 OIDC issuer 配置中提取无路径的同源 Origin。 + + :return: 合法的 HTTP(S) Origin,配置无效时返回 ``None`` + """ + try: + parsed = urlsplit(str(OidcConfig.oidc_issuer)) + if parsed.scheme not in {'http', 'https'} or not parsed.netloc or parsed.username or parsed.password: + return None + return f'{parsed.scheme}://{parsed.netloc}' + except (AttributeError, TypeError, ValueError): + return None + + @classmethod + def _is_authorize_path(cls, path: str) -> bool: + """ + 判断请求路径是否为禁止跨域的授权端点。 + + :param path: 请求路径 + :return: 是否为授权端点路径 + """ + normalized = path.rstrip('/') or '/' + return normalized in cls._NO_CORS_PATHS + + @classmethod + def _is_public_metadata_path(cls, path: str) -> bool: + """ + 判断请求路径是否为公开的 OIDC 元数据端点。 + + :param path: 请求路径 + :return: 是否为公开元数据路径 + """ + normalized = path.rstrip('/') or '/' + return normalized in cls._PUBLIC_METADATA_PATHS + + @classmethod + def _allowed_origin(cls, origin: str | None, scope: Scope | None = None) -> bool: + """ + 校验 Origin 是否存在于静态或运行时注册的跨域白名单中。 + + :param origin: 请求携带的 Origin + :param scope: 当前 ASGI 请求作用域,用于读取运行时注册白名单 + :return: Origin 是否被允许 + """ + if not origin: + return False + configured = set(OidcConfig.cors_origin_list) + current_app = scope.get('app') if scope else None + state = getattr(current_app, 'state', None) + registered = getattr(state, 'oidc_registered_cors_origins', ()) + if isinstance(registered, (str, bytes)): + return False + try: + configured.update(item for item in registered if isinstance(item, str)) + except TypeError: + return False + return origin in configured + + @staticmethod + def _without_cors(headers: list[tuple[bytes, bytes]]) -> list[tuple[bytes, bytes]]: + """ + 移除响应中的 CORS 头并清理 ``Vary: Origin`` 标记。 + + :param headers: ASGI 响应头列表 + :return: 移除 CORS 相关头后的响应头列表 + """ + cleaned: list[tuple[bytes, bytes]] = [] + for name, value in headers: + lower_name = name.lower() + if lower_name.startswith(b'access-control-'): + continue + if lower_name != b'vary': + cleaned.append((name, value)) + continue + vary_values = [item.strip() for item in value.decode('latin-1').split(',')] + vary_values = [item for item in vary_values if item and item.lower() != 'origin'] + if vary_values: + cleaned.append((name, ', '.join(vary_values).encode('latin-1'))) + return cleaned + + @classmethod + def _with_allowed_cors(cls, headers: list[tuple[bytes, bytes]], origin: str) -> list[tuple[bytes, bytes]]: + """ + 为响应追加指定 Origin 的 CORS 头。 + + :param headers: 原始 ASGI 响应头列表 + :param origin: 已校验通过的请求 Origin + :return: 清理旧 CORS 头并追加新 CORS 头后的列表 + """ + clean = cls._without_cors(headers) + clean.extend( + [ + (b'access-control-allow-origin', origin.encode('utf-8')), + (b'access-control-allow-methods', b'GET, POST, OPTIONS'), + (b'access-control-allow-headers', b'Authorization, Content-Type, X-Requested-With'), + (b'access-control-expose-headers', b'WWW-Authenticate'), + (b'vary', b'Origin'), + ] + ) + return clean + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + """ + 按 OIDC 端点类型校验 Origin 并处理 CORS 响应。 + + :param scope: 当前 ASGI 请求作用域 + :param receive: 接收 ASGI 消息的可调用对象 + :param send: 发送 ASGI 消息的可调用对象 + :return: None + """ + # 注册层已经按 OIDC_ENABLED 控制是否挂载;这里保留运行时保护, + # 避免测试或动态配置场景下开关关闭后仍处理协议请求。 + if not OidcConfig.oidc_enabled: + await self.app(scope, receive, send) + return + if scope.get('type') != 'http': + await self.app(scope, receive, send) + return + + path = str(scope.get('path', '')) + oidc_path = self._is_oidc_path(path) + interaction_path = self._is_interaction_path(path) + if not oidc_path and not interaction_path: + await self.app(scope, receive, send) + return + + origin = Headers(scope=scope).get('origin') + authorize = oidc_path and self._is_authorize_path(path) + public_metadata = oidc_path and self._is_public_metadata_path(path) + logout_navigation = path.rstrip('/') == '/oauth2/logout' and scope.get('method') in {'GET', 'POST'} + if ( + origin + and oidc_path + and not interaction_path + and not authorize + and not public_metadata + and not logout_navigation + ): + await OidcRuntimeService.ensure_cors_snapshot(scope['app']) + if interaction_path: + allowed = not origin or origin == self._issuer_origin() + else: + allowed = ( + not authorize and not logout_navigation and (public_metadata or self._allowed_origin(origin, scope)) + ) + + if authorize and origin: + response = PlainTextResponse('Authorization Endpoint 不支持 CORS', status_code=403) + await response(scope, receive, send) + return + if interaction_path and origin and not allowed: + response = PlainTextResponse('认证交互接口仅允许认证中心同源 Origin', status_code=403) + await response(scope, receive, send) + return + if oidc_path and origin and not public_metadata and not logout_navigation and not allowed: + response = PlainTextResponse('Origin 不被认证中心允许', status_code=403) + await response(scope, receive, send) + return + + if str(scope.get('method', '')).upper() == 'OPTIONS' and (oidc_path or interaction_path): + if not allowed or not origin: + response = PlainTextResponse('Origin 不被认证中心允许', status_code=403) + await response(scope, receive, send) + return + response = PlainTextResponse('', status_code=200) + for name, value in self._with_allowed_cors([], origin): + response.headers[name.decode('latin-1')] = value.decode('latin-1') + await response(scope, receive, send) + return + + async def send_wrapper(message: Message) -> None: + """ + 在下游响应开始消息中应用 OIDC CORS 策略。 + + :param message: 下游 ASGI 消息 + :return: None + """ + if message.get('type') != 'http.response.start': + await send(message) + return + headers = list(message.get('headers', [])) + if authorize or not allowed: + headers = self._without_cors(headers) + elif origin: + headers = self._with_allowed_cors(headers, origin) + await send({**message, 'headers': headers}) + + await self.app(scope, receive, send_wrapper) + + +def add_oidc_cors_middleware(app: FastAPI) -> None: + """ + 添加 OIDC 专用 CORS 中间件。 + + :param app: FastAPI 对象 + """ + app.add_middleware(OidcCorsMiddleware) diff --git a/ruoyi-fastapi-backend/middlewares/transport_crypto_middleware.py b/ruoyi-fastapi-backend/middlewares/transport_crypto_middleware.py index 623ea1cc4..6289ce2be 100644 --- a/ruoyi-fastapi-backend/middlewares/transport_crypto_middleware.py +++ b/ruoyi-fastapi-backend/middlewares/transport_crypto_middleware.py @@ -22,6 +22,20 @@ class TransportCryptoMiddleware: 传输层请求解密与响应加密中间件 """ + # 标准 OAuth/OIDC 客户端无法理解项目自定义传输信封,协议路径始终旁路, + # 即使运行时通过环境变量覆盖普通排除列表也不得重新启用加密。 + _STANDARD_OIDC_PATHS = ( + '/.well-known/openid-configuration', + '/.well-known/oauth-authorization-server', + '/oauth2/authorize', + '/oauth2/token', + '/oauth2/userinfo', + '/oauth2/jwks', + '/oauth2/revoke', + '/oauth2/introspect', + '/oauth2/logout', + ) + _ENCRYPT_REQUEST_HEADER = 'x-transport-encrypt' _ENCRYPT_RESPONSE_HEADER = 'x-body-encrypted' _ENCRYPT_ALG_HEADER = 'x-encrypt-alg' @@ -710,6 +724,7 @@ def _is_excluded_path(cls, path: str) -> bool: for excluded_path in TransportCryptoConfig.transport_crypto_exclude_paths.split(',') if excluded_path.strip() ] + excluded_paths.extend(cls._STANDARD_OIDC_PATHS) return any(path == excluded_path or path.startswith(f'{excluded_path}/') for excluded_path in excluded_paths) @classmethod diff --git a/ruoyi-fastapi-backend/module_admin/controller/role_controller.py b/ruoyi-fastapi-backend/module_admin/controller/role_controller.py index c902e4aad..6f1b2509c 100644 --- a/ruoyi-fastapi-backend/module_admin/controller/role_controller.py +++ b/ruoyi-fastapi-backend/module_admin/controller/role_controller.py @@ -346,7 +346,9 @@ async def add_system_role_user( ) -> Response: if not current_user.user.admin: await RoleService.check_role_data_scope_services(query_db, str(add_role_user.role_id), data_scope_sql) - add_role_user_result = await UserService.add_user_role_services(query_db, add_role_user) + add_role_user_result = await UserService.add_user_role_services( + query_db, add_role_user, current_user.user.user_name + ) logger.info(add_role_user_result.message) return ResponseUtil.success(msg=add_role_user_result.message) @@ -366,8 +368,11 @@ async def cancel_system_role_user( request: Request, cancel_user_role: CrudUserRoleModel, query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], ) -> Response: - cancel_user_role_result = await UserService.delete_user_role_services(query_db, cancel_user_role) + cancel_user_role_result = await UserService.delete_user_role_services( + query_db, cancel_user_role, current_user.user.user_name + ) logger.info(cancel_user_role_result.message) return ResponseUtil.success(msg=cancel_user_role_result.message) @@ -387,8 +392,11 @@ async def batch_cancel_system_role_user( request: Request, batch_cancel_user_role: Annotated[CrudUserRoleModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], ) -> Response: - batch_cancel_user_role_result = await UserService.delete_user_role_services(query_db, batch_cancel_user_role) + batch_cancel_user_role_result = await UserService.delete_user_role_services( + query_db, batch_cancel_user_role, current_user.user.user_name + ) logger.info(batch_cancel_user_role_result.message) return ResponseUtil.success(msg=batch_cancel_user_role_result.message) diff --git a/ruoyi-fastapi-backend/module_admin/controller/user_controller.py b/ruoyi-fastapi-backend/module_admin/controller/user_controller.py index 5c50a66e3..a720a5aba 100644 --- a/ruoyi-fastapi-backend/module_admin/controller/user_controller.py +++ b/ruoyi-fastapi-backend/module_admin/controller/user_controller.py @@ -584,7 +584,9 @@ async def update_system_role_user( await UserService.check_user_data_scope_services(query_db, user_id, user_data_scope_sql) await RoleService.check_role_data_scope_services(query_db, role_ids, role_data_scope_sql) add_user_role_result = await UserService.add_user_role_services( - query_db, CrudUserRoleModel(userId=user_id, roleIds=role_ids) + query_db, + CrudUserRoleModel(userId=user_id, roleIds=role_ids), + current_user.user.user_name, ) logger.info(add_user_role_result.message) diff --git a/ruoyi-fastapi-backend/module_admin/dao/role_dao.py b/ruoyi-fastapi-backend/module_admin/dao/role_dao.py index 3acc6ff4b..b71256f28 100644 --- a/ruoyi-fastapi-backend/module_admin/dao/role_dao.py +++ b/ruoyi-fastapi-backend/module_admin/dao/role_dao.py @@ -247,6 +247,18 @@ async def get_role_menu_dao(cls, db: AsyncSession, role: RoleModel) -> Sequence[ return role_menu_query_all + @classmethod + async def list_role_menu_ids(cls, db: AsyncSession, role_id: int) -> Sequence[int]: + """ + 查询角色完整菜单关联编号 + + :param db: 异步数据库会话 + :param role_id: 角色编号 + :return: 角色已关联的菜单编号序列 + """ + result = await db.execute(select(SysRoleMenu.menu_id).where(SysRoleMenu.role_id == role_id)) + return result.scalars().all() + @classmethod async def add_role_menu_dao(cls, db: AsyncSession, role_menu: RoleMenuModel) -> None: """ diff --git a/ruoyi-fastapi-backend/module_admin/dao/user_dao.py b/ruoyi-fastapi-backend/module_admin/dao/user_dao.py index 6f65aaaef..ee15fdb82 100644 --- a/ruoyi-fastapi-backend/module_admin/dao/user_dao.py +++ b/ruoyi-fastapi-backend/module_admin/dao/user_dao.py @@ -51,6 +51,18 @@ async def get_user_by_name(cls, db: AsyncSession, user_name: str) -> SysUser | N return query_user_info + @classmethod + async def get_role_ids(cls, db: AsyncSession, user_id: int) -> set[int]: + """ + 查询用户当前关联的角色ID + + :param db: orm对象 + :param user_id: 用户id + :return: 角色ID集合 + """ + result = await db.execute(select(SysUserRole.role_id).where(SysUserRole.user_id == user_id)) + return set(result.scalars().all()) + @classmethod async def get_user_by_info(cls, db: AsyncSession, user: UserModel) -> SysUser | None: """ diff --git a/ruoyi-fastapi-backend/module_admin/service/login_service.py b/ruoyi-fastapi-backend/module_admin/service/login_service.py index a13212769..5368d15d7 100644 --- a/ruoyi-fastapi-backend/module_admin/service/login_service.py +++ b/ruoyi-fastapi-backend/module_admin/service/login_service.py @@ -15,7 +15,6 @@ from common.vo import CrudResponseModel from config.env import AppConfig, JwtConfig from exceptions.exception import AuthException, LoginException, ServiceException -from module_admin.dao.login_dao import login_by_account from module_admin.dao.user_dao import UserDao from module_admin.entity.do.dept_do import SysDept from module_admin.entity.do.menu_do import SysMenu @@ -23,6 +22,10 @@ from module_admin.entity.vo.login_vo import MenuTreeModel, MetaModel, RouterModel, SmsCode, UserLogin, UserRegister from module_admin.entity.vo.user_vo import AddUserModel, CurrentUserModel, ResetUserModel, TokenData, UserInfoModel from module_admin.service.user_service import UserService +from module_identity.service.identity_service import ( + CredentialAuthenticationError, + CredentialAuthenticationService, +) from utils.client_ip_util import ClientIPUtil from utils.common_util import CamelCaseUtil from utils.jwt_util import JwtUtil @@ -81,62 +84,19 @@ async def authenticate_user( :param login_user: 登录用户对象 :return: 校验结果 """ - await cls.__check_login_ip(request) - account_lock = await request.app.state.redis.get( - f'{RedisInitKeyConfig.ACCOUNT_LOCK.key}:{login_user.user_name}' - ) - if login_user.user_name == account_lock: - logger.warning('账号已锁定,请稍后再试') - raise LoginException(data='', message='账号已锁定,请稍后再试') - # 判断请求是否来自于api文档,如果是返回指定格式的结果,用于修复api文档认证成功后token显示undefined的bug - request_from_swagger = ( - request.headers.get('referer').endswith('docs') if request.headers.get('referer') else False - ) - request_from_redoc = ( - request.headers.get('referer').endswith('redoc') if request.headers.get('referer') else False - ) - # 判断是否开启验证码,开启则验证,否则不验证(dev模式下来自API文档的登录请求不检验) - if not login_user.captcha_enabled or ( - (request_from_swagger or request_from_redoc) and AppConfig.app_env == 'dev' - ): - pass - else: - await cls.__check_login_captcha(request, login_user) - user = await login_by_account(query_db, login_user.user_name) - if not user: - logger.warning('用户不存在') - raise LoginException(data='', message='用户不存在') - if not PwdUtil.verify_password(login_user.password, user[0].password): - cache_password_error_count = await request.app.state.redis.get( - f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:{login_user.user_name}' - ) - password_error_counted = 0 - if cache_password_error_count: - password_error_counted = cache_password_error_count - password_error_count = int(password_error_counted) + 1 - await request.app.state.redis.set( - f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:{login_user.user_name}', - password_error_count, - ex=timedelta(minutes=10), + referer = request.headers.get('referer') + request_from_docs = bool(referer and referer.endswith(('docs', 'redoc'))) + try: + return await CredentialAuthenticationService.authenticate_legacy( + request.app.state.redis, + query_db, + login_user, + client_ip=ClientIPUtil.get_client_ip(request), + skip_captcha=request_from_docs and AppConfig.app_env == 'dev', ) - if password_error_count > CommonConstant.PASSWORD_ERROR_COUNT: - await request.app.state.redis.delete( - f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:{login_user.user_name}' - ) - await request.app.state.redis.set( - f'{RedisInitKeyConfig.ACCOUNT_LOCK.key}:{login_user.user_name}', - login_user.user_name, - ex=timedelta(minutes=10), - ) - logger.warning('10分钟内密码已输错超过5次,账号已锁定,请10分钟后再试') - raise LoginException(data='', message='10分钟内密码已输错超过5次,账号已锁定,请10分钟后再试') - logger.warning('密码错误') - raise LoginException(data='', message='密码错误') - if user[0].status == '1': - logger.warning('用户已停用') - raise LoginException(data='', message='用户已停用') - await request.app.state.redis.delete(f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:{login_user.user_name}') - return user + except CredentialAuthenticationError as exc: + logger.warning(exc.legacy_message) + raise LoginException(data='', message=exc.legacy_message) from exc @classmethod async def unlock_screen_services( @@ -160,39 +120,6 @@ async def unlock_screen_services( return True - @classmethod - async def __check_login_ip(cls, request: Request) -> bool: - """ - 校验用户登录ip是否在黑名单内 - - :param request: Request对象 - :return: 校验结果 - """ - black_ip_value = await request.app.state.redis.get(f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.login.blackIPList') - black_ip_list = black_ip_value.split(',') if black_ip_value else [] - if ClientIPUtil.get_client_ip(request) in black_ip_list: - logger.warning('当前IP禁止登录') - raise LoginException(data='', message='当前IP禁止登录') - return True - - @classmethod - async def __check_login_captcha(cls, request: Request, login_user: UserLogin) -> bool: - """ - 校验用户登录验证码 - - :param request: Request对象 - :param login_user: 登录用户对象 - :return: 校验结果 - """ - captcha_value = await request.app.state.redis.get(f'{RedisInitKeyConfig.CAPTCHA_CODES.key}:{login_user.uuid}') - if not captcha_value: - logger.warning('验证码已失效') - raise LoginException(data='', message='验证码已失效') - if login_user.code != str(captcha_value): - logger.warning('验证码错误') - raise LoginException(data='', message='验证码错误') - return True - @classmethod async def create_access_token(cls, data: dict, expires_delta: timedelta | None = None) -> str: """ diff --git a/ruoyi-fastapi-backend/module_admin/service/role_service.py b/ruoyi-fastapi-backend/module_admin/service/role_service.py index 4db12a93a..8cfed38e1 100644 --- a/ruoyi-fastapi-backend/module_admin/service/role_service.py +++ b/ruoyi-fastapi-backend/module_admin/service/role_service.py @@ -18,6 +18,7 @@ RolePageQueryModel, ) from module_admin.entity.vo.user_vo import UserInfoModel, UserRolePageQueryModel +from module_identity.service.identity_service import IdentitySecurityEventService from utils.common_util import CamelCaseUtil from utils.excel_util import ExcelUtil @@ -166,6 +167,60 @@ async def add_role_services(cls, query_db: AsyncSession, page_object: AddRoleMod await query_db.rollback() raise e + @staticmethod + async def _role_menu_changed( + query_db: AsyncSession, + page_object: AddRoleModel, + role_info: RoleModel, + ) -> bool: + """ + 比较角色保存前后的菜单授权集合 + + :param query_db: 角色变更使用的异步数据库会话 + :param page_object: 本次角色编辑请求 + :param role_info: 编辑前的角色信息 + :return: 菜单集合发生变化时为 True + """ + if 'menu_ids' not in page_object.model_fields_set: + return False + current_ids = set(await RoleDao.list_role_menu_ids(query_db, int(role_info.role_id))) + requested_ids = {int(menu_id) for menu_id in (page_object.menu_ids or [])} + return current_ids != requested_ids + + @staticmethod + async def _handle_edit_identity_event( + query_db: AsyncSession, + page_object: AddRoleModel, + role_info: RoleModel, + ) -> None: + """ + 在角色安全属性变化时同步失效统一认证身份状态 + + :param query_db: orm对象 + :param page_object: 编辑角色对象 + :param role_info: 当前角色信息 + :return: None + """ + if page_object.type == 'status': + if page_object.status == '1' and role_info.status != '1': + await IdentitySecurityEventService.handle_role_event( + query_db, + page_object.role_id, + 'role_disabled', + actor=page_object.update_by, + ) + return + role_menu_changed = await RoleService._role_menu_changed(query_db, page_object, role_info) + if ( + 'role_key' in page_object.model_fields_set and page_object.role_key != role_info.role_key + ) or role_menu_changed: + await IdentitySecurityEventService.handle_role_event( + query_db, + page_object.role_id, + 'role_claim_changed', + actor=page_object.update_by, + ) + @classmethod async def edit_role_services(cls, query_db: AsyncSession, page_object: AddRoleModel) -> CrudResponseModel: """ @@ -191,6 +246,7 @@ async def edit_role_services(cls, query_db: AsyncSession, page_object: AddRoleMo if not await cls.check_role_key_unique_services(query_db, page_object): raise ServiceException(message=f'修改角色{page_object.role_name}失败,角色权限已存在') try: + await cls._handle_edit_identity_event(query_db, page_object, role_info) await RoleDao.edit_role_dao(query_db, edit_role) if page_object.type != 'status': await RoleDao.delete_role_menu_dao(query_db, RoleMenuModel(roleId=page_object.role_id)) @@ -254,6 +310,12 @@ async def delete_role_services(cls, query_db: AsyncSession, page_object: DeleteR role = await cls.role_detail_services(query_db, int(role_id)) if (await RoleDao.count_user_role_dao(query_db, int(role_id))) > 0: raise ServiceException(message=f'角色{role.role_name}已分配,不能删除') + await IdentitySecurityEventService.handle_role_event( + query_db, + int(role_id), + 'role_deleted', + actor=page_object.update_by, + ) role_id_dict = { 'roleId': role_id, 'updateBy': page_object.update_by, diff --git a/ruoyi-fastapi-backend/module_admin/service/user_service.py b/ruoyi-fastapi-backend/module_admin/service/user_service.py index 2205fb6ed..c7a8cda5f 100644 --- a/ruoyi-fastapi-backend/module_admin/service/user_service.py +++ b/ruoyi-fastapi-backend/module_admin/service/user_service.py @@ -38,6 +38,7 @@ from module_admin.service.dept_service import DeptService from module_admin.service.post_service import PostService from module_admin.service.role_service import RoleService +from module_identity.service.identity_service import IdentitySecurityEventService, IdentitySubjectService from utils.common_util import CamelCaseUtil from utils.excel_util import ExcelUtil from utils.pwd_util import PwdUtil @@ -52,6 +53,37 @@ class UserService: PASSWORD_MIN_LENGTH = 6 PASSWORD_MAX_LENGTH = 20 + @staticmethod + async def _current_role_ids(query_db: AsyncSession, user_id: int) -> set[int]: + """ + 读取用户当前角色ID,用于避免无变化时递增身份版本 + + :param query_db: orm对象 + :param user_id: 用户id + :return: 用户当前角色ID集合 + """ + return await UserDao.get_role_ids(query_db, user_id) + + @staticmethod + async def _handle_role_assignment_changes( + query_db: AsyncSession, user_ids: list[int], actor: str | None = None + ) -> None: + """ + 仅在角色关联实际变化时批量触发身份安全事件 + + :param query_db: orm对象 + :param user_ids: 角色关联发生变化的用户ID列表 + :param actor: 实际执行角色分配的操作者 + :return: None + """ + if user_ids: + await IdentitySecurityEventService.handle_users_event( + query_db, + user_ids, + 'role_assignment_changed', + actor=actor, + ) + @classmethod async def validate_password_services( cls, @@ -208,6 +240,11 @@ async def add_user_services(cls, query_db: AsyncSession, page_object: AddUserMod try: add_result = await UserDao.add_user_dao(query_db, add_user) user_id = add_result.user_id + await IdentitySubjectService.create_for_new_user( + query_db, + user_id=user_id, + create_by=page_object.create_by, + ) if page_object.role_ids: for role in page_object.role_ids: await UserDao.add_user_role_dao(query_db, UserRoleModel(userId=user_id, roleId=role)) @@ -236,6 +273,58 @@ def _deal_edit_user(cls, page_object: EditUserModel, edit_user: dict[str, Any]) else: del edit_user['type'] + @classmethod + async def _handle_edit_identity_security_event( + cls, + query_db: AsyncSession, + page_object: EditUserModel, + current_user: UserInfoModel, + ) -> None: + """ + 将用户编辑类型映射为同事务OIDC安全事件 + + :param query_db: orm对象 + :param page_object: 编辑用户对象 + :param current_user: 编辑前的用户信息 + :return: None + """ + user_id = page_object.user_id + if user_id is None: + raise ServiceException(message='用户不存在') + if page_object.type == 'status': + if page_object.status == '1' and current_user.status != '1': + await IdentitySecurityEventService.handle_user_event( + query_db, + user_id, + 'user_disabled', + actor=page_object.update_by, + ) + return + if page_object.type == 'pwd': + if page_object.password is not None: + await IdentitySecurityEventService.handle_user_event( + query_db, + user_id, + 'password_changed', + actor=page_object.update_by, + ) + return + if page_object.type == 'avatar': + return + + current_roles = {int(value) for value in (current_user.role_ids or '').split(',') if value.isdigit()} + requested_roles = {int(value) for value in (page_object.role_ids or [])} + roles_changed = 'role_ids' in page_object.model_fields_set and current_roles != requested_roles + department_changed = 'dept_id' in page_object.model_fields_set and current_user.dept_id != page_object.dept_id + if roles_changed or department_changed: + event = 'role_assignment_changed' if roles_changed else 'department_changed' + await IdentitySecurityEventService.handle_user_event( + query_db, + user_id, + event, + actor=page_object.update_by, + ) + @classmethod async def edit_user_services(cls, query_db: AsyncSession, page_object: EditUserModel) -> CrudResponseModel: """ @@ -274,6 +363,7 @@ async def edit_user_services(cls, query_db: AsyncSession, page_object: EditUserM await UserDao.add_user_post_dao( query_db, UserPostModel(userId=page_object.user_id, postId=post) ) + await cls._handle_edit_identity_security_event(query_db, page_object, user_info.data) await query_db.commit() return CrudResponseModel(is_success=True, message='更新成功') except Exception as e: @@ -303,6 +393,12 @@ async def delete_user_services(cls, query_db: AsyncSession, page_object: DeleteU await UserDao.delete_user_role_dao(query_db, UserRoleModel(**user_id_dict)) await UserDao.delete_user_post_dao(query_db, UserPostModel(**user_id_dict)) await UserDao.delete_user_dao(query_db, UserModel(**user_id_dict)) + await IdentitySecurityEventService.handle_user_event( + query_db, + int(user_id), + 'user_deleted', + actor=page_object.update_by, + ) await query_db.commit() return CrudResponseModel(is_success=True, message='删除成功') except Exception as e: @@ -420,6 +516,12 @@ async def reset_user_services(cls, query_db: AsyncSession, page_object: ResetUse try: reset_user['password'] = PwdUtil.get_password_hash(page_object.password) await UserDao.edit_user_dao(query_db, reset_user) + await IdentitySecurityEventService.handle_user_event( + query_db, + page_object.user_id, + 'password_changed', + actor=page_object.update_by, + ) await query_db.commit() return CrudResponseModel(is_success=True, message='重置成功') except Exception as e: @@ -541,6 +643,20 @@ async def batch_import_user_services( exclude={'create_time', 'update_time'}, ) await UserDao.edit_user_dao(query_db, edit_user) + if edit_user_model.status == '1' and user_info.status != '1': + await IdentitySecurityEventService.handle_user_event( + query_db, + user_info.user_id, + 'user_disabled', + actor=current_user.user.user_name, + ) + if edit_user_model.dept_id != user_info.dept_id: + await IdentitySecurityEventService.handle_user_event( + query_db, + user_info.user_id, + 'department_changed', + actor=current_user.user.user_name, + ) else: add_error_result.append(f'{count}.用户账号{row["user_name"]}已存在') else: @@ -549,7 +665,12 @@ async def batch_import_user_services( await DeptService.check_dept_data_scope_services( query_db, add_user.dept_id, dept_data_scope_sql ) - await UserDao.add_user_dao(query_db, add_user) + added_user = await UserDao.add_user_dao(query_db, add_user) + await IdentitySubjectService.create_for_new_user( + query_db, + user_id=added_user.user_id, + create_by=current_user.user.user_name, + ) await query_db.commit() return CrudResponseModel(is_success=True, message='\n'.join(add_error_result)) except Exception as e: @@ -646,20 +767,33 @@ async def get_user_role_allocated_list_services( return result @classmethod - async def add_user_role_services(cls, query_db: AsyncSession, page_object: CrudUserRoleModel) -> CrudResponseModel: + async def add_user_role_services( + cls, query_db: AsyncSession, page_object: CrudUserRoleModel, actor: str | None = None + ) -> CrudResponseModel: """ 新增用户关联角色信息service :param query_db: orm对象 :param page_object: 新增用户关联角色对象 + :param actor: 实际执行角色分配的操作者 :return: 新增用户关联角色校验结果 """ if page_object.user_id and page_object.role_ids: role_id_list = page_object.role_ids.split(',') try: + requested_role_ids = {int(role_id) for role_id in role_id_list} + if await cls._current_role_ids(query_db, page_object.user_id) == requested_role_ids: + await query_db.commit() + return CrudResponseModel(is_success=True, message='分配成功') await UserDao.delete_user_role_by_user_and_role_dao(query_db, UserRoleModel(userId=page_object.user_id)) for role_id in role_id_list: await UserDao.add_user_role_dao(query_db, UserRoleModel(userId=page_object.user_id, roleId=role_id)) + await IdentitySecurityEventService.handle_user_event( + query_db, + page_object.user_id, + 'role_assignment_changed', + actor=actor, + ) await query_db.commit() return CrudResponseModel(is_success=True, message='分配成功') except Exception as e: @@ -667,7 +801,16 @@ async def add_user_role_services(cls, query_db: AsyncSession, page_object: CrudU raise e elif page_object.user_id and not page_object.role_ids: try: + if not await cls._current_role_ids(query_db, page_object.user_id): + await query_db.commit() + return CrudResponseModel(is_success=True, message='分配成功') await UserDao.delete_user_role_by_user_and_role_dao(query_db, UserRoleModel(userId=page_object.user_id)) + await IdentitySecurityEventService.handle_user_event( + query_db, + page_object.user_id, + 'role_assignment_changed', + actor=actor, + ) await query_db.commit() return CrudResponseModel(is_success=True, message='分配成功') except Exception as e: @@ -676,6 +819,7 @@ async def add_user_role_services(cls, query_db: AsyncSession, page_object: CrudU elif page_object.user_ids and page_object.role_id: user_id_list = page_object.user_ids.split(',') try: + changed_user_ids: list[int] = [] for user_id in user_id_list: user_role = await cls.detail_user_role_services( query_db, UserRoleModel(userId=user_id, roleId=page_object.role_id) @@ -683,6 +827,8 @@ async def add_user_role_services(cls, query_db: AsyncSession, page_object: CrudU if user_role: continue await UserDao.add_user_role_dao(query_db, UserRoleModel(userId=user_id, roleId=page_object.role_id)) + changed_user_ids.append(int(user_id)) + await cls._handle_role_assignment_changes(query_db, changed_user_ids, actor) await query_db.commit() return CrudResponseModel(is_success=True, message='新增成功') except Exception as e: @@ -693,21 +839,33 @@ async def add_user_role_services(cls, query_db: AsyncSession, page_object: CrudU @classmethod async def delete_user_role_services( - cls, query_db: AsyncSession, page_object: CrudUserRoleModel + cls, query_db: AsyncSession, page_object: CrudUserRoleModel, actor: str | None = None ) -> CrudResponseModel: """ 删除用户关联角色信息service :param query_db: orm对象 :param page_object: 删除用户关联角色对象 + :param actor: 实际执行取消分配的操作者 :return: 删除用户关联角色校验结果 """ if (page_object.user_id and page_object.role_id) or (page_object.user_ids and page_object.role_id): if page_object.user_id and page_object.role_id: try: + existing = await cls.detail_user_role_services( + query_db, + UserRoleModel(userId=page_object.user_id, roleId=page_object.role_id), + ) await UserDao.delete_user_role_by_user_and_role_dao( query_db, UserRoleModel(userId=page_object.user_id, roleId=page_object.role_id) ) + if existing is not None: + await IdentitySecurityEventService.handle_user_event( + query_db, + page_object.user_id, + 'role_assignment_changed', + actor=actor, + ) await query_db.commit() return CrudResponseModel(is_success=True, message='删除成功') except Exception as e: @@ -716,10 +874,18 @@ async def delete_user_role_services( elif page_object.user_ids and page_object.role_id: user_id_list = page_object.user_ids.split(',') try: + changed_user_ids: list[int] = [] for user_id in user_id_list: + existing = await cls.detail_user_role_services( + query_db, + UserRoleModel(userId=user_id, roleId=page_object.role_id), + ) await UserDao.delete_user_role_by_user_and_role_dao( query_db, UserRoleModel(userId=user_id, roleId=page_object.role_id) ) + if existing is not None: + changed_user_ids.append(int(user_id)) + await cls._handle_role_assignment_changes(query_db, changed_user_ids, actor) await query_db.commit() return CrudResponseModel(is_success=True, message='删除成功') except Exception as e: diff --git a/ruoyi-fastapi-backend/module_identity/controller/auth_center_controller.py b/ruoyi-fastapi-backend/module_identity/controller/auth_center_controller.py new file mode 100644 index 000000000..236fb8764 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/auth_center_controller.py @@ -0,0 +1,22 @@ +from fastapi.responses import Response + +from common.router import APIRouterPro +from common.vo import DataResponseModel +from config.env import OidcConfig +from utils.response_util import ResponseUtil + +auth_center_controller = APIRouterPro(tags=['认证中心状态'], order_num=1) + + +@auth_center_controller.get( + '/auth/status', + summary='获取认证中心启用状态', + description='供公共认证页面判断是否允许进入,仅返回功能开关', + response_model=DataResponseModel[dict[str, bool]], +) +async def get_auth_center_status() -> Response: + """匿名读取功能开关,不依赖交互凭据、数据库或签名密钥。""" + return ResponseUtil.success( + data={'enabled': OidcConfig.oidc_enabled}, + headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'}, + ) diff --git a/ruoyi-fastapi-backend/module_identity/controller/authorization_controller.py b/ruoyi-fastapi-backend/module_identity/controller/authorization_controller.py new file mode 100644 index 000000000..f86803dab --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/authorization_controller.py @@ -0,0 +1,388 @@ +import hashlib +from base64 import b64encode +from collections.abc import Iterable +from html import escape +from typing import Annotated, Any, cast +from urllib.parse import urlsplit + +from fastapi import Depends, HTTPException, Request +from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse, Response +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.aspect.db_session import DBSessionDependency +from common.constant import OidcAuditEvent +from common.router import APIRouterPro +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.dependencies import read_form, require_oidc_protocol_ready +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import AuthorizationService +from module_identity.service.infrastructure_service import OidcRateLimiter, RateLimitExceeded, RateLimitUnavailable +from module_identity.service.logout_confirmation_service import LogoutConfirmationService +from module_identity.service.session_service import LogoutService, LogoutServiceError, SsoSessionService +from utils.client_ip_util import ClientIPUtil +from utils.oidc_util import OidcUtil + +authorization_controller = APIRouterPro( + tags=['认证中心协议'], order_num=4, dependencies=[Depends(require_oidc_protocol_ready)] +) +_NO_STORE = {'Cache-Control': 'no-store', 'Pragma': 'no-cache'} +_AUTHORIZE_RATE_LIMIT = 30 +_AUTHORIZE_RATE_WINDOW_SECONDS = 60 +_LOGOUT_STYLE = """ +:root{font-family:-apple-system,BlinkMacSystemFont,"Segoe UI","Microsoft YaHei",sans-serif;color:#1f2937;background:#f3f4f6} +*{box-sizing:border-box}body{margin:0;min-height:100vh;display:grid;place-items:center;padding:24px} +main{width:100%;max-width:480px;background:#fff;border:1px solid #e5e7eb;border-radius:12px;padding:36px} +.brand{font-size:14px;color:#4b5563;margin:0 0 24px}h1{font-size:26px;line-height:1.4;margin:0 0 16px} +p{font-size:16px;line-height:1.75;margin:0 0 28px;color:#4b5563}form{display:flex;gap:12px;flex-wrap:wrap} +button{flex:1;min-height:44px;padding:10px 16px;border:1px solid #cbd5e1;border-radius:6px;background:#fff; +color:#1f2937;font:inherit;font-size:15px;cursor:pointer}button:hover{background:#f3f4f6} +button:focus-visible{outline:3px solid #2563eb;outline-offset:3px} +button[value=confirm]{background:#b91c1c;border-color:#b91c1c;color:#fff}button[value=confirm]:hover{background:#991b1b} +@media(max-width:400px){body{padding:16px}main{padding:24px}form{flex-direction:column}} +""" +_LOGOUT_STYLE_HASH = b64encode(hashlib.sha256(_LOGOUT_STYLE.encode()).digest()).decode() +_LOGOUT_QUERY_LIMITS = {'id_token_hint': 8192, 'post_logout_redirect_uri': 1000, 'state': 2048} +_LOGOUT_QUERY_FIELDS = frozenset(_LOGOUT_QUERY_LIMITS) + + +def _read_authorization_query(request: Request) -> dict[str, str]: + """ + 读取并拒绝重复授权 Query 参数 + + :param request: 当前 HTTP 请求 + :return: 唯一 Query 参数映射 + :raises OAuthProtocolException: Query 参数重复时抛出 + """ + + values: dict[str, str] = {} + # 拒绝重复授权参数 + for key, value in request.query_params.multi_items(): + if key in values: + raise OAuthProtocolException('invalid_request', 'Duplicate authorization parameter') + values[key] = value + return values + + +def _authorization_redis(request: Request) -> Redis: + """ + 读取应用生命周期创建的 Redis 客户端 + + :param request: 当前 HTTP 请求 + :return: Redis 客户端 + :raises OAuthProtocolException: Redis 客户端不可用时抛出 + """ + + redis = getattr(request.app.state, 'redis', None) + if redis is None: + raise OAuthProtocolException('server_error', 'Authorization service is unavailable', 503) + return cast('Redis', redis) + + +async def _enforce_authorization_rate_limit(request: Request, redis: Redis) -> None: + """ + 按 HMAC-IP 固定窗口执行授权请求限流 + + :param request: 当前 HTTP 请求 + :param redis: Redis 客户端 + :return: 无返回值 + :raises OAuthProtocolException: 客户端地址缺失、请求超限或限流服务不可用时抛出 + """ + # 使用 HMAC-IP 固定窗口限流,Redirect 校验前不执行外跳 + ip_address = ClientIPUtil.get_client_ip(request) + if not isinstance(ip_address, str) or not ip_address or ip_address == 'unknown': + raise OAuthProtocolException('temporarily_unavailable', 'Authorization service is unavailable', 503) + ip_hash = OidcUtil.hash_sensitive_identifier(ip_address, OidcConfig.oidc_token_hash_pepper) + try: + await OidcRateLimiter.enforce( + redis, + OidcRedisKey.authorize_ip_rate_limit(ip_hash), + limit=_AUTHORIZE_RATE_LIMIT, + window_seconds=_AUTHORIZE_RATE_WINDOW_SECONDS, + ) + except RateLimitExceeded as exc: + raise OAuthProtocolException( + 'slow_down', 'Too many authorization requests', 429, headers={'Retry-After': str(exc.retry_after)} + ) from exc + except RateLimitUnavailable as exc: + raise OAuthProtocolException('temporarily_unavailable', 'Authorization service is unavailable', 503) from exc + + +async def _record_authorization_failure_audit(db: AsyncSession, **fields: Any) -> None: + """ + 尝试记录授权失败审计,审计异常不影响原协议错误 + + :param db: 异步数据库会话 + :param fields: 审计字段 + :return: 无返回值 + """ + + try: + await AuditService.record_independent( + db, OidcAuditEvent.AUTHORIZE_DENIED, 'failure', risk_level='high', **fields + ) + except Exception: + return + + +def _oidc_not_found_response() -> Response: + """ + 构造 OIDC 关闭时的本地 404 响应 + + :return: OIDC 未启用时的本地 404 响应 + """ + + return JSONResponse(content={'error': 'not_found'}, status_code=404, headers=_NO_STORE) + + +def _logout_page(title: str, body: str, *, redirect_origin: str | None = None) -> HTMLResponse: + """ + 构造仅包含服务端固定内容的退出页面 + + :param title: 页面标题 + :param body: 服务端生成的页面正文 + :param redirect_origin: 已校验的客户端回调源地址 + :return: 带有内容安全策略的退出页面响应 + """ + + form_sources = "'self'" + (f' {redirect_origin}' if redirect_origin else '') + content = ( + '' + '' + f'{escape(title)}' + '
统一认证中心
' + f'

{escape(title)}

{body}
' + ) + + return HTMLResponse( + content, + headers={ + **_NO_STORE, + 'Content-Security-Policy': ( + f"default-src 'none'; style-src 'sha256-{_LOGOUT_STYLE_HASH}'; " + f"form-action {form_sources}; base-uri 'none'; frame-ancestors 'none'" + ), + 'Referrer-Policy': 'strict-origin', + 'X-Frame-Options': 'DENY', + }, + ) + + +def _confirmation_response(token: str, redirect_origin: str | None = None) -> HTMLResponse: + """ + 构造绑定一次性凭据的退出确认页面 + + :param token: 一次性退出确认凭据 + :param redirect_origin: 已校验的客户端回调源地址 + :return: 退出确认页面响应 + """ + + return _logout_page( + '退出认证中心?', + ( + '

确认后,此次登录以及关联应用的离线访问授权将失效。' + '你可以取消,继续保持登录。

' + '
' + f'' + '' + '' + '
' + ), + redirect_origin=redirect_origin, + ) + + +def _local_response(_request: Request) -> HTMLResponse: + """ + 构造不回显退出参数的同源完成页 + + :param _request: 当前 HTTP 请求 + :return: 同源退出完成页 + """ + + return _logout_page('已退出', '

此次登录已结束,可以关闭此页面。

') + + +async def _logout_parameters(request: Request) -> dict[str, str] | None: + """ + 读取并校验 GET Query 或 POST Form 退出参数 + + :param request: 当前 HTTP 请求 + :return: 通过校验的退出参数,校验失败时返回 None + """ + + if request.method == 'POST': + try: + values = await read_form(request) + except HTTPException: + return None + return values if _valid_logout_parameters(values.items()) else None + seen: set[str] = set() + for key, value in request.query_params.multi_items(): + if key not in _LOGOUT_QUERY_FIELDS or key in seen or len(value) > _LOGOUT_QUERY_LIMITS[key]: + return None + seen.add(key) + return dict(request.query_params.multi_items()) + + +def _valid_logout_parameters(values: Iterable[tuple[str, str]]) -> bool: + """ + 校验退出参数的字段白名单和长度 + + :param values: 待校验的退出参数 + :return: 参数是否全部通过校验 + """ + + return all(key in _LOGOUT_QUERY_FIELDS and len(value) <= _LOGOUT_QUERY_LIMITS[key] for key, value in values) + + +@authorization_controller.api_route( + '/oauth2/authorize', + methods=['GET', 'POST'], + summary='OAuth 授权接口', + description='用于处理 OAuth 授权请求并返回认证交互或客户端回调', + include_in_schema=False, +) +async def authorize( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + if not OidcConfig.oidc_enabled: + return _oidc_not_found_response() + if request.method == 'POST': + if request.query_params: + raise OAuthProtocolException('invalid_request', 'Authorization parameters must be in the form body') + try: + raw = await read_form(request) + except HTTPException as exc: + raise OAuthProtocolException('invalid_request', 'Invalid authorization form', exc.status_code) from exc + else: + raw = _read_authorization_query(request) + redis = _authorization_redis(request) + try: + await _enforce_authorization_rate_limit(request, redis) + except OAuthProtocolException as exc: + await _record_authorization_failure_audit(query_db, client_id=raw.get('client_id'), failure_code=exc.error) + raise + try: + result = await AuthorizationService.process_authorization_request( + query_db, + redis, + raw, + sso_cookie=request.cookies.get(OidcConfig.oidc_sso_cookie_name), + ) + except OAuthProtocolException as exc: + if exc.error == 'not_found' and not exc.can_redirect: + return _oidc_not_found_response() + raise + return RedirectResponse(result.location, status_code=result.status_code, headers=_NO_STORE) + + +@authorization_controller.api_route( + '/oauth2/logout', + methods=['GET', 'POST'], + summary='退出认证中心接口', + description='用于退出认证中心登录', + include_in_schema=False, +) +async def logout( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + if not OidcConfig.oidc_enabled: + return _oidc_not_found_response() + parameters = await _logout_parameters(request) + if parameters is None: + return JSONResponse(content={'error': 'invalid_request'}, status_code=400, headers=_NO_STORE) + redis = getattr(request.app.state, 'redis', None) + try: + if redis is None: + raise LogoutServiceError('退出服务暂不可用') + address = ClientIPUtil.get_client_ip(request) + rate_scope = OidcUtil.logout_rate_scope(address, getattr(OidcConfig, 'oidc_token_hash_pepper', '')) + await OidcRateLimiter.enforce(redis, OidcRedisKey.logout_rate_limit(rate_scope), limit=60, window_seconds=60) + token, nonce = await LogoutConfirmationService.issue( + redis, parameters, request.cookies.get(OidcConfig.oidc_sso_cookie_name) + ) + redirect_origin = await LogoutConfirmationService.form_redirect_origin(query_db, parameters) + response = _confirmation_response(token, redirect_origin) + response.set_cookie( + LogoutConfirmationService.COOKIE_NAME, + nonce, + max_age=LogoutConfirmationService.TTL_SECONDS, + path='/', + secure=True, + httponly=True, + samesite='strict', + ) + return response + except RateLimitExceeded as exc: + headers = dict(_NO_STORE) + headers['Retry-After'] = str(exc.retry_after) + return JSONResponse(content={'error': 'temporarily_unavailable'}, status_code=429, headers=headers) + except RateLimitUnavailable: + return JSONResponse(content={'error': 'temporarily_unavailable'}, status_code=503, headers=_NO_STORE) + except Exception: + return JSONResponse(content={'error': 'temporarily_unavailable'}, status_code=503, headers=_NO_STORE) + + +@authorization_controller.post( + '/oauth2/logout/confirm', + summary='确认退出认证中心接口', + description='用于确认或取消当前浏览器的认证中心退出请求', + include_in_schema=False, +) +async def confirm_logout( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + if not OidcConfig.oidc_enabled: + return _oidc_not_found_response() + issuer = urlsplit(OidcConfig.oidc_issuer) + if request.headers.get('origin') != f'{issuer.scheme}://{issuer.netloc}': + return JSONResponse({'error': 'invalid_request'}, status_code=403, headers=_NO_STORE) + try: + form = await read_form(request) + if set(form) != {'confirmation', 'decision'} or form['decision'] not in {'confirm', 'cancel'}: + raise ValueError('确认信息无效') + redis = request.app.state.redis + parameters = await LogoutConfirmationService.consume( + redis, + form['confirmation'], + request.cookies.get(OidcConfig.oidc_sso_cookie_name), + request.cookies.get(LogoutConfirmationService.COOKIE_NAME), + ) + except (HTTPException, ValueError, KeyError, TypeError): + return JSONResponse({'error': 'invalid_request'}, status_code=400, headers=_NO_STORE) + except Exception: + return JSONResponse({'error': 'temporarily_unavailable'}, status_code=503, headers=_NO_STORE) + if form['decision'] == 'cancel': + response = _logout_page('已取消退出', '

登录状态保持不变,可以关闭此页面。

') + else: + try: + result = await LogoutService.execute_logout( + query_db, + redis, + confirmed=True, + id_token_hint=parameters.get('id_token_hint'), + cookie=request.cookies.get(OidcConfig.oidc_sso_cookie_name), + post_logout_redirect_uri=parameters.get('post_logout_redirect_uri'), + state=parameters.get('state'), + ) + except Exception: + return JSONResponse({'error': 'temporarily_unavailable'}, status_code=503, headers=_NO_STORE) + response = ( + _local_response(request) + if result.redirect_uri is None + else RedirectResponse( + OidcUtil.append_state(result.redirect_uri, result.state), status_code=303, headers=_NO_STORE + ) + ) + response.delete_cookie(**SsoSessionService.cookie_parameters()) + response.delete_cookie( + LogoutConfirmationService.COOKIE_NAME, path='/', secure=True, httponly=True, samesite='strict' + ) + + return response diff --git a/ruoyi-fastapi-backend/module_identity/controller/discovery_controller.py b/ruoyi-fastapi-backend/module_identity/controller/discovery_controller.py new file mode 100644 index 000000000..9618915ce --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/discovery_controller.py @@ -0,0 +1,112 @@ +from typing import Annotated, Any + +from fastapi import Depends, Request +from fastapi.responses import JSONResponse, Response +from sqlalchemy.ext.asyncio import AsyncSession + +from common.aspect.db_session import DBSessionDependency +from common.router import APIRouterPro +from config.env import OidcConfig +from module_identity.dependencies import require_oidc_protocol_ready +from module_identity.service.discovery_service import DiscoveryService +from module_identity.service.key_service import KeyService, KeyServiceError +from utils.oidc_util import OidcUtil + +discovery_controller = APIRouterPro( + tags=['认证中心发现'], order_num=1, dependencies=[Depends(require_oidc_protocol_ready)] +) +_CACHE_CONTROL = 'public, max-age=300' + + +def _disabled_response() -> Response: + """ + 返回协议端点关闭时的标准 404 响应 + + :return: 协议端点关闭时的标准 404 响应 + """ + + return JSONResponse( + content={'error': 'not_found'}, status_code=404, headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'} + ) + + +def _etag_response(request: Request, payload: dict[str, Any]) -> Response: + """ + 按 If-None-Match 返回标准 JSON 或 304 响应 + + :param request: 当前 HTTP 请求 + :param payload: 待返回的 JSON 数据 + :return: 标准 JSON 响应或 304 响应 + """ + + etag = OidcUtil.json_etag(payload) + headers = {'Cache-Control': _CACHE_CONTROL, 'ETag': etag} + if_none_match = request.headers.get('if-none-match', '') + if '*' in {item.strip() for item in if_none_match.split(',')} or etag in { + item.strip() for item in if_none_match.split(',') + }: + return Response(status_code=304, headers=headers) + return JSONResponse(content=payload, headers=headers) + + +@discovery_controller.get( + '/.well-known/openid-configuration', + summary='获取 OpenID 配置接口', + description='用于返回 OpenID Connect Discovery 元数据', + include_in_schema=False, +) +async def openid_configuration( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + if not OidcConfig.oidc_enabled: + return _disabled_response() + try: + payload = await DiscoveryService.openid_metadata_with_scopes(query_db) + except Exception: + return JSONResponse( + content={'error': 'temporarily_unavailable'}, status_code=503, headers={'Cache-Control': 'no-store'} + ) + return _etag_response(request, payload) + + +@discovery_controller.get( + '/.well-known/oauth-authorization-server', + summary='获取 OAuth 授权服务器元数据接口', + description='用于返回 OAuth 授权服务器元数据', + include_in_schema=False, +) +async def oauth_authorization_server_metadata( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + if not OidcConfig.oidc_enabled: + return _disabled_response() + try: + payload = await DiscoveryService.oauth_metadata_with_scopes(query_db) + except Exception: + return JSONResponse( + content={'error': 'temporarily_unavailable'}, status_code=503, headers={'Cache-Control': 'no-store'} + ) + return _etag_response(request, payload) + + +@discovery_controller.get( + '/oauth2/jwks', + summary='获取签名公钥接口', + description='用于返回不含私钥材料的签名公钥集合', + include_in_schema=False, +) +async def jwks(request: Request, query_db: Annotated[AsyncSession, DBSessionDependency()]) -> Response: + if not OidcConfig.oidc_enabled: + return _disabled_response() + try: + # 仅返回签名公钥,不暴露私钥材料 + payload = await KeyService.build_jwks(query_db) + except KeyServiceError: + return JSONResponse( + content={'error': 'temporarily_unavailable'}, + status_code=503, + headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'}, + ) + return _etag_response(request, payload) diff --git a/ruoyi-fastapi-backend/module_identity/controller/interaction_controller.py b/ruoyi-fastapi-backend/module_identity/controller/interaction_controller.py new file mode 100644 index 000000000..f6d2ad011 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/interaction_controller.py @@ -0,0 +1,317 @@ +import json +from typing import Annotated, TypeVar + +from fastapi import Depends, Header, Request +from fastapi.responses import Response +from pydantic import BaseModel, ValidationError +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.aspect.db_session import DBSessionDependency +from common.router import APIRouterPro +from common.vo import DataResponseModel +from config.env import OidcConfig +from exceptions.exception import OidcInteractionException +from module_identity.dependencies import require_oidc_protocol_ready +from module_identity.entity.vo.interaction_vo import ( + CaptchaResponseModel, + ChangePasswordModel, + InteractionConsentModel, + InteractionLoginModel, + InteractionResultModel, + InteractionViewModel, +) +from module_identity.service.authorization_service import InteractionCompletionService +from module_identity.service.consent_service import InteractionConsentService +from module_identity.service.interaction_service import ( + InteractionFlowService, + InteractionLoginOutcome, + InteractionLoginService, + InteractionService, +) +from module_identity.service.session_service import SsoSessionError, SsoSessionService +from utils.client_ip_util import ClientIPUtil +from utils.oidc_util import OidcUtil +from utils.response_util import ResponseUtil + +interaction_controller = APIRouterPro( + tags=['认证中心交互'], order_num=5, dependencies=[Depends(require_oidc_protocol_ready)] +) +_MODEL = TypeVar('_MODEL', bound=BaseModel) +_MAX_INTERACTION_BODY_BYTES = 16 * 1024 +_NO_STORE = {'Cache-Control': 'no-store', 'Pragma': 'no-cache'} + + +def _redis(request: Request) -> Redis: + """ + 读取应用生命周期创建的 Redis 客户端 + + :param request: 当前 HTTP 请求 + :return: 应用 Redis 客户端 + :raises OidcInteractionException: Redis 客户端不可用 + """ + + value = getattr(request.app.state, 'redis', None) + if value is None: + raise OidcInteractionException(error='server_error', status_code=503, message='认证服务不可用') + return value + + +def _not_found() -> Response: + """ + 构造认证中心未启用时的响应 + + :return: 认证中心未启用的 404 响应 + """ + + return _failure_response('认证服务未启用', 404) + + +def _invalid_body() -> Response: + """ + 构造认证交互请求无效时的响应 + + :return: 请求无效的 400 响应 + """ + + return _failure_response('认证交互请求无效', 400) + + +def _failure_response(msg: str, status_code: int, headers: dict[str, str] | None = None) -> Response: + """ + 构造项目统一错误响应并保留安全 HTTP 状态码 + + :param msg: 错误消息 + :param status_code: HTTP 状态码 + :param headers: 可选响应头 + :return: 统一错误响应 + """ + + response = ResponseUtil.failure(msg=msg, headers=headers or _NO_STORE) + response.status_code = status_code + + return response + + +def _login_response(outcome: InteractionLoginOutcome) -> Response: + """ + 构造认证登录或改密响应 + + :param outcome: 登录流程结果 + :return: 统一认证响应 + """ + + if outcome.failure_message is not None: + return ResponseUtil.failure(msg=outcome.failure_message, headers=_NO_STORE) + response = ResponseUtil.success(data=outcome.result, headers=_NO_STORE) + if outcome.cookie is not None: + try: + OidcUtil.parse_sso_cookie(outcome.cookie) + except ValueError as exc: + raise SsoSessionError(str(exc)) from exc + response.set_cookie(value=outcome.cookie, **SsoSessionService.cookie_parameters(max_age=outcome.cookie_max_age)) + return response + + +async def _safe_json_body(request: Request, model: type[_MODEL]) -> _MODEL: + """ + 读取并校验认证交互 JSON 请求体 + + :param request: 当前 HTTP 请求 + :param model: 请求体模型类型 + :return: 校验后的请求体模型 + :raises OidcInteractionException: 请求体格式或内容无效 + """ + # 仅接受 JSON,避免协议请求被宽松解析 + content_type = request.headers.get('content-type', '').split(';', 1)[0].strip().lower() + if content_type != 'application/json': + raise OidcInteractionException(error='invalid_request', status_code=400, message='认证交互请求无效') + try: + chunks: list[bytes] = [] + size = 0 + async for chunk in request.stream(): + if not isinstance(chunk, (bytes, bytearray)): + raise ValueError('请求体必须为字节数据') + # 限制请求体大小,避免认证交互接口被大请求消耗资源 + size += len(chunk) + if size > _MAX_INTERACTION_BODY_BYTES: + raise ValueError('请求体大小超过允许范围') + chunks.append(bytes(chunk)) + raw = b''.join(chunks) + except Exception as exc: + raise OidcInteractionException(error='invalid_request', status_code=400, message='认证交互请求无效') from exc + if len(raw) > _MAX_INTERACTION_BODY_BYTES: + raise OidcInteractionException(error='invalid_request', status_code=400, message='认证交互请求无效') + try: + value = json.loads(raw.decode('utf-8'), object_pairs_hook=OidcUtil.json_object_pairs) + if not isinstance(value, dict): + raise ValueError('JSON 请求体必须为对象') + return model.model_validate(value) + except (UnicodeDecodeError, json.JSONDecodeError, TypeError, ValueError, ValidationError) as exc: + raise OidcInteractionException(error='invalid_request', status_code=400, message='认证交互请求无效') from exc + + +@interaction_controller.get( + '/auth/interaction/{interaction_id}', + summary='获取认证交互页面接口', + description='用于获取认证交互页面信息', + response_model=DataResponseModel[InteractionViewModel], +) +async def get_interaction( + request: Request, + interaction_id: str, + query_db: Annotated[AsyncSession, DBSessionDependency()], + csrf_token: str | None = Header(default=None, alias='X-CSRF-Token'), +) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + redis = _redis(request) + await InteractionFlowService.csrf_record(redis, interaction_id, csrf_token) + page = await InteractionService.get(redis, interaction_id, query_db) + page['captchaEnabled'] = await InteractionFlowService.captcha_enabled(redis) + + return ResponseUtil.success(data=InteractionViewModel.model_validate(page), headers=_NO_STORE) + + +@interaction_controller.get( + '/auth/interaction/{interaction_id}/captcha', + summary='获取认证交互验证码接口', + description='用于获取认证交互验证码', + response_model=DataResponseModel[CaptchaResponseModel], +) +async def captcha(request: Request, interaction_id: str) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + outcome = await InteractionFlowService.captcha(_redis(request), interaction_id, ClientIPUtil.get_client_ip(request)) + if outcome.rate_limited: + return _failure_response( + '请求过于频繁,请稍后再试', + 429, + {**_NO_STORE, 'Retry-After': str(outcome.retry_after)}, + ) + if outcome.unavailable: + return _failure_response('认证服务暂不可用', 503) + return ResponseUtil.success(data=outcome.result, headers=_NO_STORE) + + +@interaction_controller.post( + '/auth/interaction/{interaction_id}/login', + summary='提交认证中心登录接口', + description='用于提交认证中心登录', + response_model=DataResponseModel[InteractionResultModel], +) +async def login_endpoint( + request: Request, + interaction_id: str, + query_db: Annotated[AsyncSession, DBSessionDependency()], + csrf_token: str | None = Header(default=None, alias='X-CSRF-Token'), +) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + try: + body = await _safe_json_body(request, InteractionLoginModel) + except OidcInteractionException: + return _invalid_body() + outcome = await InteractionLoginService.login( + _redis(request), + interaction_id, + body, + query_db, + csrf_token, + ClientIPUtil.get_client_ip(request), + request.headers.get('user-agent'), + ) + + return _login_response(outcome) + + +@interaction_controller.post( + '/auth/interaction/{interaction_id}/change-password', + summary='提交认证中心改密接口', + description='用于提交认证中心密码修改', + response_model=DataResponseModel[InteractionResultModel], +) +async def change_password_endpoint( + request: Request, + interaction_id: str, + query_db: Annotated[AsyncSession, DBSessionDependency()], + csrf_token: str | None = Header(default=None, alias='X-CSRF-Token'), +) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + try: + body = await _safe_json_body(request, ChangePasswordModel) + except OidcInteractionException: + return _invalid_body() + outcome = await InteractionLoginService.change_password(_redis(request), interaction_id, body, query_db, csrf_token) + + return _login_response(outcome) + + +@interaction_controller.post( + '/auth/interaction/{interaction_id}/consent', + summary='提交认证中心授权同意接口', + description='用于提交认证中心授权同意', + response_model=DataResponseModel[InteractionResultModel], +) +async def consent_endpoint( + request: Request, + interaction_id: str, + query_db: Annotated[AsyncSession, DBSessionDependency()], + csrf_token: str | None = Header(default=None, alias='X-CSRF-Token'), +) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + try: + body = await _safe_json_body(request, InteractionConsentModel) + except OidcInteractionException: + return _invalid_body() + result = await InteractionConsentService.consent(_redis(request), interaction_id, body, query_db, csrf_token) + + return ResponseUtil.success(data=result, headers=_NO_STORE) + + +@interaction_controller.post( + '/auth/interaction/{interaction_id}/cancel', + summary='取消认证中心授权接口', + description='用于取消认证中心授权', + response_model=DataResponseModel[InteractionResultModel], +) +async def cancel( + request: Request, + interaction_id: str, + query_db: Annotated[AsyncSession, DBSessionDependency()], + csrf_token: str | None = Header(default=None, alias='X-CSRF-Token'), +) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + result = await InteractionConsentService.cancel(_redis(request), interaction_id, query_db, csrf_token) + + return ResponseUtil.success(data=result, headers=_NO_STORE) + + +@interaction_controller.post( + '/auth/interaction/{interaction_id}/complete', + summary='完成认证中心交互接口', + description='用于完成认证交互并返回外部客户端跳转地址', + response_model=DataResponseModel[InteractionResultModel], +) +async def complete( + request: Request, + interaction_id: str, + query_db: Annotated[AsyncSession, DBSessionDependency()], + csrf_token: str | None = Header(default=None, alias='X-CSRF-Token'), +) -> Response: + if not OidcConfig.oidc_enabled: + return _not_found() + await InteractionFlowService.csrf_record(_redis(request), interaction_id, csrf_token) + result = await InteractionCompletionService.complete(_redis(request), interaction_id, query_db) + + return ResponseUtil.success( + data=InteractionResultModel( + next_action='redirect', + interaction_id=interaction_id, + redirect_url=result.location, + ), + headers=_NO_STORE, + ) diff --git a/ruoyi-fastapi-backend/module_identity/controller/oauth_audit_controller.py b/ruoyi-fastapi-backend/module_identity/controller/oauth_audit_controller.py new file mode 100644 index 000000000..ea3a13023 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/oauth_audit_controller.py @@ -0,0 +1,61 @@ +from typing import Annotated + +from fastapi import Form, Query, Request, Response +from fastapi.responses import StreamingResponse +from sqlalchemy.ext.asyncio import AsyncSession + +from common.annotation.log_annotation import Log +from common.annotation.rate_limit_annotation import ApiRateLimit, ApiRateLimitPreset +from common.aspect.db_session import DBSessionDependency +from common.aspect.interface_auth import UserInterfaceAuthDependency +from common.aspect.pre_auth import PreAuthDependency +from common.constant import ApiNamespace +from common.enums import BusinessType +from common.router import APIRouterPro +from common.vo import PageResponseModel +from module_identity.entity.vo.oauth_session_vo import AuditModel, AuditPageQueryModel +from module_identity.service.audit_service import AuditService +from utils.common_util import bytes2file_response +from utils.response_util import ResponseUtil + +oauth_audit_controller = APIRouterPro( + prefix='/monitor/oauth/audit', order_num=25, tags=['监控管理-OAuth 审计'], dependencies=[PreAuthDependency()] +) + + +@oauth_audit_controller.get( + '/list', + summary='获取 OAuth 审计分页列表接口', + description='用于获取 OAuth 审计分页列表', + response_model=PageResponseModel[AuditModel], + dependencies=[UserInterfaceAuthDependency('monitor:oauthAudit:list')], +) +async def list_oauth_audit( + query: Annotated[AuditPageQueryModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + rows, total = await AuditService.list_admin_page(query_db, query) + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_audit_controller.post( + '/export', + summary='导出 OAuth 审计接口', + description='用于导出 OAuth 审计数据', + response_class=StreamingResponse, + dependencies=[UserInterfaceAuthDependency('monitor:oauthAudit:export')], +) +@ApiRateLimit(namespace=ApiNamespace.MONITOR_OAUTH_AUDIT_EXPORT, preset=ApiRateLimitPreset.USER_RESOURCE_EXPORT) +@Log(title='OAuth 审计', business_type=BusinessType.EXPORT) +async def export_oauth_audit( + request: Request, + query: Annotated[AuditPageQueryModel, Form()], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + data = await AuditService.export_admin(query_db, query) + + return ResponseUtil.streaming( + data=bytes2file_response(data), + media_type='application/vnd.openxmlformats-officedocument.spreadsheetml.sheet', + headers={'Content-Disposition': 'attachment; filename="oauth-audit.xlsx"'}, + ) diff --git a/ruoyi-fastapi-backend/module_identity/controller/oauth_client_controller.py b/ruoyi-fastapi-backend/module_identity/controller/oauth_client_controller.py new file mode 100644 index 000000000..a6b194e6c --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/oauth_client_controller.py @@ -0,0 +1,305 @@ +from typing import Annotated + +from fastapi import Body, Path, Query, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + +from common.annotation.log_annotation import Log +from common.annotation.rate_limit_annotation import ApiRateLimit, ApiRateLimitPreset +from common.aspect.db_session import DBSessionDependency +from common.aspect.interface_auth import UserInterfaceAuthDependency +from common.aspect.pre_auth import CurrentUserDependency, PreAuthDependency +from common.constant import ApiNamespace +from common.enums import BusinessType +from common.router import APIRouterPro +from common.vo import DataResponseModel, PageResponseModel, ResponseBaseModel +from exceptions.exception import ServiceException +from module_admin.entity.vo.user_vo import CurrentUserModel +from module_identity.entity.vo.oauth_client_vo import ( + ClientCreateModel, + ClientPageQueryModel, + ClientSecretResponseModel, + ClientStatusModel, + ClientUpdateModel, + ClientUriModel, + ClientViewModel, + SecretRotationModel, +) +from module_identity.service.oauth_management_service import OAuthClientManagementService +from module_identity.service.runtime_service import OidcRuntimeService +from utils.log_util import logger +from utils.oidc_util import OidcUtil +from utils.response_util import ResponseUtil + +oauth_client_controller = APIRouterPro( + prefix='/system/oauth/client', order_num=21, tags=['系统管理-OAuth Client'], dependencies=[PreAuthDependency()] +) +_MAX_BATCH_SIZE = 100 + + +def _actor(current_user: CurrentUserModel) -> str: + """ + 提取管理操作者标识 + + :param current_user: 当前登录用户 + :return: 安全截断后的用户名 + :raises ServiceException: 当前用户不可用 + """ + + try: + return OidcUtil.actor_name( + getattr(getattr(current_user, 'user', None), 'user_name', None), error_message='当前操作者不可用' + ) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + + +@oauth_client_controller.get( + '/list', + summary='获取 OAuth 客户端分页列表接口', + description='用于获取 OAuth 客户端分页列表', + response_model=PageResponseModel[ClientViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthClient:list')], +) +async def get_system_oauth_client_list( + client_query: Annotated[ClientPageQueryModel, Query()], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + rows = await OAuthClientManagementService.list_clients(query_db, client_query) + total = await OAuthClientManagementService.count_clients(query_db, client_query) + logger.info('OAuth 客户端列表查询成功') + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_client_controller.get( + '/{client_id}', + summary='查询 OAuth 客户端详情接口', + description='用于查询 OAuth 客户端详情', + response_model=DataResponseModel[ClientViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthClient:query')], +) +async def query_system_oauth_client( + client_id: Annotated[str, Path(min_length=1, max_length=64)], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + return ResponseUtil.success(data=await OAuthClientManagementService.detail(query_db, client_id)) + + +@oauth_client_controller.post( + '', + summary='新增 OAuth 客户端接口', + description='用于新增 OAuth 客户端', + response_model=DataResponseModel[ClientViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthClient:add')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_CREATE, preset=ApiRateLimitPreset.USER_COMMON_MUTATION) +@Log(title='OAuth 客户端管理', business_type=BusinessType.INSERT) +async def add_system_oauth_client( + request: Request, + payload: ClientCreateModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthClientManagementService.create_client( + query_db, + payload, + _actor(current_user), + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + logger.info('OAuth 客户端创建成功') + + return ResponseUtil.success(data=result) + + +@oauth_client_controller.put( + '', + summary='编辑 OAuth 客户端接口', + description='用于编辑 OAuth 客户端', + response_model=DataResponseModel[ClientViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthClient:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_UPDATE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 客户端管理', business_type=BusinessType.UPDATE) +async def edit_system_oauth_client( + request: Request, + payload: ClientUpdateModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthClientManagementService.update_client( + query_db, + payload, + _actor(current_user), + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(data=result) + + +@oauth_client_controller.delete( + '/{client_ids}', + summary='批量停用 OAuth 客户端接口', + description='用于批量停用 OAuth 客户端', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthClient:remove')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_DISABLE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +@Log(title='OAuth 客户端管理', business_type=BusinessType.DELETE) +async def delete_system_oauth_clients( + request: Request, + client_ids: Annotated[str, Path(min_length=1, max_length=6500)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + actor = _actor(current_user) + try: + batch_ids = OidcUtil.split_batch(client_ids, 'client_ids', max_size=_MAX_BATCH_SIZE) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + await OAuthClientManagementService.disable_clients( + query_db, + batch_ids, + actor, + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(msg='OAuth 客户端已停用') + + +@oauth_client_controller.put( + '/changeStatus', + summary='启停 OAuth 客户端接口', + description='用于启用或停用 OAuth 客户端', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthClient:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_STATUS, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 客户端管理', business_type=BusinessType.UPDATE) +async def change_system_oauth_client_status( + request: Request, + payload: ClientStatusModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthClientManagementService.change_client_status( + query_db, + payload, + _actor(current_user), + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(data=result) + + +@oauth_client_controller.post( + '/{client_id}/secret', + summary='轮换 OAuth 客户端密钥接口', + description='用于轮换 OAuth 客户端密钥', + response_model=DataResponseModel[ClientSecretResponseModel], + dependencies=[UserInterfaceAuthDependency('system:oauthClient:rotateSecret')], +) +@ApiRateLimit( + namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_SECRET_ROTATE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION +) +@Log(title='OAuth 客户端密钥管理', business_type=BusinessType.UPDATE) +async def rotate_system_oauth_client_secret( + request: Request, + client_id: Annotated[str, Path(min_length=1, max_length=64)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], + payload: SecretRotationModel = Body(default_factory=SecretRotationModel), +) -> Response: + result = await OAuthClientManagementService.rotate_secret( + query_db, + client_id, + actor=_actor(current_user), + not_before=payload.not_before, + expires_at=payload.expires_at, + retirement_seconds=payload.retirement_seconds, + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(msg='客户端密钥生成成功,请立即保存', data=result) + + +@oauth_client_controller.delete( + '/{client_id}/secret/{secret_id}', + summary='撤销 OAuth 客户端密钥接口', + description='用于撤销 OAuth 客户端密钥', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthClient:rotateSecret')], +) +@ApiRateLimit( + namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_SECRET_REVOKE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION +) +@Log(title='OAuth 客户端密钥管理', business_type=BusinessType.DELETE) +async def revoke_system_oauth_client_secret( + request: Request, + client_id: Annotated[str, Path(min_length=1, max_length=64)], + secret_id: Annotated[str, Path(min_length=1, max_length=64)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthClientManagementService.revoke_secret( + query_db, + client_id, + secret_id, + _actor(current_user), + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(data=result, msg='客户端密钥已撤销') + + +@oauth_client_controller.post( + '/{client_id}/uri', + summary='新增 OAuth 客户端 URI 接口', + description='用于新增 OAuth 客户端 URI', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthClient:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_URI_ADD, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 客户端管理', business_type=BusinessType.INSERT) +async def add_system_oauth_client_uri( + request: Request, + client_id: Annotated[str, Path(min_length=1, max_length=64)], + payload: ClientUriModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthClientManagementService.add_uri( + query_db, + client_id, + payload, + _actor(current_user), + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(data=result, msg='客户端 URI 已增加') + + +@oauth_client_controller.delete( + '/{client_id}/uri/{uri_id}', + summary='停用 OAuth 客户端 URI 接口', + description='用于停用 OAuth 客户端 URI', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthClient:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_CLIENT_URI_REMOVE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 客户端管理', business_type=BusinessType.DELETE) +async def delete_system_oauth_client_uri( + request: Request, + client_id: Annotated[str, Path(min_length=1, max_length=64)], + uri_id: Annotated[int, Path(gt=0)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthClientManagementService.remove_uri( + query_db, + client_id, + uri_id, + _actor(current_user), + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(data=result, msg='客户端 URI 已停用') diff --git a/ruoyi-fastapi-backend/module_identity/controller/oauth_resource_controller.py b/ruoyi-fastapi-backend/module_identity/controller/oauth_resource_controller.py new file mode 100644 index 000000000..8ecd455e2 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/oauth_resource_controller.py @@ -0,0 +1,309 @@ +from typing import Annotated + +from fastapi import Path, Query, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + +from common.annotation.log_annotation import Log +from common.annotation.rate_limit_annotation import ApiRateLimit, ApiRateLimitPreset +from common.aspect.db_session import DBSessionDependency +from common.aspect.interface_auth import UserInterfaceAuthDependency +from common.aspect.pre_auth import CurrentUserDependency, PreAuthDependency +from common.constant import ApiNamespace +from common.enums import BusinessType +from common.router import APIRouterPro +from common.vo import DataResponseModel, PageResponseModel, ResponseBaseModel +from exceptions.exception import ServiceException +from module_admin.entity.vo.user_vo import CurrentUserModel +from module_identity.entity.vo.oauth_resource_vo import ( + ResourceCreateModel, + ResourcePageQueryModel, + ResourceStatusModel, + ResourceUpdateModel, + ResourceViewModel, + ScopeModel, + ScopePageQueryModel, + ScopeStatusModel, +) +from module_identity.service.oauth_management_service import OAuthResourceManagementService +from module_identity.service.runtime_service import OidcRuntimeService +from utils.log_util import logger +from utils.oidc_util import OidcUtil +from utils.response_util import ResponseUtil + +oauth_resource_controller = APIRouterPro( + prefix='/system/oauth/resource', order_num=22, tags=['系统管理-OAuth Resource'], dependencies=[PreAuthDependency()] +) +oauth_scope_controller = APIRouterPro( + prefix='/system/oauth/scope', order_num=23, tags=['系统管理-OAuth Scope'], dependencies=[PreAuthDependency()] +) +_MAX_BATCH_SIZE = 100 + + +def _actor(current_user: CurrentUserModel) -> str: + """ + 提取管理操作者标识 + + :param current_user: 当前登录用户 + :return: 安全截断后的用户名 + :raises ServiceException: 当前用户不可用 + """ + + try: + return OidcUtil.actor_name( + getattr(getattr(current_user, 'user', None), 'user_name', None), error_message='当前操作者不可用' + ) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + + +@oauth_resource_controller.get( + '/list', + summary='获取 OAuth 资源分页列表接口', + description='用于获取 OAuth 资源分页列表', + response_model=PageResponseModel[ResourceViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthResource:list')], +) +async def get_system_oauth_resource_list( + resource_query: Annotated[ResourcePageQueryModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + rows = await OAuthResourceManagementService.list_resources(query_db, resource_query) + total = await OAuthResourceManagementService.count_resources(query_db, resource_query) + logger.info('OAuth 资源列表查询成功') + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_resource_controller.get( + '/{resource_id}', + summary='查询 OAuth 资源详情接口', + description='用于查询 OAuth 资源详情', + response_model=DataResponseModel[ResourceViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthResource:list')], +) +async def query_system_oauth_resource( + resource_id: Annotated[str, Path(min_length=1, max_length=64)], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + return ResponseUtil.success(data=await OAuthResourceManagementService.detail_resource(query_db, resource_id)) + + +@oauth_resource_controller.post( + '', + summary='新增 OAuth 资源接口', + description='用于新增 OAuth 资源', + response_model=DataResponseModel[ResourceViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthResource:add')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_RESOURCE_CREATE, preset=ApiRateLimitPreset.USER_COMMON_MUTATION) +@Log(title='OAuth 资源管理', business_type=BusinessType.INSERT) +async def add_system_oauth_resource( + request: Request, + payload: ResourceCreateModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + return ResponseUtil.success( + data=await OAuthResourceManagementService.create_resource( + query_db, payload, _actor(current_user), after_commit=OidcRuntimeService.cors_snapshot_callback(request.app) + ) + ) + + +@oauth_resource_controller.put( + '', + summary='编辑 OAuth 资源接口', + description='用于编辑 OAuth 资源', + response_model=DataResponseModel[ResourceViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthResource:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_RESOURCE_UPDATE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 资源管理', business_type=BusinessType.UPDATE) +async def edit_system_oauth_resource( + request: Request, + payload: ResourceUpdateModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + return ResponseUtil.success( + data=await OAuthResourceManagementService.update_resource( + query_db, payload, _actor(current_user), after_commit=OidcRuntimeService.cors_snapshot_callback(request.app) + ) + ) + + +@oauth_resource_controller.delete( + '/{resource_ids}', + summary='批量停用 OAuth 资源接口', + description='用于批量停用 OAuth 资源', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthResource:remove')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_RESOURCE_DISABLE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +@Log(title='OAuth 资源管理', business_type=BusinessType.DELETE) +async def delete_system_oauth_resources( + request: Request, + resource_ids: Annotated[str, Path(min_length=1, max_length=6500)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + actor = _actor(current_user) + try: + batch_ids = OidcUtil.split_batch(resource_ids, 'resource_ids', max_size=_MAX_BATCH_SIZE) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + await OAuthResourceManagementService.disable_resources( + query_db, + batch_ids, + actor, + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(msg='OAuth 资源已停用') + + +@oauth_resource_controller.put( + '/changeStatus', + summary='启停 OAuth 资源接口', + description='用于启用或停用 OAuth 资源', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthResource:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_RESOURCE_STATUS, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 资源管理', business_type=BusinessType.UPDATE) +async def change_system_oauth_resource_status( + request: Request, + payload: ResourceStatusModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthResourceManagementService.change_resource_status( + query_db, payload, _actor(current_user), after_commit=OidcRuntimeService.cors_snapshot_callback(request.app) + ) + + return ResponseUtil.success(data=result) + + +@oauth_scope_controller.get( + '/list', + summary='获取 OAuth 作用域分页列表接口', + description='用于获取 OAuth 作用域分页列表', + response_model=PageResponseModel[ScopeModel], + dependencies=[UserInterfaceAuthDependency('system:oauthScope:list')], +) +async def get_system_oauth_scope_list( + scope_query: Annotated[ScopePageQueryModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + rows = await OAuthResourceManagementService.list_scopes(query_db, scope_query) + total = await OAuthResourceManagementService.count_scopes(query_db, scope_query) + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_scope_controller.get( + '/{scope_code}', + summary='查询 OAuth 作用域详情接口', + description='用于查询 OAuth 作用域详情', + response_model=DataResponseModel[ScopeModel], + dependencies=[UserInterfaceAuthDependency('system:oauthScope:list')], +) +async def query_system_oauth_scope( + scope_code: Annotated[str, Path(min_length=1, max_length=100)], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + return ResponseUtil.success(data=await OAuthResourceManagementService.detail_scope(query_db, scope_code)) + + +@oauth_scope_controller.post( + '', + summary='新增 OAuth 作用域接口', + description='用于新增 OAuth 作用域', + response_model=DataResponseModel[ScopeModel], + dependencies=[UserInterfaceAuthDependency('system:oauthScope:add')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_SCOPE_CREATE, preset=ApiRateLimitPreset.USER_COMMON_MUTATION) +@Log(title='OAuth 权限范围管理', business_type=BusinessType.INSERT) +async def add_system_oauth_scope( + request: Request, + payload: ScopeModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + return ResponseUtil.success( + data=await OAuthResourceManagementService.create_scope( + query_db, payload, _actor(current_user), after_commit=OidcRuntimeService.cors_snapshot_callback(request.app) + ) + ) + + +@oauth_scope_controller.put( + '', + summary='编辑 OAuth 作用域接口', + description='用于编辑 OAuth 作用域', + response_model=DataResponseModel[ScopeModel], + dependencies=[UserInterfaceAuthDependency('system:oauthScope:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_SCOPE_UPDATE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 权限范围管理', business_type=BusinessType.UPDATE) +async def edit_system_oauth_scope( + request: Request, + payload: ScopeModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + return ResponseUtil.success( + data=await OAuthResourceManagementService.update_scope( + query_db, payload, _actor(current_user), after_commit=OidcRuntimeService.cors_snapshot_callback(request.app) + ) + ) + + +@oauth_scope_controller.delete( + '/{scope_codes}', + summary='批量停用 OAuth 作用域接口', + description='用于批量停用 OAuth 作用域', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthScope:remove')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_SCOPE_DISABLE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +@Log(title='OAuth 权限范围管理', business_type=BusinessType.DELETE) +async def delete_system_oauth_scopes( + request: Request, + scope_codes: Annotated[str, Path(min_length=1, max_length=10000)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + actor = _actor(current_user) + try: + batch_ids = OidcUtil.split_batch(scope_codes, 'scope_codes', max_size=_MAX_BATCH_SIZE) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + await OAuthResourceManagementService.disable_scopes( + query_db, + batch_ids, + actor, + after_commit=OidcRuntimeService.cors_snapshot_callback(request.app), + ) + + return ResponseUtil.success(msg='OAuth 权限范围已停用') + + +@oauth_scope_controller.put( + '/changeStatus', + summary='启停 OAuth 作用域接口', + description='用于启用或停用 OAuth 作用域', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthScope:edit')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_SCOPE_STATUS, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +@Log(title='OAuth 权限范围管理', business_type=BusinessType.UPDATE) +async def change_system_oauth_scope_status( + request: Request, + payload: ScopeStatusModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OAuthResourceManagementService.change_scope_status( + query_db, payload, _actor(current_user), after_commit=OidcRuntimeService.cors_snapshot_callback(request.app) + ) + + return ResponseUtil.success(data=result) diff --git a/ruoyi-fastapi-backend/module_identity/controller/oauth_session_controller.py b/ruoyi-fastapi-backend/module_identity/controller/oauth_session_controller.py new file mode 100644 index 000000000..0bbcde415 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/oauth_session_controller.py @@ -0,0 +1,235 @@ +from typing import Annotated + +from fastapi import Path, Query, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + +from common.annotation.log_annotation import Log +from common.annotation.rate_limit_annotation import ApiRateLimit, ApiRateLimitPreset +from common.aspect.db_session import DBSessionDependency +from common.aspect.interface_auth import UserInterfaceAuthDependency +from common.aspect.pre_auth import CurrentUserDependency, PreAuthDependency +from common.constant import ApiNamespace +from common.enums import BusinessType +from common.router import APIRouterPro +from common.vo import DataResponseModel, PageResponseModel, ResponseBaseModel +from exceptions.exception import ServiceException +from module_admin.entity.vo.user_vo import CurrentUserModel +from module_identity.entity.vo.oauth_session_vo import ( + AccessPolicyModel, + AccessPolicyPageQueryModel, + GrantAccessModel, + GrantModel, + GrantPageQueryModel, + SessionPageQueryModel, + SessionRevokeModel, + SsoSessionModel, +) +from module_identity.service.oauth_session_management_service import OAuthSessionManagementService +from utils.oidc_util import OidcUtil +from utils.response_util import ResponseUtil + +oauth_session_controller = APIRouterPro( + prefix='/system/oauth/session', order_num=22, tags=['系统管理-OAuth Session'], dependencies=[PreAuthDependency()] +) +oauth_grant_controller = APIRouterPro( + prefix='/system/oauth/grant', order_num=23, tags=['系统管理-OAuth Grant'], dependencies=[PreAuthDependency()] +) +_MAX_BATCH_SIZE = 100 + + +def _actor(user: CurrentUserModel) -> str: + """ + 提取管理操作者标识 + + :param user: 当前登录用户 + :return: 安全截断后的用户名 + :raises ServiceException: 当前用户不可用 + """ + + try: + return OidcUtil.actor_name( + getattr(getattr(user, 'user', None), 'user_name', None), error_message='当前操作者不可用' + ) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + + +@oauth_session_controller.get( + '/list', + summary='获取 OAuth 会话分页列表接口', + description='用于获取 OAuth 会话分页列表', + response_model=PageResponseModel[SsoSessionModel], + dependencies=[UserInterfaceAuthDependency('system:oauthSession:list')], +) +async def list_oauth_sessions( + query: Annotated[SessionPageQueryModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + rows, total = await OAuthSessionManagementService.list_sessions(query_db, query) + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_session_controller.delete( + '/user/{user_id}', + summary='撤销用户 OAuth 会话接口', + description='用于撤销指定用户的 OAuth 会话', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthSession:revoke')], +) +@ApiRateLimit( + namespace=ApiNamespace.SYSTEM_OAUTH_SESSION_USER_REVOKE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION +) +@Log(title='OAuth 会话管理', business_type=BusinessType.DELETE) +async def revoke_user_oauth_sessions( + request: Request, + user_id: Annotated[int, Path(gt=0)], + payload: SessionRevokeModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + count = await OAuthSessionManagementService.revoke_user( + query_db, request.app.state.redis, user_id, _actor(current_user), payload.reason + ) + + return ResponseUtil.success(msg='单点登录会话已撤销', data={'count': count}) + + +@oauth_session_controller.get( + '/{sid}', + summary='查询 OAuth 会话详情接口', + description='用于查询 OAuth 会话详情', + response_model=DataResponseModel[SsoSessionModel], + dependencies=[UserInterfaceAuthDependency('system:oauthSession:list')], +) +async def get_oauth_session( + sid: Annotated[str, Path(min_length=1, max_length=64)], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + data = await OAuthSessionManagementService.get_session(query_db, sid) + if data is None: + raise ServiceException(message='会话不存在') + return ResponseUtil.success(data=data) + + +@oauth_session_controller.delete( + '/{sids}', + summary='批量撤销 OAuth 会话接口', + description='用于批量撤销 OAuth 会话', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthSession:revoke')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_SESSION_REVOKE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +@Log(title='OAuth 会话管理', business_type=BusinessType.DELETE) +async def revoke_oauth_sessions( + request: Request, + sids: Annotated[str, Path(min_length=1, max_length=6500)], + payload: SessionRevokeModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + try: + batch_ids = OidcUtil.split_batch(sids, 'sids', max_size=_MAX_BATCH_SIZE, validate_all_first=True) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + count = await OAuthSessionManagementService.revoke_sessions( + query_db, request.app.state.redis, batch_ids, _actor(current_user), payload.reason + ) + + return ResponseUtil.success(msg='单点登录会话已撤销', data={'count': count}) + + +@oauth_grant_controller.get( + '/list', + summary='获取 OAuth 授权分页列表接口', + description='用于获取 OAuth 授权分页列表', + response_model=PageResponseModel[GrantModel], + dependencies=[UserInterfaceAuthDependency('system:oauthGrant:list')], +) +async def list_oauth_grants( + query: Annotated[GrantPageQueryModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + rows, total = await OAuthSessionManagementService.list_grants(query_db, query) + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_grant_controller.get( + '/access/list', + summary='获取用户应用访问策略分页列表接口', + description='用于查询独立访问策略,包含尚未授权过的用户与应用', + response_model=PageResponseModel[AccessPolicyModel], + dependencies=[UserInterfaceAuthDependency('system:oauthGrant:list')], +) +async def list_oauth_access_policies( + query: Annotated[AccessPolicyPageQueryModel, Query()], query_db: Annotated[AsyncSession, DBSessionDependency()] +) -> Response: + rows, total = await OAuthSessionManagementService.list_access_policies(query_db, query) + + return ResponseUtil.success(rows=rows, dict_content={'total': total}) + + +@oauth_grant_controller.get( + '/{grant_id}', + summary='查询 OAuth 授权详情接口', + description='用于查询 OAuth 授权详情', + response_model=DataResponseModel[GrantModel], + dependencies=[UserInterfaceAuthDependency('system:oauthGrant:list')], +) +async def get_oauth_grant( + grant_id: Annotated[str, Path(min_length=1, max_length=64)], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + data = await OAuthSessionManagementService.get_grant(query_db, grant_id) + if data is None: + raise ServiceException(message='授权记录不存在') + return ResponseUtil.success(data=data) + + +@oauth_grant_controller.delete( + '/{grant_ids}', + summary='批量撤销 OAuth 授权接口', + description='撤销选中记录所属用户对应用的全部现有授权,重新同意后可再次访问', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthGrant:revoke')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_GRANT_REVOKE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +@Log(title='OAuth 授权管理', business_type=BusinessType.DELETE) +async def revoke_oauth_grants( + request: Request, + grant_ids: Annotated[str, Path(min_length=1, max_length=6500)], + payload: SessionRevokeModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + try: + batch_ids = OidcUtil.split_batch(grant_ids, 'grant_ids', max_size=_MAX_BATCH_SIZE, validate_all_first=True) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + count = await OAuthSessionManagementService.revoke_grants(query_db, batch_ids, _actor(current_user), payload.reason) + + return ResponseUtil.success(msg='OAuth 授权已撤销', data={'count': count}) + + +@oauth_grant_controller.put( + '/user/{user_id}/client/{client_id}/access', + summary='更新用户应用访问策略接口', + description='用于禁止用户访问应用或解除禁止,解除后仍需重新授权', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthGrant:revoke')], +) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_GRANT_ACCESS, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +@Log(title='OAuth 授权管理', business_type=BusinessType.UPDATE) +async def set_oauth_client_access( + request: Request, + user_id: Annotated[int, Path(gt=0)], + client_id: Annotated[str, Path(min_length=1, max_length=128)], + payload: GrantAccessModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + count = await OAuthSessionManagementService.set_access( + query_db, user_id, client_id, payload.blocked, _actor(current_user), payload.reason + ) + return ResponseUtil.success( + msg='已禁止用户访问该应用' if payload.blocked else '已解除禁止,请重新授权', + data={'count': count, 'accessStatus': 'blocked' if payload.blocked else 'allowed'}, + ) diff --git a/ruoyi-fastapi-backend/module_identity/controller/oidc_key_controller.py b/ruoyi-fastapi-backend/module_identity/controller/oidc_key_controller.py new file mode 100644 index 000000000..ab47b47cb --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/oidc_key_controller.py @@ -0,0 +1,144 @@ +from typing import Annotated, Literal + +from fastapi import Path, Query, Request, Response +from sqlalchemy.ext.asyncio import AsyncSession + +from common.annotation.log_annotation import Log +from common.annotation.rate_limit_annotation import ApiRateLimit, ApiRateLimitPreset +from common.aspect.db_session import DBSessionDependency +from common.aspect.interface_auth import UserInterfaceAuthDependency +from common.aspect.pre_auth import CurrentUserDependency, PreAuthDependency +from common.constant import ApiNamespace +from common.enums import BusinessType +from common.router import APIRouterPro +from common.vo import DataResponseModel, PageResponseModel, ResponseBaseModel +from exceptions.exception import ServiceException +from module_admin.entity.vo.user_vo import CurrentUserModel +from module_identity.entity.vo.oidc_key_vo import OidcKeyRotateModel, OidcKeyViewModel +from module_identity.service.key_service import OidcKeyManagementService +from module_identity.service.runtime_service import OidcRuntimeService +from utils.oidc_util import OidcUtil +from utils.response_util import ResponseUtil + +oidc_key_controller = APIRouterPro( + prefix='/system/oauth/key', order_num=24, tags=['系统管理-OIDC 签名密钥'], dependencies=[PreAuthDependency()] +) + + +def _actor(current_user: CurrentUserModel) -> str: + """ + 提取管理操作者标识 + + :param current_user: 当前登录用户 + :return: 安全截断后的用户名 + :raises ServiceException: 当前用户不可用 + """ + + try: + return OidcUtil.actor_name( + getattr(getattr(current_user, 'user', None), 'user_name', None), error_message='当前操作者不可用' + ) + except ValueError as exc: + raise ServiceException(message=str(exc)) from exc + + +@oidc_key_controller.get( + '/list', + summary='获取 OIDC 签名密钥分页列表接口', + description='用于获取 OIDC 签名密钥分页列表', + response_model=PageResponseModel[OidcKeyViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthKey:list')], +) +async def list_oidc_keys( + query_db: Annotated[AsyncSession, DBSessionDependency()], + status: Annotated[Literal['pending', 'active', 'retiring', 'retired', 'compromised'] | None, Query()] = None, + page_num: Annotated[int, Query(ge=1)] = 1, + page_size: Annotated[int, Query(ge=1, le=200)] = 10, +) -> Response: + rows, total = await OidcKeyManagementService.list_page(query_db, status, page_num, page_size) + readiness = await OidcRuntimeService.inspect_readiness(query_db) + + return ResponseUtil.success(rows=rows, dict_content={'total': total, **readiness.as_dict()}) + + +@oidc_key_controller.post( + '/rotate', + summary='轮换 OIDC 签名密钥接口', + description='用于轮换 OIDC 签名密钥', + response_model=DataResponseModel[OidcKeyViewModel], + dependencies=[UserInterfaceAuthDependency('system:oauthKey:rotate')], +) +@Log(title='OIDC 签名密钥', business_type=BusinessType.INSERT) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_KEY_ROTATE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +async def rotate_oidc_key( + request: Request, + payload: OidcKeyRotateModel, + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OidcKeyManagementService.rotate(query_db, payload, _actor(current_user)) + + return ResponseUtil.success(msg='OIDC 签名密钥创建成功', data=result) + + +@oidc_key_controller.put( + '/{kid}/activate', + summary='激活 OIDC 签名密钥接口', + description='用于激活 OIDC 签名密钥', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthKey:activate')], +) +@Log(title='OIDC 签名密钥', business_type=BusinessType.UPDATE) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_KEY_ACTIVATE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +async def activate_oidc_key( + request: Request, + kid: Annotated[str, Path(min_length=1, max_length=100)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OidcKeyManagementService.activate( + query_db, kid, _actor(current_user), getattr(request.app.state, 'redis', None) + ) + await OidcRuntimeService.refresh_readiness(request.app, query_db) + + return ResponseUtil.success(msg='OIDC 签名密钥已激活', data={'changed': result}) + + +@oidc_key_controller.put( + '/{kid}/retire', + summary='退役 OIDC 签名密钥接口', + description='用于退役 OIDC 签名密钥', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthKey:retire')], +) +@Log(title='OIDC 签名密钥', business_type=BusinessType.UPDATE) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_KEY_RETIRE, preset=ApiRateLimitPreset.USER_SECURITY_MUTATION) +async def retire_oidc_key( + request: Request, + kid: Annotated[str, Path(min_length=1, max_length=100)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OidcKeyManagementService.retire(query_db, kid, _actor(current_user)) + + return ResponseUtil.success(msg='OIDC 签名密钥已退役', data={'changed': result}) + + +@oidc_key_controller.delete( + '/{kid}', + summary='删除 OIDC 签名密钥接口', + description='用于删除 OIDC 签名密钥', + response_model=ResponseBaseModel, + dependencies=[UserInterfaceAuthDependency('system:oauthKey:retire')], +) +@Log(title='OIDC 签名密钥', business_type=BusinessType.DELETE) +@ApiRateLimit(namespace=ApiNamespace.SYSTEM_OAUTH_KEY_DELETE, preset=ApiRateLimitPreset.USER_DESTRUCTIVE_MUTATION) +async def delete_oidc_key( + request: Request, + kid: Annotated[str, Path(min_length=1, max_length=100)], + query_db: Annotated[AsyncSession, DBSessionDependency()], + current_user: Annotated[CurrentUserModel, CurrentUserDependency()], +) -> Response: + result = await OidcKeyManagementService.delete(query_db, kid, _actor(current_user)) + + return ResponseUtil.success(msg='OIDC 签名密钥已删除', data={'changed': result}) diff --git a/ruoyi-fastapi-backend/module_identity/controller/token_controller.py b/ruoyi-fastapi-backend/module_identity/controller/token_controller.py new file mode 100644 index 000000000..91b663d8a --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/controller/token_controller.py @@ -0,0 +1,410 @@ +from typing import Annotated, cast + +from fastapi import Depends, HTTPException, Request +from fastapi.responses import JSONResponse, Response +from jwt.exceptions import PyJWTError +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.aspect.db_session import DBSessionDependency +from common.constant import OidcAuditEvent +from common.router import APIRouterPro +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.dependencies import ( + AccessTokenContext, + OidcAccessTokenDependency, + load_access_verification_key, + read_form, + require_oidc_protocol_ready, +) +from module_identity.redis_keys import OidcRedisKey +from module_identity.security.jwt_profile import JwtProfileError +from module_identity.service.audit_service import AuditService +from module_identity.service.infrastructure_service import OidcRateLimiter, RateLimitExceeded, RateLimitUnavailable +from module_identity.service.token_protocol_service import ( + IntrospectionService, + RevocationError, + RevocationService, + UserInfoService, +) +from module_identity.service.token_service import RefreshTokenReuseDetected, TokenService +from utils.oidc_util import OidcUtil + +token_controller = APIRouterPro(tags=['认证中心协议'], order_num=2, dependencies=[Depends(require_oidc_protocol_ready)]) +_NO_STORE = {'Cache-Control': 'no-store', 'Pragma': 'no-cache'} +_FORM_LIMITS = { + 'grant_type': 50, + 'client_id': 64, + 'client_secret': 512, + 'code': 4096, + 'redirect_uri': 1000, + 'code_verifier': 128, + 'refresh_token': 4096, + 'scope': 2000, + 'resource': 1000, + 'token': 8192, + 'token_type_hint': 32, +} +_RATE_LIMIT = 60 +_RATE_WINDOW_SECONDS = 60 +_JWT_DOT_COUNT = 2 +_NOT_FOUND = 404 + + +def _invalid_userinfo_response() -> JSONResponse: + """ + 构造不泄露原因的 UserInfo 401 响应 + + :return: 不泄露校验细节的 UserInfo 401 响应 + """ + # 不向客户端泄露 UserInfo 校验失败原因 + return JSONResponse( + content={'error': 'invalid_token'}, + status_code=401, + headers={**_NO_STORE, 'WWW-Authenticate': 'Bearer error="invalid_token"'}, + ) + + +def _request_redis(request: Request) -> Redis: + """ + 读取应用 Redis 客户端 + + :param request: 当前 HTTP 请求 + :return: Redis 客户端 + :raises OAuthProtocolException: Redis 客户端不可用时抛出 + """ + + redis = getattr(request.app.state, 'redis', None) + if redis is None: + raise OAuthProtocolException('server_error', 'Token endpoint is unavailable', 503) + return cast('Redis', redis) + + +async def _pre_auth_rate_limit(request: Request, redis: Redis, form: dict[str, str], endpoint: str) -> None: + """ + 在 Client Secret 校验前按可见 Client ID 限流 + + :param request: 当前 HTTP 请求 + :param redis: Redis 客户端 + :param form: 已解析的协议表单 + :param endpoint: 当前协议端点标识 + :return: 无返回值 + :raises RateLimitExceeded: 请求超过限流阈值时抛出 + :raises RateLimitUnavailable: 限流服务不可用时抛出 + """ + # 在 Client Secret 校验前按可见 Client ID 限流 + client_id = form.get('client_id') + authorization = request.headers.get('authorization') + if authorization: + try: + client_id, _ = OidcUtil.parse_basic_credentials(authorization) + except ValueError: + client_id = None + if not isinstance(client_id, str) or not client_id: + client_id = 'anonymous' + key_builders = { + 'token': OidcRedisKey.token_client_rate_limit, + 'revoke': OidcRedisKey.revoke_client_rate_limit, + 'introspect': OidcRedisKey.introspect_client_rate_limit, + } + try: + key = key_builders[endpoint](client_id) + except (KeyError, TypeError, ValueError): + key = key_builders[endpoint]('anonymous') + await OidcRateLimiter.enforce(redis, key, limit=_RATE_LIMIT, window_seconds=_RATE_WINDOW_SECONDS) + + +def _validate_protocol_form(form: dict[str, str], allowed: set[str]) -> None: + """ + 校验协议表单字段白名单和长度 + + :param form: 已解析的协议表单 + :param allowed: 当前端点允许的字段集合 + :return: 无返回值 + :raises OAuthProtocolException: 包含未知字段或超长字段时抛出 + """ + # 拒绝未知或超长的协议表单字段 + if any(key not in allowed for key in form) or any( + key in _FORM_LIMITS and len(value) > _FORM_LIMITS[key] for key, value in form.items() + ): + raise OAuthProtocolException('invalid_request', 'Invalid request', 400) + + +def _rate_limit_error(error: RateLimitExceeded | RateLimitUnavailable) -> JSONResponse: + """ + 构造协议限流错误响应 + + :param error: 限流异常 + :return: 裸标准限流错误响应 + """ + # 限流响应不回显敏感字段 + if isinstance(error, RateLimitExceeded): + response = JSONResponse(content={'error': 'temporarily_unavailable'}, status_code=429, headers=_NO_STORE) + response.headers['Retry-After'] = str(error.retry_after) + return response + return JSONResponse(content={'error': 'temporarily_unavailable'}, status_code=503, headers=_NO_STORE) + + +async def _record_independent_audit(db: AsyncSession, event_type: str, **fields: object) -> None: + """ + 尝试记录协议失败审计,审计异常不改变标准错误响应 + + :param db: 异步数据库会话 + :param event_type: 审计事件类型 + :param fields: 审计字段 + :return: 无返回值 + """ + + try: + await AuditService.record_independent(db, event_type, 'failure', **fields) + except Exception: + return + + +async def _verification_key_or_none(token: str, query_db: AsyncSession) -> object: + """ + 为 JWT 撤销和内省加载本地公钥 + + :param token: 待读取的访问令牌 + :param query_db: 异步数据库会话 + :return: 本地公钥或 None + """ + # Opaque Token 保持幂等语义,JWT 校验失败按无公钥处理 + if not isinstance(token, str) or token.count('.') != _JWT_DOT_COUNT: + return None + try: + return await load_access_verification_key(token, query_db) + except (JwtProfileError, PyJWTError, TypeError, ValueError): + return None + + +async def _verification_key_loader(query_db: AsyncSession, token: str) -> object: + """ + 提供给 Service 事务入口的本地公钥加载回调 + + :param query_db: 异步数据库会话 + :param token: 待读取的访问令牌 + :return: 本地公钥或 None + """ + + return await _verification_key_or_none(token, query_db) + + +def _oauth_error(error: OAuthProtocolException) -> JSONResponse: + """ + 生成裸标准 OAuth 错误响应 + + :param error: OAuth 协议异常 + :return: 裸标准 OAuth 错误响应 + """ + + content = {'error': error.error} + if error.error_description: + content['error_description'] = error.error_description + headers = dict(_NO_STORE) + if error.error == 'invalid_client': + headers['WWW-Authenticate'] = 'Basic realm="oauth2/token"' + headers.update(error.headers) + + return JSONResponse(content=content, status_code=error.status_code, headers=headers) + + +def _http_error(error: HTTPException) -> JSONResponse: + """ + 将输入边界 HTTP 错误转换为裸 JSON 响应 + + :param error: 输入边界 HTTP 异常 + :return: 裸 JSON 错误响应 + """ + + code = 'not_found' if error.status_code == _NOT_FOUND else 'invalid_request' + + return JSONResponse(content={'error': code}, status_code=error.status_code, headers=_NO_STORE) + + +@token_controller.post( + '/oauth2/token', + summary='获取 OAuth Token 接口', + description='用于签发 OAuth Token', + include_in_schema=False, +) +async def token( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + try: + form = await read_form(request) + _validate_protocol_form( + form, + { + 'grant_type', + 'client_id', + 'client_secret', + 'code', + 'redirect_uri', + 'code_verifier', + 'refresh_token', + 'scope', + 'resource', + }, + ) + redis = _request_redis(request) + await _pre_auth_rate_limit(request, redis, form, 'token') + client_secret = form.pop('client_secret', None) + result = await TokenService.issue_token_request( + query_db, + redis, + form, + authorization=request.headers.get('authorization'), + client_id=form.get('client_id'), + client_secret=client_secret, + ) + return JSONResponse(content=result.as_dict(), headers=_NO_STORE) + except RefreshTokenReuseDetected as exc: + return _oauth_error(exc) + except (RateLimitExceeded, RateLimitUnavailable) as exc: + return _rate_limit_error(exc) + except OAuthProtocolException as exc: + await _record_independent_audit( + query_db, + OidcAuditEvent.INVALID_CLIENT if exc.error == 'invalid_client' else OidcAuditEvent.TOKEN_FAILED, + failure_code=exc.error, + ) + return _oauth_error(exc) + except HTTPException as exc: + return _http_error(exc) + except Exception: + await _record_independent_audit(query_db, OidcAuditEvent.TOKEN_FAILED, failure_code='server_error') + return JSONResponse( + content={'error': 'server_error', 'error_description': 'Token endpoint is unavailable'}, + status_code=500, + headers=_NO_STORE, + ) + + +@token_controller.post( + '/oauth2/revoke', + summary='撤销 OAuth Token 接口', + description='用于撤销 OAuth Token', + include_in_schema=False, +) +async def revoke( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + try: + form = await read_form(request) + _validate_protocol_form(form, {'client_id', 'client_secret', 'token', 'token_type_hint'}) + redis = _request_redis(request) + await _pre_auth_rate_limit(request, redis, form, 'revoke') + committed = await RevocationService.revoke_request( + query_db, + redis, + form.get('token', ''), + authorization=request.headers.get('authorization'), + client_id=form.get('client_id'), + client_secret=form.get('client_secret'), + verification_key_loader=_verification_key_loader, + token_type_hint=form.get('token_type_hint'), + ) + if not committed: + return JSONResponse(content={'error': 'temporarily_unavailable'}, status_code=503, headers=_NO_STORE) + return Response(status_code=200, headers=_NO_STORE) + except OAuthProtocolException as exc: + await _record_independent_audit( + query_db, + OidcAuditEvent.INVALID_CLIENT if exc.error == 'invalid_client' else OidcAuditEvent.TOKEN_FAILED, + failure_code=exc.error, + ) + return _oauth_error(exc) + except (RateLimitExceeded, RateLimitUnavailable) as exc: + return _rate_limit_error(exc) + except RevocationError as exc: + await _record_independent_audit( + query_db, + OidcAuditEvent.INVALID_CLIENT if exc.error == 'invalid_client' else OidcAuditEvent.TOKEN_FAILED, + failure_code=exc.error, + ) + headers = dict(_NO_STORE) + if exc.error == 'invalid_client': + headers['WWW-Authenticate'] = 'Basic realm="oauth2/token"' + return JSONResponse( + content={'error': exc.error, 'error_description': exc.description}, + status_code=401 if exc.error == 'invalid_client' else 400, + headers=headers, + ) + except HTTPException as exc: + return _http_error(exc) + except Exception: + await _record_independent_audit(query_db, OidcAuditEvent.TOKEN_FAILED, failure_code='server_error') + return JSONResponse(content={'error': 'server_error'}, status_code=500, headers=_NO_STORE) + + +@token_controller.post( + '/oauth2/introspect', + summary='查询 OAuth Token 状态接口', + description='用于查询 OAuth Token 状态', + include_in_schema=False, +) +async def introspect( + request: Request, + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + try: + form = await read_form(request) + _validate_protocol_form(form, {'client_id', 'client_secret', 'token', 'token_type_hint'}) + redis = _request_redis(request) + await _pre_auth_rate_limit(request, redis, form, 'introspect') + result = await IntrospectionService.introspect_request( + query_db, + redis, + form.get('token', ''), + authorization=request.headers.get('authorization'), + client_id=form.get('client_id'), + client_secret=form.get('client_secret'), + verification_key_loader=_verification_key_loader, + token_type_hint=form.get('token_type_hint'), + ) + return JSONResponse(content=result, headers=_NO_STORE) + except OAuthProtocolException as exc: + await _record_independent_audit( + query_db, + OidcAuditEvent.INVALID_CLIENT if exc.error == 'invalid_client' else OidcAuditEvent.TOKEN_FAILED, + failure_code=exc.error, + ) + return _oauth_error(exc) + except (RateLimitExceeded, RateLimitUnavailable) as exc: + return _rate_limit_error(exc) + except HTTPException as exc: + return _http_error(exc) + except Exception: + await _record_independent_audit(query_db, OidcAuditEvent.TOKEN_FAILED, failure_code='server_error') + return JSONResponse(content={'error': 'server_error'}, status_code=500, headers=_NO_STORE) + + +@token_controller.get( + '/oauth2/userinfo', + summary='获取 UserInfo 接口', + description='用于返回当前访问令牌对应的用户信息', + include_in_schema=False, +) +@token_controller.post( + '/oauth2/userinfo', + summary='获取 UserInfo 接口', + description='用于返回当前访问令牌对应的用户信息', + include_in_schema=False, +) +async def userinfo( + request: Request, + token_context: Annotated[AccessTokenContext, OidcAccessTokenDependency()], + query_db: Annotated[AsyncSession, DBSessionDependency()], +) -> Response: + if not OidcConfig.oidc_enabled: + return JSONResponse(content={'error': 'not_found'}, status_code=404, headers=_NO_STORE) + try: + claims = token_context.claims + output = await UserInfoService.build(query_db, claims, getattr(request.app.state, 'redis', None)) + return JSONResponse(content=output, headers=_NO_STORE) + except Exception: + return _invalid_userinfo_response() diff --git a/ruoyi-fastapi-backend/module_identity/dao/identity_subject_dao.py b/ruoyi-fastapi-backend/module_identity/dao/identity_subject_dao.py new file mode 100644 index 000000000..dc4ec8beb --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/identity_subject_dao.py @@ -0,0 +1,195 @@ +from collections.abc import Iterable, Sequence +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import select, update +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from module_admin.entity.do.user_do import SysUser +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from utils.time_util import TimezoneUtil + + +class IdentitySubjectDao: + """ + Identity Subject 数据库操作层 + """ + + @classmethod + async def get_by_user_id(cls, db: AsyncSession, user_id: int) -> SysIdentitySubject | None: + """ + 按用户编号查询 Identity Subject + + :param db: orm对象 + :param user_id: 用户编号 + :return: Identity Subject,不存在时返回 None + """ + + result = await db.execute(select(SysIdentitySubject).where(SysIdentitySubject.user_id == user_id)) + + return result.scalars().first() + + @classmethod + async def get_by_subject_id(cls, db: AsyncSession, subject_id: str) -> SysIdentitySubject | None: + """ + 按主体标识查询 Identity Subject + + :param db: orm对象 + :param subject_id: 主体标识 + :return: Identity Subject,不存在时返回 None + """ + + result = await db.execute(select(SysIdentitySubject).where(SysIdentitySubject.subject_id == subject_id)) + + return result.scalars().first() + + @classmethod + async def list_for_users_for_update(cls, db: AsyncSession, user_ids: Sequence[int]) -> Sequence[SysIdentitySubject]: + """ + 按用户编号锁定查询 Identity Subject + + :param db: orm对象 + :param user_ids: 用户编号序列 + :return: 按用户编号排序的 Identity Subject 序列 + """ + + result = await db.execute( + select(SysIdentitySubject) + .where(SysIdentitySubject.user_id.in_(user_ids)) + .order_by(SysIdentitySubject.user_id) + .with_for_update() + ) + + return result.scalars().all() + + @classmethod + async def create_for_user( + cls, db: AsyncSession, user_id: int, create_by: str | None = None, subject_id: str | None = None + ) -> SysIdentitySubject: + """ + 新增用户 Identity Subject + + :param db: orm对象 + :param user_id: 用户编号 + :param create_by: 创建人标识 + :param subject_id: 主体标识 + :return: 新建或已存在的 Identity Subject + """ + + existing = await cls.get_by_user_id(db, user_id) + if existing is not None: + return existing + subject = SysIdentitySubject( + user_id=user_id, + subject_id=subject_id or str(uuid4()), + auth_version=1, + create_by=create_by, + create_time=TimezoneUtil.utc_now(), + ) + try: + async with db.begin_nested(): + db.add(subject) + await db.flush() + except IntegrityError: + existing = await cls.get_by_user_id(db, user_id) + if existing is not None: + return existing + raise + return subject + + @classmethod + async def backfill_for_users( + cls, db: AsyncSession, user_ids: Iterable[int], create_by: str = 'migration' + ) -> Sequence[SysIdentitySubject]: + """ + 批量补齐用户 Identity Subject + + :param db: orm对象 + :param user_ids: 用户编号序列 + :param create_by: 创建人标识 + :return: 补齐后的 Identity Subject 序列 + """ + + ids = list(dict.fromkeys(user_ids)) + if not ids: + return [] + result = await db.execute(select(SysIdentitySubject.user_id).where(SysIdentitySubject.user_id.in_(ids))) + existing = set(result.scalars().all()) + rows: list[SysIdentitySubject] = [] + for user_id in ids: + if user_id in existing: + continue + rows.append(await cls.create_for_user(db, user_id, create_by=create_by)) + return rows + + @classmethod + async def increment_auth_version(cls, db: AsyncSession, user_id: int, expected_version: int | None = None) -> bool: + """ + 递增单个用户认证版本 + + :param db: orm对象 + :param user_id: 用户编号 + :param expected_version: 期望认证版本 + :return: 是否更新成功 + """ + + conditions = [SysIdentitySubject.user_id == user_id] + if expected_version is not None: + conditions.append(SysIdentitySubject.auth_version == expected_version) + result = await db.execute( + update(SysIdentitySubject) + .where(*conditions) + .values(auth_version=SysIdentitySubject.auth_version + 1, update_time=TimezoneUtil.utc_now()) + ) + + return bool(result.rowcount) + + @classmethod + async def increment_auth_versions( + cls, db: AsyncSession, user_ids: Sequence[int], update_by: str, now: datetime | None = None + ) -> int: + """ + 批量递增用户认证版本 + + :param db: orm对象 + :param user_ids: 用户编号序列 + :param update_by: 更新人标识 + :param now: 当前时间 + :return: 成功更新的 Identity Subject 数量 + """ + + result = await db.execute( + update(SysIdentitySubject) + .where(SysIdentitySubject.user_id.in_(user_ids)) + .values( + auth_version=SysIdentitySubject.auth_version + 1, + update_by=update_by, + update_time=now or TimezoneUtil.utc_now(), + ) + ) + await db.flush() + + return result.rowcount or 0 + + @classmethod + async def list_missing_user_ids(cls, db: AsyncSession, user_ids: Iterable[int] | None = None) -> list[int]: + """ + 查询缺少 Identity Subject 的用户编号 + + :param db: orm对象 + :param user_ids: 用户编号序列 + :return: 缺少 Identity Subject 的用户编号列表 + """ + + if user_ids is None: + result = await db.execute(select(SysUser.user_id).where(SysUser.del_flag == '0')) + ids = list(result.scalars().all()) + else: + ids = list(dict.fromkeys(user_ids)) + if not ids: + return [] + result = await db.execute(select(SysIdentitySubject.user_id).where(SysIdentitySubject.user_id.in_(ids))) + existing = set(result.scalars().all()) + + return [user_id for user_id in ids if user_id not in existing] diff --git a/ruoyi-fastapi-backend/module_identity/dao/identity_user_dao.py b/ruoyi-fastapi-backend/module_identity/dao/identity_user_dao.py new file mode 100644 index 000000000..d9eadb030 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/identity_user_dao.py @@ -0,0 +1,80 @@ +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from module_admin.entity.do.dept_do import SysDept +from module_admin.entity.do.role_do import SysRole +from module_admin.entity.do.user_do import SysUser, SysUserRole + + +class IdentityUserDao: + """ + Identity User 数据库操作层 + """ + + @classmethod + async def get_active_user(cls, db: AsyncSession, user_id: int) -> SysUser | None: + """ + 按用户编号查询活跃用户 + + :param db: orm对象 + :param user_id: 用户编号 + :return: SysUser,不存在时返回 None + """ + + result = await db.execute( + select(SysUser).where(SysUser.status == '0', SysUser.del_flag == '0', SysUser.user_id == user_id) + ) + + return result.scalars().first() + + @classmethod + async def get_user(cls, db: AsyncSession, user_id: int) -> SysUser | None: + """ + 按用户编号查询用户 + + :param db: orm对象 + :param user_id: 用户编号 + :return: SysUser,不存在时返回 None + """ + + return await db.scalar(select(SysUser).where(SysUser.user_id == user_id)) + + @classmethod + async def list_user_ids_by_role_id(cls, db: AsyncSession, role_id: int) -> list[int]: + """ + 按角色编号查询用户编号 + + :param db: orm对象 + :param role_id: 角色编号 + :return: 用户编号列表 + """ + + result = await db.execute( + select(SysUserRole.user_id).where(SysUserRole.role_id == role_id).order_by(SysUserRole.user_id) + ) + + return list(result.scalars().all()) + + @classmethod + async def get_claim_attributes(cls, db: AsyncSession, user_id: int) -> tuple[list[str], SysDept | None]: + """ + 查询用户声明属性和部门 + + :param db: orm对象 + :param user_id: 用户编号 + :return: 声明名称列表和 SysDept,部门不存在时为 None + """ + + role_result = await db.execute( + select(SysRole.role_key) + .join(SysUserRole, SysUserRole.role_id == SysRole.role_id) + .where(SysUserRole.user_id == user_id, SysRole.status == '0', SysRole.del_flag == '0') + .order_by(SysRole.role_id) + ) + dept_result = await db.execute( + select(SysDept) + .join(SysUser, SysUser.dept_id == SysDept.dept_id) + .where(SysUser.user_id == user_id, SysDept.status == '0', SysDept.del_flag == '0') + ) + + return list(role_result.scalars().all()), dept_result.scalars().first() diff --git a/ruoyi-fastapi-backend/module_identity/dao/oauth_access_policy_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oauth_access_policy_dao.py new file mode 100644 index 000000000..298def5f3 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oauth_access_policy_dao.py @@ -0,0 +1,177 @@ +from sqlalchemy import Select, func, select, tuple_, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_admin.entity.do.user_do import SysUser +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthAccessPolicy +from module_identity.entity.vo.oauth_session_vo import AccessPolicyPageQueryModel +from utils.time_util import TimezoneUtil + + +class OAuthAccessPolicyDao: + """ + 用户应用访问控制数据库操作层 + """ + + @staticmethod + async def lock_client(db: AsyncSession, client_pk: int) -> None: + """ + 锁定 Client,统一授权签发、撤销和访问控制的事务顺序 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :return: None + """ + + await db.execute( + update(SysOAuthClient) + .where(SysOAuthClient.client_pk == client_pk) + .values(policy_version=SysOAuthClient.policy_version) + ) + await db.execute( + select(SysOAuthClient) + .where(SysOAuthClient.client_pk == client_pk) + .with_for_update() + .execution_options(populate_existing=True) + ) + + @staticmethod + async def get( + db: AsyncSession, user_id: int, client_pk: int, *, for_update: bool = False + ) -> SysOAuthAccessPolicy | None: + """ + 查询用户对 Client 的访问策略,未配置时默认允许 + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :param for_update: 是否使用当前读锁定策略 + :return: 访问策略,不存在时返回 None + """ + + query = select(SysOAuthAccessPolicy).where( + SysOAuthAccessPolicy.user_id == user_id, SysOAuthAccessPolicy.client_pk == client_pk + ) + if for_update: + query = query.with_for_update() + result = await db.execute(query.execution_options(populate_existing=True)) + + return result.scalars().first() + + @classmethod + async def is_blocked(cls, db: AsyncSession, user_id: int, client_pk: int, *, for_update: bool = False) -> bool: + """ + 判断用户是否被禁止访问 Client + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :param for_update: 是否使用当前读锁定策略 + :return: 是否禁止访问 + """ + + row = await cls.get(db, user_id, client_pk, for_update=for_update) + return row is not None and row.access_status == 'blocked' + + @classmethod + async def set_status( + cls, db: AsyncSession, user_id: int, client_pk: int, blocked: bool, actor: str, reason: str + ) -> SysOAuthAccessPolicy: + """ + 更新访问策略;调用方须先锁定 Client 并负责事务提交 + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :param blocked: 是否禁止访问 + :param actor: 操作人 + :param reason: 操作原因 + :return: 更新后的访问策略 + """ + + row = await cls.get(db, user_id, client_pk, for_update=True) + if row is None: + row = SysOAuthAccessPolicy(user_id=user_id, client_pk=client_pk) + db.add(row) + row.access_status = 'blocked' if blocked else 'allowed' + row.reason = reason + row.update_by = actor + row.update_time = TimezoneUtil.utc_now() + await db.flush() + + return row + + @staticmethod + async def list_for_grants(db: AsyncSession, targets: list[tuple[int, int]]) -> list[SysOAuthAccessPolicy]: + """ + 批量查询管理列表中用户与 Client 的访问策略 + + :param db: orm对象 + :param targets: 用户编号和 Client 主键列表 + :return: 用户访问策略列表 + """ + + result = await db.execute( + select(SysOAuthAccessPolicy).where( + tuple_(SysOAuthAccessPolicy.user_id, SysOAuthAccessPolicy.client_pk).in_(targets) + ) + ) + return list(result.scalars().all()) + + @staticmethod + def _query(query: AccessPolicyPageQueryModel) -> Select: + """ + 构建独立访问策略查询,不依赖用户已有授权记录 + + :param query: 访问策略分页查询参数 + :return: 包含用户与客户端名称的查询语句 + """ + + statement = ( + select(SysOAuthAccessPolicy, SysUser.user_name, SysOAuthClient.client_id, SysOAuthClient.client_name) + .join(SysUser, SysUser.user_id == SysOAuthAccessPolicy.user_id) + .join(SysOAuthClient, SysOAuthClient.client_pk == SysOAuthAccessPolicy.client_pk) + ) + if query.user_id is not None: + statement = statement.where(SysOAuthAccessPolicy.user_id == query.user_id) + if query.client_id: + statement = statement.where(SysOAuthClient.client_id == query.client_id) + if query.access_status: + statement = statement.where(SysOAuthAccessPolicy.access_status == query.access_status) + return statement + + @classmethod + async def list_page( + cls, db: AsyncSession, query: AccessPolicyPageQueryModel + ) -> list[tuple[SysOAuthAccessPolicy, str, str, str]]: + """ + 分页查询访问策略及其关联名称 + + :param db: orm对象 + :param query: 访问策略分页查询参数 + :return: 访问策略、用户名、客户端标识和名称的列表 + """ + + result = await db.execute( + cls._query(query) + .order_by( + SysOAuthAccessPolicy.update_time.desc(), + SysOAuthAccessPolicy.user_id, + SysOAuthAccessPolicy.client_pk, + ) + .offset((query.page_num - 1) * query.page_size) + .limit(query.page_size) + ) + return list(result.all()) + + @classmethod + async def count(cls, db: AsyncSession, query: AccessPolicyPageQueryModel) -> int: + """ + 统计匹配条件的独立访问策略数量 + + :param db: orm对象 + :param query: 访问策略分页查询参数 + :return: 匹配的记录总数 + """ + + return int(await db.scalar(select(func.count()).select_from(cls._query(query).subquery())) or 0) diff --git a/ruoyi-fastapi-backend/module_identity/dao/oauth_audit_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oauth_audit_dao.py new file mode 100644 index 000000000..8a69a2e3c --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oauth_audit_dao.py @@ -0,0 +1,254 @@ +from collections.abc import Iterable, Sequence +from datetime import datetime + +from sqlalchemy import delete, func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditArchive, SysOAuthAuditLog + + +class OAuthAuditDao: + """ + OAuth Audit 数据库操作层 + """ + + @classmethod + async def append(cls, db: AsyncSession, event: SysOAuthAuditLog) -> SysOAuthAuditLog: + """ + 新增 OAuth Audit 日志并刷新主键 + + :param db: orm对象 + :param event: OAuth Audit 日志对象 + :return: 已写入的 OAuth Audit 日志对象 + """ + + db.add(event) + await db.flush() + + return event + + @classmethod + async def bulk_append(cls, db: AsyncSession, events: Iterable[SysOAuthAuditLog]) -> None: + """ + 批量新增 OAuth Audit 日志 + + :param db: orm对象 + :param events: OAuth Audit 日志对象序列 + :return: None + """ + + rows = list(events) + if rows: + db.add_all(rows) + await db.flush() + + @classmethod + async def list_page( + cls, + db: AsyncSession, + offset: int = 0, + limit: int = 50, + client_id: str | None = None, + user_id: int | None = None, + event_type: str | None = None, + result: str | None = None, + before: datetime | None = None, + ) -> Sequence[SysOAuthAuditLog]: + """ + 分页查询 OAuth Audit 日志 + + :param db: orm对象 + :param offset: 分页偏移量 + :param limit: 分页大小 + :param client_id: Client 公开标识 + :param user_id: 用户编号 + :param event_type: 审计事件类型 + :param result: 审计结果 + :param before: 归档截止时间 + :return: OAuth Audit 日志序列 + """ + + conditions = [] + if client_id: + conditions.append(SysOAuthAuditLog.client_id == client_id) + if user_id is not None: + conditions.append(SysOAuthAuditLog.user_id == user_id) + if event_type: + conditions.append(SysOAuthAuditLog.event_type == event_type) + if result: + conditions.append(SysOAuthAuditLog.result == result) + if before is not None: + conditions.append(SysOAuthAuditLog.create_time < before) + query = select(SysOAuthAuditLog).where(*conditions).order_by(SysOAuthAuditLog.event_id.desc()) + query = query.offset(max(offset, 0)).limit(min(max(limit, 1), 500)) + rows = await db.execute(query) + + return rows.scalars().all() + + @classmethod + def _conditions( + cls, + *, + client_id: str | None = None, + user_id: int | None = None, + event_type: str | None = None, + result: str | None = None, + risk_level: str | None = None, + start_time: datetime | None = None, + end_time: datetime | None = None, + ) -> list: + """ + 构造 OAuth Audit 查询条件 + + :param client_id: Client 公开标识 + :param user_id: 用户编号 + :param event_type: 审计事件类型 + :param result: 审计结果 + :param risk_level: 审计风险级别 + :param start_time: 查询开始时间 + :param end_time: 查询结束时间 + :return: SQLAlchemy 条件列表 + """ + + conditions = [] + if client_id: + conditions.append(SysOAuthAuditLog.client_id == client_id) + if user_id is not None: + conditions.append(SysOAuthAuditLog.user_id == user_id) + if event_type: + conditions.append(SysOAuthAuditLog.event_type == event_type) + if result: + conditions.append(SysOAuthAuditLog.result == result) + if risk_level: + conditions.append(SysOAuthAuditLog.risk_level == risk_level) + if start_time is not None: + conditions.append(SysOAuthAuditLog.create_time >= start_time) + if end_time is not None: + conditions.append(SysOAuthAuditLog.create_time <= end_time) + return conditions + + @classmethod + async def list_admin_page( + cls, + db: AsyncSession, + *, + offset: int, + limit: int, + client_id: str | None = None, + user_id: int | None = None, + event_type: str | None = None, + result: str | None = None, + risk_level: str | None = None, + start_time: datetime | None = None, + end_time: datetime | None = None, + ) -> Sequence[SysOAuthAuditLog]: + """ + 分页查询管理端 OAuth Audit 日志 + + :param db: orm对象 + :param offset: 分页偏移量 + :param limit: 分页大小 + :param client_id: Client 公开标识 + :param user_id: 用户编号 + :param event_type: 审计事件类型 + :param result: 审计结果 + :param risk_level: 审计风险级别 + :param start_time: 查询开始时间 + :param end_time: 查询结束时间 + :return: 管理端 OAuth Audit 日志序列 + """ + + query = ( + select(SysOAuthAuditLog) + .where( + *cls._conditions( + client_id=client_id, + user_id=user_id, + event_type=event_type, + result=result, + risk_level=risk_level, + start_time=start_time, + end_time=end_time, + ) + ) + .order_by(SysOAuthAuditLog.event_id.desc()) + .offset(max(offset, 0)) + .limit(min(max(limit, 1), 5000)) + ) + rows = await db.execute(query) + + return rows.scalars().all() + + @classmethod + async def count_admin( + cls, + db: AsyncSession, + *, + client_id: str | None = None, + user_id: int | None = None, + event_type: str | None = None, + result: str | None = None, + risk_level: str | None = None, + start_time: datetime | None = None, + end_time: datetime | None = None, + ) -> int: + """ + 统计管理端 OAuth Audit 日志数量 + + :param db: orm对象 + :param client_id: Client 公开标识 + :param user_id: 用户编号 + :param event_type: 审计事件类型 + :param result: 审计结果 + :param risk_level: 审计风险级别 + :param start_time: 查询开始时间 + :param end_time: 查询结束时间 + :return: 符合条件的 OAuth Audit 日志数量 + """ + + result = await db.execute( + select(func.count()) + .select_from(SysOAuthAuditLog) + .where( + *cls._conditions( + client_id=client_id, + user_id=user_id, + event_type=event_type, + result=result, + risk_level=risk_level, + start_time=start_time, + end_time=end_time, + ) + ) + ) + + return int(result.scalar_one()) + + @classmethod + async def archive_before(cls, db: AsyncSession, before: datetime, limit: int = 1000) -> int: + """ + 归档指定时间前的 OAuth Audit 日志 + + :param db: orm对象 + :param before: 归档截止时间 + :param limit: 单次归档数量上限 + :return: 已归档的 OAuth Audit 日志数量 + """ + + result = await db.execute( + select(SysOAuthAuditLog) + .where(SysOAuthAuditLog.create_time < before) + .order_by(SysOAuthAuditLog.event_id) + .limit(min(max(limit, 1), 5000)) + .with_for_update() + ) + rows = result.scalars().all() + if not rows: + return 0 + fields = tuple(column.name for column in SysOAuthAuditLog.__table__.columns) + db.add_all([SysOAuthAuditArchive(**{field: getattr(row, field) for field in fields}) for row in rows]) + await db.flush() + event_ids = [row.event_id for row in rows] + await db.execute(delete(SysOAuthAuditLog).where(SysOAuthAuditLog.event_id.in_(event_ids))) + + return len(event_ids) diff --git a/ruoyi-fastapi-backend/module_identity/dao/oauth_client_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oauth_client_dao.py new file mode 100644 index 000000000..ce6304408 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oauth_client_dao.py @@ -0,0 +1,719 @@ +from collections.abc import Mapping, Sequence +from datetime import datetime + +from sqlalchemy import delete, func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.oauth_client_do import SysOAuthClient, SysOAuthClientSecret, SysOAuthClientUri +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.oauth_client_vo import ClientPageQueryModel +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + + +class OAuthClientDao: + """ + OAuth Client 数据库操作层 + """ + + @classmethod + async def add_client(cls, db: AsyncSession, client: SysOAuthClient) -> None: + """ + 新增 OAuth Client 并刷新主键 + + :param db: orm对象 + :param client: OAuth Client 对象 + :return: None + """ + + db.add(client) + await db.flush() + + @classmethod + async def add_uri(cls, db: AsyncSession, uri: SysOAuthClientUri) -> None: + """ + 新增 OAuth Client 回调地址并刷新主键 + + :param db: orm对象 + :param uri: 回调地址 + :return: None + """ + + db.add(uri) + await db.flush() + + @classmethod + async def persist_client_policy_change(cls, db: AsyncSession, client: SysOAuthClient) -> None: + """ + 刷新 OAuth Client 策略变更 + + :param db: orm对象 + :param client: OAuth Client 对象 + :return: None + """ + + await db.flush() + + @classmethod + async def get_by_client_id( + cls, db: AsyncSession, client_id: str, active_only: bool = True, for_update: bool = False + ) -> SysOAuthClient | None: + """ + 按公开标识查询 OAuth Client + + :param db: orm对象 + :param client_id: Client 公开标识 + :param active_only: 是否仅查询启用记录 + :param for_update: 是否锁定查询结果 + :return: OAuth Client,不存在时返回 None + """ + + conditions = [SysOAuthClient.client_id == client_id] + if active_only: + conditions.append(SysOAuthClient.status == '0') + query = select(SysOAuthClient).where(*conditions) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def get_by_pk( + cls, db: AsyncSession, client_pk: int, active_only: bool = True, for_update: bool = False + ) -> SysOAuthClient | None: + """ + 按内部主键查询 OAuth Client + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param active_only: 是否仅查询启用记录 + :param for_update: 是否锁定查询结果 + :return: OAuth Client,不存在时返回 None + """ + + conditions = [SysOAuthClient.client_pk == client_pk] + if active_only: + conditions.append(SysOAuthClient.status == '0') + query = select(SysOAuthClient).where(*conditions) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def id_map(cls, db: AsyncSession, client_pks: Sequence[int]) -> Mapping[int, str]: + """ + 批量查询 Client 主键与公开标识映射 + + :param db: orm对象 + :param client_pks: Client 内部主键序列 + :return: Client 内部主键到公开标识的映射 + """ + + if not client_pks: + return {} + result = await db.execute( + select(SysOAuthClient.client_pk, SysOAuthClient.client_id).where(SysOAuthClient.client_pk.in_(client_pks)) + ) + + return {int(client_pk): client_id for client_pk, client_id in result.all()} + + @classmethod + async def id_for_pk(cls, db: AsyncSession, client_pk: int) -> str | None: + """ + 按内部主键查询 Client 公开标识 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :return: Client 公开标识,不存在时返回 None + """ + + result = await db.execute(select(SysOAuthClient.client_id).where(SysOAuthClient.client_pk == client_pk)) + + return result.scalar_one_or_none() + + @classmethod + async def has_disabled_bound_resource(cls, db: AsyncSession, client_id: str, audience: str) -> bool: + """ + 查询 Client 是否绑定停用 Resource + + :param db: orm对象 + :param client_id: Client 公开标识 + :param audience: Resource 受众 + :return: 是否绑定停用 Resource + """ + + result = await db.execute( + select(SysOAuthResource.resource_pk) + .join(SysOAuthClientResource, SysOAuthClientResource.resource_pk == SysOAuthResource.resource_pk) + .join(SysOAuthClient, SysOAuthClient.client_pk == SysOAuthClientResource.client_pk) + .where( + SysOAuthClient.client_id == client_id, + SysOAuthResource.audience == audience, + SysOAuthResource.status != '0', + ) + ) + + return result.scalar_one_or_none() is not None + + @classmethod + async def list_secrets( + cls, db: AsyncSession, client_pk: int, active_only: bool = True, for_update: bool = False + ) -> Sequence[SysOAuthClientSecret]: + """ + 查询 OAuth Client 密钥列表 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param active_only: 是否仅查询启用记录 + :param for_update: 是否锁定查询结果 + :return: OAuth Client Secret 序列 + """ + + conditions = [SysOAuthClientSecret.client_pk == client_pk] + if active_only: + now = TimezoneUtil.utc_now() + conditions.extend( + [ + SysOAuthClientSecret.status.in_(['active', 'retiring']), + SysOAuthClientSecret.not_before <= now, + (SysOAuthClientSecret.expires_at.is_(None) | (SysOAuthClientSecret.expires_at > now)), + ] + ) + query = select(SysOAuthClientSecret).where(*conditions).order_by(SysOAuthClientSecret.create_time) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().all() + + @classmethod + async def get_secret_for_update(cls, db: AsyncSession, secret_id: str) -> SysOAuthClientSecret | None: + """ + 按密钥标识锁定查询 Client Secret + + :param db: orm对象 + :param secret_id: Client Secret 标识 + :return: Client Secret,不存在时返回 None + """ + + result = await db.execute( + select(SysOAuthClientSecret).where(SysOAuthClientSecret.secret_id == secret_id).with_for_update() + ) + + return result.scalars().first() + + @classmethod + async def mark_secret_used(cls, db: AsyncSession, secret_id: str) -> bool: + """ + 标记 Client Secret 已使用 + + :param db: orm对象 + :param secret_id: Client Secret 标识 + :return: 是否更新成功 + """ + + result = await db.execute( + update(SysOAuthClientSecret) + .where(SysOAuthClientSecret.secret_id == secret_id) + .values(last_used_at=TimezoneUtil.utc_now()) + ) + + return bool(result.rowcount) + + @classmethod + async def list_uris( + cls, db: AsyncSession, client_pk: int, uri_type: str | None = None, active_only: bool = True + ) -> Sequence[SysOAuthClientUri]: + """ + 查询 OAuth Client 回调地址列表 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param uri_type: 回调地址类型 + :param active_only: 是否仅查询启用记录 + :return: OAuth Client 回调地址序列 + """ + + conditions = [SysOAuthClientUri.client_pk == client_pk] + if uri_type: + conditions.append(SysOAuthClientUri.uri_type == uri_type) + if active_only: + conditions.append(SysOAuthClientUri.status == '0') + result = await db.execute(select(SysOAuthClientUri).where(*conditions).order_by(SysOAuthClientUri.uri_id)) + + return result.scalars().all() + + @classmethod + async def find_exact_uri( + cls, db: AsyncSession, client_pk: int, uri_type: str, uri: str + ) -> SysOAuthClientUri | None: + """ + 按完整地址查询 OAuth Client 回调地址 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param uri_type: 回调地址类型 + :param uri: 回调地址 + :return: Client 回调地址,不存在时返回 None + """ + + result = await db.execute( + select(SysOAuthClientUri).where( + SysOAuthClientUri.client_pk == client_pk, + SysOAuthClientUri.uri_type == uri_type, + SysOAuthClientUri.uri_hash == OidcUtil.sha256_digest(uri), + SysOAuthClientUri.uri == uri, + SysOAuthClientUri.status == '0', + ) + ) + + return result.scalars().first() + + @classmethod + async def find_uri( + cls, db: AsyncSession, client_pk: int, uri_type: str, uri_hash: str, uri: str + ) -> SysOAuthClientUri | None: + """ + 按地址摘要查询 OAuth Client 回调地址 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param uri_type: 回调地址类型 + :param uri_hash: 回调地址摘要 + :param uri: 回调地址 + :return: Client 回调地址,不存在时返回 None + """ + + result = await db.execute( + select(SysOAuthClientUri).where( + SysOAuthClientUri.client_pk == client_pk, + SysOAuthClientUri.uri_type == uri_type, + SysOAuthClientUri.uri_hash == uri_hash, + SysOAuthClientUri.uri == uri, + ) + ) + + return result.scalars().first() + + @classmethod + async def get_uri_for_update(cls, db: AsyncSession, uri_id: int) -> SysOAuthClientUri | None: + """ + 按内部主键锁定查询 Client 回调地址 + + :param db: orm对象 + :param uri_id: 回调地址内部主键 + :return: Client 回调地址,不存在时返回 None + """ + + result = await db.execute(select(SysOAuthClientUri).where(SysOAuthClientUri.uri_id == uri_id).with_for_update()) + + return result.scalars().first() + + @classmethod + async def list_scope_bindings(cls, db: AsyncSession, client_pk: int) -> Sequence[SysOAuthClientScope]: + """ + 查询 Client 权限范围绑定列表 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :return: Client Scope 绑定序列 + """ + + result = await db.execute(select(SysOAuthClientScope).where(SysOAuthClientScope.client_pk == client_pk)) + + return result.scalars().all() + + @classmethod + async def list_resource_bindings(cls, db: AsyncSession, client_pk: int) -> Sequence[SysOAuthClientResource]: + """ + 查询 Client Resource 绑定列表 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :return: Client Resource 绑定序列 + """ + + result = await db.execute(select(SysOAuthClientResource).where(SysOAuthClientResource.client_pk == client_pk)) + + return result.scalars().all() + + @classmethod + async def list_scopes(cls, db: AsyncSession, client_pk: int) -> Sequence[SysOAuthScope]: + """ + 查询 Client 权限范围列表 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :return: OAuth Scope 序列 + """ + + result = await db.execute( + select(SysOAuthScope) + .join(SysOAuthClientScope, SysOAuthClientScope.scope_pk == SysOAuthScope.scope_pk) + .where(SysOAuthClientScope.client_pk == client_pk, SysOAuthScope.status == '0') + .order_by(SysOAuthScope.scope_pk) + ) + + return result.scalars().all() + + @classmethod + async def list_resources(cls, db: AsyncSession, client_pk: int) -> Sequence[SysOAuthResource]: + """ + 查询 Client Resource 列表 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :return: OAuth Resource 序列 + """ + + result = await db.execute( + select(SysOAuthResource) + .join(SysOAuthClientResource, SysOAuthClientResource.resource_pk == SysOAuthResource.resource_pk) + .where(SysOAuthClientResource.client_pk == client_pk, SysOAuthResource.status == '0') + .order_by(SysOAuthResource.resource_pk) + ) + + return result.scalars().all() + + @classmethod + async def list_scope_definitions(cls, db: AsyncSession, active_only: bool = True) -> Sequence[SysOAuthScope]: + """ + 查询启用或全部 Scope 定义 + + :param db: orm对象 + :param active_only: 是否仅查询启用记录 + :return: OAuth Scope 定义序列 + """ + + query = select(SysOAuthScope) + if active_only: + query = query.where(SysOAuthScope.status == '0') + result = await db.execute(query.order_by(SysOAuthScope.scope_pk)) + + return result.scalars().all() + + @classmethod + async def get_active_scopes_by_codes(cls, db: AsyncSession, scope_codes: Sequence[str]) -> Sequence[SysOAuthScope]: + """ + 按编码批量查询启用 Scope + + :param db: orm对象 + :param scope_codes: Scope 编码序列 + :return: 启用的 OAuth Scope 序列 + """ + + if not scope_codes: + return () + result = await db.execute( + select(SysOAuthScope).where(SysOAuthScope.scope_code.in_(scope_codes), SysOAuthScope.status == '0') + ) + + return result.scalars().all() + + @classmethod + async def get_active_resources_by_ids( + cls, db: AsyncSession, resource_ids: Sequence[str] + ) -> Sequence[SysOAuthResource]: + """ + 按资源标识批量查询启用 Resource + + :param db: orm对象 + :param resource_ids: Resource 公开标识序列 + :return: 启用的 OAuth Resource 序列 + """ + + if not resource_ids: + return () + result = await db.execute( + select(SysOAuthResource).where( + SysOAuthResource.resource_id.in_(resource_ids), SysOAuthResource.status == '0' + ) + ) + + return result.scalars().all() + + @classmethod + async def replace_bindings( + cls, + db: AsyncSession, + client_pk: int, + scope_rows: Sequence[SysOAuthScope], + resource_rows: Sequence[SysOAuthResource], + uri_values: Mapping[str, Sequence[str]], + pre_authorized: set[str], + now: datetime, + *, + allowed_role_keys: Sequence[str] = (), + ) -> None: + """ + 替换 Client 的 Scope、Resource 与回调地址绑定 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param scope_rows: Scope 绑定对象序列 + :param resource_rows: Resource 绑定对象序列 + :param uri_values: 回调地址值序列 + :param pre_authorized: 是否预授权 + :param now: 当前时间 + :param allowed_role_keys: 允许向客户端发布的角色权限字符,空列表表示不发布角色 + :return: None + """ + + await db.execute(delete(SysOAuthClientScope).where(SysOAuthClientScope.client_pk == client_pk)) + await db.execute(delete(SysOAuthClientResource).where(SysOAuthClientResource.client_pk == client_pk)) + await db.execute(delete(SysOAuthClientUri).where(SysOAuthClientUri.client_pk == client_pk)) + db.add_all( + [ + SysOAuthClientScope( + client_pk=client_pk, + scope_pk=scope.scope_pk, + is_default=0, + pre_authorized=int(scope.scope_code in pre_authorized), + claim_filter={'claims': ['roles'], 'allowed_role_keys': list(allowed_role_keys)} + if scope.scope_code == 'roles' + else None, + create_time=now, + ) + for scope in scope_rows + ] + ) + db.add_all( + [ + SysOAuthClientResource( + client_pk=client_pk, + resource_pk=resource.resource_pk, + is_default=int(index == 0), + create_time=now, + ) + for index, resource in enumerate(resource_rows) + ] + ) + db.add_all( + [ + SysOAuthClientUri( + client_pk=client_pk, + uri_type=uri_type, + uri=uri, + uri_hash=OidcUtil.sha256_digest(uri), + is_default=int(index == 0), + status='0', + create_time=now, + ) + for uri_type, uris in uri_values.items() + for index, uri in enumerate(uris) + ] + ) + + @classmethod + async def add_secret(cls, db: AsyncSession, secret: SysOAuthClientSecret) -> None: + """ + 新增 OAuth Client Secret 并刷新主键 + + :param db: orm对象 + :param secret: OAuth Client Secret 对象 + :return: None + """ + + db.add(secret) + await db.flush() + + @classmethod + async def find_active_introspection_client(cls, db: AsyncSession, client_id: str) -> SysOAuthClient | None: + """ + 按公开标识查询启用的内省 OAuth Client + + :param db: orm对象 + :param client_id: Client 公开标识 + :return: 内省 OAuth Client,不存在时返回 None + """ + + result = await db.execute( + select(SysOAuthClient).where( + SysOAuthClient.client_id == client_id, + SysOAuthClient.status == '0', + SysOAuthClient.client_type == 'confidential', + SysOAuthClient.token_endpoint_auth_method == 'client_secret_basic', + ) + ) + + return result.scalars().first() + + @classmethod + async def lock_clients_for_resource(cls, db: AsyncSession, resource_pk: int) -> Sequence[SysOAuthClient]: + """ + 锁定查询 Resource 关联的 OAuth Client + + :param db: orm对象 + :param resource_pk: Resource 内部主键 + :return: 按 Resource 关联的 OAuth Client 序列 + """ + + direct = await db.execute( + select(SysOAuthClient.client_pk) + .join(SysOAuthClientResource, SysOAuthClientResource.client_pk == SysOAuthClient.client_pk) + .where(SysOAuthClientResource.resource_pk == resource_pk) + ) + via_scope = await db.execute( + select(SysOAuthClientScope.client_pk) + .join(SysOAuthScope, SysOAuthScope.scope_pk == SysOAuthClientScope.scope_pk) + .where(SysOAuthScope.resource_pk == resource_pk) + ) + client_pks = {row[0] for row in (*direct.all(), *via_scope.all())} + if not client_pks: + return () + result = await db.execute( + select(SysOAuthClient).where(SysOAuthClient.client_pk.in_(client_pks)).with_for_update() + ) + + return result.scalars().all() + + @classmethod + async def lock_clients_for_scope(cls, db: AsyncSession, scope_pk: int) -> Sequence[SysOAuthClient]: + """ + 锁定查询 Scope 关联的 OAuth Client + + :param db: orm对象 + :param scope_pk: Scope 内部主键 + :return: 按 Scope 关联的 OAuth Client 序列 + """ + + result = await db.execute( + select(SysOAuthClient) + .join(SysOAuthClientScope, SysOAuthClientScope.client_pk == SysOAuthClient.client_pk) + .where(SysOAuthClientScope.scope_pk == scope_pk) + .with_for_update() + ) + + return result.scalars().all() + + @classmethod + async def revoke_client_credentials(cls, db: AsyncSession, client_pk: int, now: datetime) -> None: + """ + 撤销 OAuth Client 的全部凭据 + + :param db: orm对象 + :param client_pk: Client 内部主键 + :param now: 当前时间 + :return: None + """ + + await db.execute( + update(SysOAuthGrant) + .where(SysOAuthGrant.client_pk == client_pk, SysOAuthGrant.status == 'active') + .values(status='revoked', revoked_at=now, revoke_reason='client_disabled') + ) + await db.execute( + update(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.client_pk == client_pk, SysOAuthRefreshToken.status == 'active') + .values(status='revoked', revoked_at=now, revoke_reason='client_disabled') + ) + + @classmethod + async def get_client_detail_rows( + cls, db: AsyncSession, client_id: str + ) -> tuple[SysOAuthClient, list[tuple[SysOAuthScope, int]], list[str], Sequence[SysOAuthClientUri]] | None: + """ + 查询 OAuth Client 详情及关联数据 + + :param db: orm对象 + :param client_id: Client 公开标识 + :return: OAuth Client、Scope 绑定统计、Resource 标识和回调地址组成的元组,不存在时返回 None + """ + + client = await cls.get_by_client_id(db, client_id, active_only=False) + if client is None: + return None + scopes = await db.execute( + select(SysOAuthScope, SysOAuthClientScope.pre_authorized) + .join(SysOAuthClientScope, SysOAuthClientScope.scope_pk == SysOAuthScope.scope_pk) + .where(SysOAuthClientScope.client_pk == client.client_pk) + .order_by(SysOAuthScope.scope_pk) + ) + resources = await db.execute( + select(SysOAuthResource.resource_id) + .join(SysOAuthClientResource, SysOAuthClientResource.resource_pk == SysOAuthResource.resource_pk) + .where(SysOAuthClientResource.client_pk == client.client_pk) + .order_by(SysOAuthResource.resource_pk) + ) + uris = await cls.list_uris(db, client.client_pk) + + return client, scopes.all(), list(resources.scalars().all()), uris + + @classmethod + async def list_clients_page(cls, db: AsyncSession, page: ClientPageQueryModel) -> Sequence[SysOAuthClient]: + """ + 分页查询 OAuth Client + + :param db: orm对象 + :param page: 分页查询条件对象 + :return: OAuth Client 序列 + """ + + conditions = [] + if page.client_name: + conditions.append(SysOAuthClient.client_name.contains(page.client_name, autoescape=True, escape='\\')) + if page.client_type: + conditions.append(SysOAuthClient.client_type == page.client_type) + if page.status: + conditions.append(SysOAuthClient.status == page.status) + result = await db.execute( + select(SysOAuthClient) + .where(*conditions) + .order_by(SysOAuthClient.client_pk) + .offset((page.page_num - 1) * page.page_size) + .limit(page.page_size) + ) + + return result.scalars().all() + + @classmethod + async def count_clients(cls, db: AsyncSession, page: ClientPageQueryModel) -> int: + """ + 统计 OAuth Client 数量 + + :param db: orm对象 + :param page: 分页查询条件对象 + :return: OAuth Client 数量 + """ + + conditions = [] + if page.client_name: + conditions.append(SysOAuthClient.client_name.contains(page.client_name, autoescape=True, escape='\\')) + if page.client_type: + conditions.append(SysOAuthClient.client_type == page.client_type) + if page.status: + conditions.append(SysOAuthClient.status == page.status) + result = await db.execute(select(func.count()).select_from(SysOAuthClient).where(*conditions)) + + return int(result.scalar_one()) + + @classmethod + async def list_cors_origins(cls, db: AsyncSession) -> tuple[str, ...]: + """ + 查询 OAuth Client 跨域来源 + + :param db: orm对象 + :return: 去重后的 CORS 来源元组 + """ + + result = await db.execute( + select(SysOAuthClientUri.uri) + .join(SysOAuthClient, SysOAuthClient.client_pk == SysOAuthClientUri.client_pk) + .where( + SysOAuthClient.status == '0', + SysOAuthClientUri.status == '0', + SysOAuthClientUri.uri_type == 'cors_origin', + ) + .order_by(SysOAuthClientUri.uri) + ) + + return tuple(dict.fromkeys(result.scalars().all())) diff --git a/ruoyi-fastapi-backend/module_identity/dao/oauth_grant_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oauth_grant_dao.py new file mode 100644 index 000000000..92fab90d8 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oauth_grant_dao.py @@ -0,0 +1,555 @@ +from collections.abc import Sequence +from dataclasses import dataclass +from datetime import datetime +from uuid import uuid4 + +from sqlalchemy import case, func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthAccessPolicy, SysOAuthGrant, SysOAuthRefreshToken +from utils.time_util import TimezoneUtil + + +@dataclass(frozen=True, slots=True) +class OAuthGrantSnapshot: + """ + OAuth Grant 可恢复快照 + """ + + grant_id: str + user_id: int + subject_id: str + client_pk: int + granted_scopes: tuple[str, ...] + granted_resources: tuple[str, ...] + client_policy_version: int + status: str + consented_at: datetime + expires_at: datetime | None + revoked_at: datetime | None + revoke_reason: str | None + last_used_at: datetime | None + remembered_scopes: tuple[str, ...] = () + remembered_resources: tuple[str, ...] = () + + +class OAuthGrantDao: + """ + OAuth Grant 数据库操作层 + """ + + @classmethod + async def get_by_grant_id_for_update( + cls, db: AsyncSession, grant_id: str, *, refresh: bool = False + ) -> SysOAuthGrant | None: + """ + 按授权标识锁定查询 OAuth Grant + + :param db: orm对象 + :param grant_id: Grant 公开标识 + :param refresh: 是否使用数据库当前值刷新已有对象 + :return: OAuth Grant,不存在时返回 None + """ + + query = select(SysOAuthGrant).where(SysOAuthGrant.grant_id == grant_id).with_for_update() + if refresh: + query = query.execution_options(populate_existing=True) + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def get_active_for_user_client( + cls, db: AsyncSession, user_id: int, client_pk: int, for_update: bool = False + ) -> SysOAuthGrant | None: + """ + 查询用户与 Client 的活跃 OAuth Grant + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :param for_update: 是否锁定查询结果 + :return: 活跃 OAuth Grant,不存在时返回 None + """ + + query = select(SysOAuthGrant).where( + SysOAuthGrant.user_id == user_id, + SysOAuthGrant.client_pk == client_pk, + SysOAuthGrant.status == 'active', + (SysOAuthGrant.expires_at.is_(None) | (SysOAuthGrant.expires_at > TimezoneUtil.utc_now())), + ) + if for_update: + query = query.with_for_update().execution_options(populate_existing=True) + result = await db.execute(query.order_by(SysOAuthGrant.consented_at.desc())) + + return result.scalars().first() + + @classmethod + async def list_active_for_user_client( + cls, db: AsyncSession, user_id: int, client_pk: int + ) -> Sequence[SysOAuthGrant]: + """ + 查询用户与 Client 的活跃 OAuth Grant 列表 + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :return: 活跃 OAuth Grant 序列 + """ + + result = await db.execute( + select(SysOAuthGrant).where( + SysOAuthGrant.user_id == user_id, + SysOAuthGrant.client_pk == client_pk, + SysOAuthGrant.status == 'active', + ) + ) + + return result.scalars().all() + + @classmethod + async def merge_active_grant( + cls, + db: AsyncSession, + user_id: int, + subject_id: str, + client_pk: int, + granted_scopes: list[str], + granted_resources: list[str], + client_policy_version: int, + expires_at: datetime | None = None, + *, + remember_consent: bool = True, + ) -> SysOAuthGrant: + """ + 合并用户与 Client 的活跃 OAuth Grant + + :param db: orm对象 + :param user_id: 用户编号 + :param subject_id: 主体标识 + :param client_pk: Client 内部主键 + :param granted_scopes: 授权 Scope 编码序列 + :param granted_resources: 授权 Resource 标识序列 + :param client_policy_version: Client 策略版本 + :param expires_at: 授权过期时间 + :param remember_consent: 是否允许后续请求复用本次同意,独立于离线授权 + :return: 新增或更新后的 OAuth Grant + """ + + await OAuthAccessPolicyDao.lock_client(db, client_pk) + grant = await cls.get_active_for_user_client(db, user_id, client_pk, for_update=True) + now = TimezoneUtil.utc_now() + if grant is None: + grant = SysOAuthGrant( + grant_id=str(uuid4()), + user_id=user_id, + subject_id=subject_id, + client_pk=client_pk, + granted_scopes=list(dict.fromkeys(granted_scopes)), + granted_resources=list(dict.fromkeys(granted_resources)), + remembered_scopes=list(dict.fromkeys(granted_scopes)) if remember_consent else [], + remembered_resources=list(dict.fromkeys(granted_resources)) if remember_consent else [], + client_policy_version=client_policy_version, + status='active', + consented_at=now, + expires_at=expires_at, + ) + db.add(grant) + else: + grant.subject_id = subject_id + grant.granted_scopes = list(dict.fromkeys([*(grant.granted_scopes or []), *granted_scopes])) + grant.granted_resources = list(dict.fromkeys([*(grant.granted_resources or []), *granted_resources])) + grant.remembered_scopes = ( + list(dict.fromkeys([*(grant.remembered_scopes or []), *granted_scopes])) if remember_consent else [] + ) + grant.remembered_resources = ( + list(dict.fromkeys([*(grant.remembered_resources or []), *granted_resources])) + if remember_consent + else [] + ) + grant.client_policy_version = client_policy_version + grant.consented_at = now + grant.expires_at = expires_at + grant.revoked_at = None + grant.revoke_reason = None + await db.flush() + + return grant + + @staticmethod + def snapshot(grant: SysOAuthGrant) -> OAuthGrantSnapshot: + """ + 生成 OAuth Grant 可恢复快照 + + :param grant: OAuth Grant 对象 + :return: OAuth Grant 快照 + """ + + return OAuthGrantSnapshot( + grant_id=grant.grant_id, + user_id=grant.user_id, + subject_id=grant.subject_id, + client_pk=grant.client_pk, + granted_scopes=tuple(grant.granted_scopes or ()), + granted_resources=tuple(grant.granted_resources or ()), + client_policy_version=grant.client_policy_version, + status=grant.status, + consented_at=grant.consented_at, + expires_at=grant.expires_at, + revoked_at=grant.revoked_at, + revoke_reason=grant.revoke_reason, + last_used_at=grant.last_used_at, + remembered_scopes=tuple(grant.remembered_scopes or ()), + remembered_resources=tuple(grant.remembered_resources or ()), + ) + + @classmethod + async def restore_snapshot( + cls, + db: AsyncSession, + snapshot: OAuthGrantSnapshot, + expected: OAuthGrantSnapshot, + ) -> bool: + """ + 恢复 OAuth Grant 可恢复快照 + + :param db: orm对象 + :param snapshot: OAuth Grant 快照 + :param expected: 期望状态 + :return: 是否恢复成功 + """ + + await OAuthAccessPolicyDao.lock_client(db, snapshot.client_pk) + grant = await cls.get_by_grant_id_for_update(db, snapshot.grant_id, refresh=True) + if grant is None or cls.snapshot(grant) != expected: + return False + grant.subject_id = snapshot.subject_id + grant.granted_scopes = list(snapshot.granted_scopes) + grant.granted_resources = list(snapshot.granted_resources) + grant.client_policy_version = snapshot.client_policy_version + grant.status = snapshot.status + grant.consented_at = snapshot.consented_at + grant.expires_at = snapshot.expires_at + grant.revoked_at = snapshot.revoked_at + grant.revoke_reason = snapshot.revoke_reason + grant.last_used_at = snapshot.last_used_at + grant.remembered_scopes = list(snapshot.remembered_scopes) + grant.remembered_resources = list(snapshot.remembered_resources) + await db.flush() + + return True + + @classmethod + async def revoke_snapshot(cls, db: AsyncSession, snapshot: OAuthGrantSnapshot, reason: str) -> bool: + """ + 撤销快照对应的 OAuth Grant + + :param db: orm对象 + :param snapshot: OAuth Grant 快照 + :param reason: 撤销或标记原因 + :return: 是否撤销成功 + """ + + await OAuthAccessPolicyDao.lock_client(db, snapshot.client_pk) + grant = await cls.get_by_grant_id_for_update(db, snapshot.grant_id, refresh=True) + if grant is None or cls.snapshot(grant) != snapshot: + return False + current = TimezoneUtil.utc_now() + grant.status = 'revoked' + grant.revoked_at = current + grant.revoke_reason = reason + await db.execute( + update(SysOAuthRefreshToken) + .where( + SysOAuthRefreshToken.grant_id == snapshot.grant_id, + SysOAuthRefreshToken.status.in_(['active', 'used', 'rotated']), + ) + .values(status='revoked', revoked_at=current, revoke_reason=reason) + ) + await db.flush() + + return True + + @classmethod + async def get_valid_for_user_client( + cls, db: AsyncSession, user_id: int, client_pk: int, for_update: bool = False + ) -> SysOAuthGrant | None: + """ + 查询用户与 Client 的有效 OAuth Grant + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :param for_update: 是否锁定查询结果 + :return: 有效 OAuth Grant,不存在时返回 None + """ + + query = ( + select(SysOAuthGrant) + .join(SysOAuthClient, SysOAuthClient.client_pk == SysOAuthGrant.client_pk) + .where( + SysOAuthGrant.user_id == user_id, + SysOAuthGrant.client_pk == client_pk, + SysOAuthGrant.status == 'active', + SysOAuthGrant.client_policy_version == SysOAuthClient.policy_version, + (SysOAuthGrant.expires_at.is_(None) | (SysOAuthGrant.expires_at > TimezoneUtil.utc_now())), + ) + .order_by(SysOAuthGrant.consented_at.desc()) + ) + if for_update: + query = query.with_for_update().execution_options(populate_existing=True) + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def revoke(cls, db: AsyncSession, grant_id: str, reason: str | None = None) -> bool: + """ + 撤销 OAuth Grant + + :param db: orm对象 + :param grant_id: Grant 公开标识 + :param reason: 撤销或标记原因 + :return: 是否撤销成功 + """ + + grant = await db.execute(select(SysOAuthGrant).where(SysOAuthGrant.grant_id == grant_id).with_for_update()) + row = grant.scalars().first() + if row is None: + return False + await db.execute( + update(SysOAuthRefreshToken) + .where( + SysOAuthRefreshToken.grant_id == grant_id, + SysOAuthRefreshToken.status.in_(['active', 'used', 'rotated']), + ) + .values(status='revoked', revoked_at=TimezoneUtil.utc_now(), revoke_reason=reason) + ) + result = await db.execute( + update(SysOAuthGrant) + .where(SysOAuthGrant.grant_id == grant_id, SysOAuthGrant.status == 'active') + .values( + status='revoked', + revoked_at=TimezoneUtil.utc_now(), + revoke_reason=reason, + remembered_scopes=[], + remembered_resources=[], + ) + ) + + return bool(result.rowcount) + + @staticmethod + async def targets(db: AsyncSession, grant_ids: list[str]) -> list[tuple[int, int]]: + """ + 将选中授权归并为用户和 Client,按 Client 排序以统一批量锁序 + + :param db: orm对象 + :param grant_ids: 选中的授权标识 + :return: 去重后的 Client 主键和用户编号列表 + """ + + result = await db.execute( + select(SysOAuthGrant.client_pk, SysOAuthGrant.user_id) + .where(SysOAuthGrant.grant_id.in_(grant_ids)) + .distinct() + .order_by(SysOAuthGrant.client_pk, SysOAuthGrant.user_id) + ) + return [(row.client_pk, row.user_id) for row in result] + + @classmethod + async def revoke_for_user_client(cls, db: AsyncSession, user_id: int, client_pk: int, reason: str) -> list[str]: + """ + 撤销用户对 Client 的全部现有授权;调用方须先锁定 Client + + :param db: orm对象 + :param user_id: 用户编号 + :param client_pk: Client 内部主键 + :param reason: 撤销原因 + :return: 实际撤销的授权标识列表 + """ + + result = await db.execute( + select(SysOAuthGrant.grant_id) + .where( + SysOAuthGrant.user_id == user_id, + SysOAuthGrant.client_pk == client_pk, + SysOAuthGrant.status == 'active', + ) + .order_by(SysOAuthGrant.grant_id) + .with_for_update() + ) + grant_ids = list(result.scalars().all()) + for grant_id in grant_ids: + await cls.revoke(db, grant_id, reason) + return grant_ids + + @classmethod + async def revoke_for_user(cls, db: AsyncSession, user_id: int, reason: str | None = None) -> int: + """ + 撤销用户的 OAuth Grant + + :param db: orm对象 + :param user_id: 用户编号 + :param reason: 撤销或标记原因 + :return: 已撤销的 OAuth Grant 数量 + """ + + result = await db.execute( + update(SysOAuthGrant) + .where(SysOAuthGrant.user_id == user_id, SysOAuthGrant.status == 'active') + .values(status='revoked', revoked_at=TimezoneUtil.utc_now(), revoke_reason=reason) + ) + + return result.rowcount or 0 + + @classmethod + async def list_for_user(cls, db: AsyncSession, user_id: int, status: str | None = None) -> Sequence[SysOAuthGrant]: + """ + 查询用户 OAuth Grant 列表 + + :param db: orm对象 + :param user_id: 用户编号 + :param status: 状态过滤值 + :return: 用户 OAuth Grant 序列 + """ + + query = select(SysOAuthGrant).where(SysOAuthGrant.user_id == user_id) + if status: + query = query.where(SysOAuthGrant.status == status) + result = await db.execute(query.order_by(SysOAuthGrant.consented_at.desc())) + + return result.scalars().all() + + @classmethod + async def get_by_grant_id(cls, db: AsyncSession, grant_id: str, for_update: bool = False) -> SysOAuthGrant | None: + """ + 按授权标识查询 OAuth Grant + + :param db: orm对象 + :param grant_id: Grant 公开标识 + :param for_update: 是否锁定查询结果 + :return: OAuth Grant,不存在时返回 None + """ + + query = select(SysOAuthGrant).where(SysOAuthGrant.grant_id == grant_id) + if for_update: + query = query.with_for_update().execution_options(populate_existing=True) + if not for_update: + return await db.scalar(query) + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def list_page( + cls, + db: AsyncSession, + *, + user_id: int | None = None, + client_id: str | None = None, + status: str | None = None, + access_status: str | None = None, + offset: int = 0, + limit: int = 200, + ) -> Sequence[SysOAuthGrant]: + """ + 分页查询 OAuth Grant + + :param db: orm对象 + :param user_id: 用户编号 + :param client_id: Client 公开标识 + :param status: 状态过滤值 + :param access_status: 用户对应用的访问策略 + :param offset: 分页偏移量 + :param limit: 分页大小 + :return: OAuth Grant 序列 + """ + + conditions = [] + if user_id is not None: + conditions.append(SysOAuthGrant.user_id == user_id) + if client_id is not None: + conditions.append(SysOAuthClient.client_id == client_id) + if status: + effective_status = case( + ((SysOAuthGrant.status == 'active') & (SysOAuthGrant.expires_at <= TimezoneUtil.utc_now()), 'expired'), + else_=SysOAuthGrant.status, + ) + conditions.append(effective_status == status) + if access_status: + blocked = ( + select(SysOAuthAccessPolicy.user_id) + .where( + SysOAuthAccessPolicy.user_id == SysOAuthGrant.user_id, + SysOAuthAccessPolicy.client_pk == SysOAuthGrant.client_pk, + SysOAuthAccessPolicy.access_status == 'blocked', + ) + .exists() + ) + conditions.append(blocked if access_status == 'blocked' else ~blocked) + query = ( + select(SysOAuthGrant) + .join(SysOAuthClient, SysOAuthClient.client_pk == SysOAuthGrant.client_pk) + .where(*conditions) + .order_by(SysOAuthGrant.consented_at.desc()) + .offset(max(offset, 0)) + .limit(min(max(limit, 1), 200)) + ) + result = await db.execute(query) + + return result.scalars().all() + + @classmethod + async def count( + cls, + db: AsyncSession, + *, + user_id: int | None = None, + client_id: str | None = None, + status: str | None = None, + access_status: str | None = None, + ) -> int: + """ + 统计 OAuth Grant 数量 + + :param db: orm对象 + :param user_id: 用户编号 + :param client_id: Client 公开标识 + :param status: 状态过滤值 + :param access_status: 用户对应用的访问策略 + :return: OAuth Grant 数量 + """ + + conditions = [] + if user_id is not None: + conditions.append(SysOAuthGrant.user_id == user_id) + if client_id is not None: + conditions.append(SysOAuthClient.client_id == client_id) + if status: + effective_status = case( + ((SysOAuthGrant.status == 'active') & (SysOAuthGrant.expires_at <= TimezoneUtil.utc_now()), 'expired'), + else_=SysOAuthGrant.status, + ) + conditions.append(effective_status == status) + if access_status: + blocked = ( + select(SysOAuthAccessPolicy.user_id) + .where( + SysOAuthAccessPolicy.user_id == SysOAuthGrant.user_id, + SysOAuthAccessPolicy.client_pk == SysOAuthGrant.client_pk, + SysOAuthAccessPolicy.access_status == 'blocked', + ) + .exists() + ) + conditions.append(blocked if access_status == 'blocked' else ~blocked) + result = await db.execute( + select(func.count()) + .select_from(SysOAuthGrant) + .join(SysOAuthClient, SysOAuthClient.client_pk == SysOAuthGrant.client_pk) + .where(*conditions) + ) + + return int(result.scalar_one()) diff --git a/ruoyi-fastapi-backend/module_identity/dao/oauth_resource_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oauth_resource_dao.py new file mode 100644 index 000000000..192a274bd --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oauth_resource_dao.py @@ -0,0 +1,407 @@ +from collections.abc import Sequence +from datetime import datetime + +from sqlalchemy import func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.oauth_resource_vo import ResourcePageQueryModel, ScopePageQueryModel + + +class OAuthResourceDao: + """ + OAuth Resource 数据库操作层 + """ + + @classmethod + async def add_resource(cls, db: AsyncSession, resource: SysOAuthResource) -> None: + """ + 新增 OAuth Resource 并刷新主键 + + :param db: orm对象 + :param resource: OAuth Resource 对象 + :return: None + """ + + db.add(resource) + await db.flush() + + @classmethod + async def add_scope(cls, db: AsyncSession, scope: SysOAuthScope) -> None: + """ + 新增 OAuth Scope 并刷新主键 + + :param db: orm对象 + :param scope: OAuth Scope 对象 + :return: None + """ + + db.add(scope) + await db.flush() + + @classmethod + async def persist_resource_change(cls, db: AsyncSession, resource: SysOAuthResource) -> None: + """ + 刷新 OAuth Resource 变更 + + :param db: orm对象 + :param resource: OAuth Resource 对象 + :return: None + """ + + await db.flush() + + @classmethod + async def persist_scope_change(cls, db: AsyncSession, scope: SysOAuthScope) -> None: + """ + 刷新 OAuth Scope 变更 + + :param db: orm对象 + :param scope: OAuth Scope 对象 + :return: None + """ + + await db.flush() + + @classmethod + async def get_resource( + cls, db: AsyncSession, resource_id: str, *, for_update: bool = False + ) -> SysOAuthResource | None: + """ + 按资源标识查询 OAuth Resource + + :param db: orm对象 + :param resource_id: Resource 公开标识 + :param for_update: 是否锁定查询结果 + :return: OAuth Resource,不存在时返回 None + """ + + query = select(SysOAuthResource).where(SysOAuthResource.resource_id == resource_id) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def find_resource_duplicate(cls, db: AsyncSession, resource_id: str, audience: str) -> bool: + """ + 按资源标识和受众查询重复 Resource + + :param db: orm对象 + :param resource_id: Resource 公开标识 + :param audience: Resource 受众 + :return: 是否存在重复 Resource + """ + + result = await db.execute( + select(SysOAuthResource.resource_pk).where( + (SysOAuthResource.resource_id == resource_id) | (SysOAuthResource.audience == audience) + ) + ) + + return result.scalar_one_or_none() is not None + + @classmethod + async def list_resources_page(cls, db: AsyncSession, page: ResourcePageQueryModel) -> Sequence[SysOAuthResource]: + """ + 分页查询 OAuth Resource + + :param db: orm对象 + :param page: 分页查询条件对象 + :return: OAuth Resource 序列 + """ + + conditions = [] + if page.resource_name: + conditions.append(SysOAuthResource.resource_name.contains(page.resource_name, autoescape=True, escape='\\')) + if page.status: + conditions.append(SysOAuthResource.status == page.status) + result = await db.execute( + select(SysOAuthResource) + .where(*conditions) + .order_by(SysOAuthResource.resource_pk) + .offset((page.page_num - 1) * page.page_size) + .limit(page.page_size) + ) + + return result.scalars().all() + + @classmethod + async def count_resources(cls, db: AsyncSession, page: ResourcePageQueryModel) -> int: + """ + 统计 OAuth Resource 数量 + + :param db: orm对象 + :param page: 分页查询条件对象 + :return: OAuth Resource 数量 + """ + + conditions = [] + if page.resource_name: + conditions.append(SysOAuthResource.resource_name.contains(page.resource_name, autoescape=True, escape='\\')) + if page.status: + conditions.append(SysOAuthResource.status == page.status) + result = await db.execute(select(func.count()).select_from(SysOAuthResource).where(*conditions)) + + return int(result.scalar_one()) + + @classmethod + async def active_resource(cls, db: AsyncSession, resource_id: str | None) -> SysOAuthResource | None: + """ + 按资源标识查询启用 OAuth Resource + + :param db: orm对象 + :param resource_id: Resource 公开标识 + :return: 启用 OAuth Resource,不存在时返回 None + """ + + if not resource_id: + return None + result = await db.execute( + select(SysOAuthResource).where(SysOAuthResource.resource_id == resource_id, SysOAuthResource.status == '0') + ) + + return result.scalars().first() + + @classmethod + async def active_by_audiences(cls, db: AsyncSession, audiences: Sequence[str]) -> Sequence[SysOAuthResource]: + """ + 按受众批量查询启用 OAuth Resource + + :param db: orm对象 + :param audiences: Resource 受众序列 + :return: 启用 OAuth Resource 序列 + """ + + if not audiences: + return () + result = await db.execute( + select(SysOAuthResource).where(SysOAuthResource.audience.in_(audiences), SysOAuthResource.status == '0') + ) + + return result.scalars().all() + + @classmethod + async def get_scope(cls, db: AsyncSession, scope_code: str, *, for_update: bool = False) -> SysOAuthScope | None: + """ + 按 Scope 编码查询 OAuth Scope + + :param db: orm对象 + :param scope_code: Scope 编码 + :param for_update: 是否锁定查询结果 + :return: OAuth Scope,不存在时返回 None + """ + + query = select(SysOAuthScope).where(SysOAuthScope.scope_code == scope_code) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def find_scope_duplicate(cls, db: AsyncSession, scope_code: str) -> bool: + """ + 按 Scope 编码查询重复 OAuth Scope + + :param db: orm对象 + :param scope_code: Scope 编码 + :return: 是否存在重复 Scope + """ + + result = await db.execute(select(SysOAuthScope.scope_pk).where(SysOAuthScope.scope_code == scope_code)) + + return result.scalar_one_or_none() is not None + + @classmethod + async def list_scopes_page(cls, db: AsyncSession, page: ScopePageQueryModel) -> Sequence[SysOAuthScope]: + """ + 分页查询 OAuth Scope + + :param db: orm对象 + :param page: 分页查询条件对象 + :return: OAuth Scope 序列 + """ + + conditions = [] + if page.scope_name: + conditions.append(SysOAuthScope.scope_name.contains(page.scope_name, autoescape=True, escape='\\')) + if page.scope_type: + conditions.append(SysOAuthScope.scope_type == page.scope_type) + if page.status: + conditions.append(SysOAuthScope.status == page.status) + result = await db.execute( + select(SysOAuthScope) + .where(*conditions) + .order_by(SysOAuthScope.scope_pk) + .offset((page.page_num - 1) * page.page_size) + .limit(page.page_size) + ) + + return result.scalars().all() + + @classmethod + async def count_scopes(cls, db: AsyncSession, page: ScopePageQueryModel) -> int: + """ + 统计 OAuth Scope 数量 + + :param db: orm对象 + :param page: 分页查询条件对象 + :return: OAuth Scope 数量 + """ + + conditions = [] + if page.scope_name: + conditions.append(SysOAuthScope.scope_name.contains(page.scope_name, autoescape=True, escape='\\')) + if page.scope_type: + conditions.append(SysOAuthScope.scope_type == page.scope_type) + if page.status: + conditions.append(SysOAuthScope.status == page.status) + result = await db.execute(select(func.count()).select_from(SysOAuthScope).where(*conditions)) + + return int(result.scalar_one()) + + @classmethod + async def resource_id_for_scope(cls, db: AsyncSession, resource_pk: int) -> str | None: + """ + 按 Resource 主键查询资源标识 + + :param db: orm对象 + :param resource_pk: Resource 内部主键 + :return: Resource 公开标识,不存在时返回 None + """ + + result = await db.execute( + select(SysOAuthResource.resource_id).where(SysOAuthResource.resource_pk == resource_pk) + ) + + return result.scalar_one_or_none() + + @classmethod + async def lock_scope_clients(cls, db: AsyncSession, scope_pk: int) -> Sequence[SysOAuthClient]: + """ + 锁定查询 Scope 关联的 OAuth Client + + :param db: orm对象 + :param scope_pk: Scope 内部主键 + :return: Scope 关联的 OAuth Client 序列 + """ + + result = await db.execute( + select(SysOAuthClient) + .join(SysOAuthClientScope, SysOAuthClientScope.client_pk == SysOAuthClient.client_pk) + .where(SysOAuthClientScope.scope_pk == scope_pk) + .with_for_update() + ) + + return result.scalars().all() + + @classmethod + async def client_ids_for_resource(cls, db: AsyncSession, resource_pk: int) -> Sequence[tuple[int, str]]: + """ + 查询 Resource 直接绑定的 Client 主键和公开标识 + + :param db: orm对象 + :param resource_pk: Resource 内部主键 + :return: Client 内部主键和公开标识元组序列 + """ + + result = await db.execute( + select(SysOAuthClient.client_pk, SysOAuthClient.client_id) + .join(SysOAuthClientResource, SysOAuthClientResource.client_pk == SysOAuthClient.client_pk) + .where(SysOAuthClientResource.resource_pk == resource_pk) + .order_by(SysOAuthClient.client_id) + ) + + return result.all() + + @classmethod + async def client_ids_for_resource_scope(cls, db: AsyncSession, resource_pk: int) -> Sequence[tuple[int, str]]: + """ + 查询 Resource 通过 Scope 绑定的 Client 主键和公开标识 + + :param db: orm对象 + :param resource_pk: Resource 内部主键 + :return: 通过 Scope 关联的 Client 内部主键和公开标识元组序列 + """ + + result = await db.execute( + select(SysOAuthClient.client_pk, SysOAuthClient.client_id) + .join(SysOAuthClientScope, SysOAuthClientScope.client_pk == SysOAuthClient.client_pk) + .join(SysOAuthScope, SysOAuthScope.scope_pk == SysOAuthClientScope.scope_pk) + .where(SysOAuthScope.resource_pk == resource_pk) + .order_by(SysOAuthClient.client_id) + ) + + return result.all() + + @classmethod + async def active_grants(cls, db: AsyncSession, client_pks: Sequence[int]) -> Sequence[SysOAuthGrant]: + """ + 按 Client 主键查询活跃 OAuth Grant + + :param db: orm对象 + :param client_pks: Client 内部主键序列 + :return: 活跃 OAuth Grant 序列 + """ + + if not client_pks: + return () + result = await db.execute( + select(SysOAuthGrant).where(SysOAuthGrant.status == 'active', SysOAuthGrant.client_pk.in_(client_pks)) + ) + + return result.scalars().all() + + @classmethod + async def active_refresh_tokens(cls, db: AsyncSession, client_pks: Sequence[int]) -> Sequence[SysOAuthRefreshToken]: + """ + 按 Client 主键查询活跃 Refresh Token + + :param db: orm对象 + :param client_pks: Client 内部主键序列 + :return: 活跃 OAuth Refresh Token 序列 + """ + + if not client_pks: + return () + result = await db.execute( + select(SysOAuthRefreshToken).where( + SysOAuthRefreshToken.status == 'active', SysOAuthRefreshToken.client_pk.in_(client_pks) + ) + ) + + return result.scalars().all() + + @classmethod + async def revoke_credentials(cls, db: AsyncSession, client_pks: Sequence[int], now: datetime, reason: str) -> None: + """ + 撤销 Client 主键对应的 Grant 和 Refresh Token + + :param db: orm对象 + :param client_pks: Client 内部主键序列 + :param now: 当前时间 + :param reason: 撤销或标记原因 + :return: None + """ + + if not client_pks: + return + await db.execute( + update(SysOAuthGrant) + .where(SysOAuthGrant.client_pk.in_(client_pks), SysOAuthGrant.status == 'active') + .values(status='revoked', revoked_at=now, revoke_reason=reason) + ) + await db.execute( + update(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.client_pk.in_(client_pks), SysOAuthRefreshToken.status == 'active') + .values(status='revoked', revoked_at=now, revoke_reason=reason) + ) diff --git a/ruoyi-fastapi-backend/module_identity/dao/oauth_token_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oauth_token_dao.py new file mode 100644 index 000000000..cb31104cf --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oauth_token_dao.py @@ -0,0 +1,294 @@ +from collections.abc import Sequence +from datetime import datetime + +from sqlalchemy import select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.oauth_grant_do import SysOAuthRefreshToken +from utils.time_util import TimezoneUtil + + +class OAuthTokenDao: + """ + OAuth Token 数据库操作层 + """ + + @classmethod + async def create(cls, db: AsyncSession, row: SysOAuthRefreshToken) -> SysOAuthRefreshToken: + """ + 新增 OAuth Refresh Token 并刷新主键 + + :param db: orm对象 + :param row: OAuth Refresh Token 对象 + :return: 已写入的 OAuth Refresh Token + """ + + db.add(row) + await db.flush() + + return row + + @classmethod + async def get_by_token_id( + cls, db: AsyncSession, token_id: str, for_update: bool = False + ) -> SysOAuthRefreshToken | None: + """ + 按 Token 标识查询 OAuth Refresh Token + + :param db: orm对象 + :param token_id: Token 公开标识 + :param for_update: 是否锁定查询结果 + :return: OAuth Refresh Token,不存在时返回 None + """ + + query = select(SysOAuthRefreshToken).where(SysOAuthRefreshToken.token_id == token_id) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def get_by_hash( + cls, db: AsyncSession, token_hash: str, for_update: bool = False + ) -> SysOAuthRefreshToken | None: + """ + 按 Token 摘要查询 OAuth Refresh Token + + :param db: orm对象 + :param token_hash: Token 摘要 + :param for_update: 是否锁定查询结果 + :return: OAuth Refresh Token,不存在时返回 None + """ + + query = select(SysOAuthRefreshToken).where(SysOAuthRefreshToken.token_hash == token_hash) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def lock_family(cls, db: AsyncSession, family_id: str) -> Sequence[SysOAuthRefreshToken]: + """ + 锁定查询 OAuth Refresh Token 家族 + + :param db: orm对象 + :param family_id: Token 家族标识 + :return: OAuth Refresh Token 序列 + """ + + result = await db.execute( + select(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.family_id == family_id) + .order_by(SysOAuthRefreshToken.issued_at) + .with_for_update() + ) + + return result.scalars().all() + + @classmethod + async def family_is_active(cls, db: AsyncSession, family_id: str) -> bool: + """ + 查询 OAuth Refresh Token 家族是否活跃 + + :param db: orm对象 + :param family_id: Token 家族标识 + :return: OAuth Refresh Token 家族是否活跃 + """ + + rows = await cls.lock_family(db, family_id) + + return bool(rows) and not any(row.status in {'revoked', 'reuse_detected', 'family_revoked'} for row in rows) + + @classmethod + async def list_for_sid_for_update(cls, db: AsyncSession, sid: str) -> Sequence[SysOAuthRefreshToken]: + """ + 按 Session 标识锁定查询 Refresh Token + + :param db: orm对象 + :param sid: SSO Session 标识 + :return: SSO Session 关联的 Refresh Token 序列 + """ + + result = await db.execute(select(SysOAuthRefreshToken).where(SysOAuthRefreshToken.sid == sid).with_for_update()) + + return result.scalars().all() + + @classmethod + async def mark_used( + cls, + db: AsyncSession, + token_id: str, + replaced_by_token_id: str | None = None, + *, + now: datetime | None = None, + ) -> bool: + """ + 标记 OAuth Refresh Token 已使用 + + :param db: orm对象 + :param token_id: Token 公开标识 + :param replaced_by_token_id: 替换 Token 标识 + :param now: 当前时间 + :return: 是否更新成功 + """ + + values: dict[str, object] = {'status': 'used', 'last_used_at': now or TimezoneUtil.utc_now()} + if replaced_by_token_id is not None: + values['replaced_by_token_id'] = replaced_by_token_id + result = await db.execute( + update(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.token_id == token_id, SysOAuthRefreshToken.status == 'active') + .values(**values) + ) + + return bool(result.rowcount) + + @classmethod + async def mark_reuse_detected(cls, db: AsyncSession, token_id: str, reason: str = 'refresh_token_reuse') -> bool: + """ + 标记 OAuth Refresh Token 重用 + + :param db: orm对象 + :param token_id: Token 公开标识 + :param reason: 撤销或标记原因 + :return: 是否更新成功 + """ + + now = TimezoneUtil.utc_now() + result = await db.execute( + update(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.token_id == token_id) + .values(status='reuse_detected', reuse_detected_at=now, revoked_at=now, revoke_reason=reason) + ) + + return bool(result.rowcount) + + @classmethod + async def refresh_token_family_reuse( + cls, + db: AsyncSession, + family_id: str, + offending_token_id: str, + reason: str = 'refresh_token_reuse', + *, + now: datetime | None = None, + ) -> int: + """ + 处理 OAuth Refresh Token 家族重用 + + :param db: orm对象 + :param family_id: Token 家族标识 + :param offending_token_id: 检测到重用的 Token 标识 + :param reason: 撤销或标记原因 + :param now: 当前时间 + :return: 已撤销的 OAuth Refresh Token 数量 + """ + + await cls.lock_family(db, family_id) + current = now or TimezoneUtil.utc_now() + other_result = await db.execute( + update(SysOAuthRefreshToken) + .where( + SysOAuthRefreshToken.family_id == family_id, + SysOAuthRefreshToken.token_id != offending_token_id, + SysOAuthRefreshToken.status.not_in(['revoked', 'expired', 'reuse_detected']), + ) + .values(status='revoked', revoked_at=current, revoke_reason=reason) + ) + offending_result = await db.execute( + update(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.family_id == family_id, SysOAuthRefreshToken.token_id == offending_token_id) + .values(status='reuse_detected', reuse_detected_at=current, revoked_at=current, revoke_reason=reason) + ) + + return (other_result.rowcount or 0) + (offending_result.rowcount or 0) + + @classmethod + async def revoke_family(cls, db: AsyncSession, family_id: str, reason: str | None = None) -> int: + """ + 撤销 OAuth Refresh Token 家族 + + :param db: orm对象 + :param family_id: Token 家族标识 + :param reason: 撤销或标记原因 + :return: 已撤销的 OAuth Refresh Token 数量 + """ + + result = await db.execute( + update(SysOAuthRefreshToken) + .where( + SysOAuthRefreshToken.family_id == family_id, SysOAuthRefreshToken.status.not_in(['revoked', 'expired']) + ) + .values(status='revoked', revoked_at=TimezoneUtil.utc_now(), revoke_reason=reason) + ) + await db.flush() + + return result.rowcount or 0 + + @classmethod + async def revoke_for_user(cls, db: AsyncSession, user_id: int, reason: str | None = None) -> int: + """ + 撤销用户的 OAuth Refresh Token + + :param db: orm对象 + :param user_id: 用户编号 + :param reason: 撤销或标记原因 + :return: 已撤销的 OAuth Refresh Token 数量 + """ + + result = await db.execute( + update(SysOAuthRefreshToken) + .where(SysOAuthRefreshToken.user_id == user_id, SysOAuthRefreshToken.status.not_in(['revoked', 'expired'])) + .values(status='revoked', revoked_at=TimezoneUtil.utc_now(), revoke_reason=reason) + ) + + return result.rowcount or 0 + + @classmethod + async def revoke_for_users( + cls, db: AsyncSession, user_ids: Sequence[int], reason: str | None = None, now: datetime | None = None + ) -> int: + """ + 批量撤销用户的 OAuth Refresh Token + + :param db: orm对象 + :param user_ids: 用户编号序列 + :param reason: 撤销或标记原因 + :param now: 当前时间 + :return: 已撤销的 OAuth Refresh Token 数量 + """ + + result = await db.execute( + update(SysOAuthRefreshToken) + .where( + SysOAuthRefreshToken.user_id.in_(user_ids), + SysOAuthRefreshToken.status.not_in(['revoked', 'expired', 'reuse_detected']), + ) + .values(status='revoked', revoked_at=now or TimezoneUtil.utc_now(), revoke_reason=reason) + ) + await db.flush() + + return result.rowcount or 0 + + @classmethod + async def expire_due(cls, db: AsyncSession) -> int: + """ + 处理已到期的 OAuth Refresh Token + + :param db: orm对象 + :return: 已过期的 OAuth Refresh Token 数量 + """ + + result = await db.execute( + update(SysOAuthRefreshToken) + .where( + SysOAuthRefreshToken.status == 'active', + (SysOAuthRefreshToken.idle_expires_at <= TimezoneUtil.utc_now()) + | (SysOAuthRefreshToken.absolute_expires_at <= TimezoneUtil.utc_now()), + ) + .values(status='expired') + ) + + return result.rowcount or 0 diff --git a/ruoyi-fastapi-backend/module_identity/dao/oidc_key_dao.py b/ruoyi-fastapi-backend/module_identity/dao/oidc_key_dao.py new file mode 100644 index 000000000..c10dfe690 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/oidc_key_dao.py @@ -0,0 +1,314 @@ +from collections.abc import Sequence +from datetime import datetime + +from sqlalchemy import delete, func, select, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.oidc_key_do import SysOidcSigningKey +from utils.time_util import TimezoneUtil + + +class OidcKeyDao: + """ + OIDC Key 数据库操作层 + """ + + @classmethod + async def get_active( + cls, db: AsyncSession, alg: str = 'RS256', for_update: bool = False + ) -> SysOidcSigningKey | None: + """ + 按签名算法查询活跃 OIDC Signing Key + + :param db: orm对象 + :param alg: 签名算法 + :param for_update: 是否锁定查询结果 + :return: 活跃 OIDC Signing Key,不存在时返回 None + """ + + query = select(SysOidcSigningKey).where( + SysOidcSigningKey.status == 'active', + SysOidcSigningKey.alg == alg, + SysOidcSigningKey.signing_start_at <= TimezoneUtil.utc_now(), + ) + if for_update: + query = query.with_for_update() + result = await db.execute(query.order_by(SysOidcSigningKey.signing_start_at.desc())) + + return result.scalars().first() + + @classmethod + async def list_published(cls, db: AsyncSession, now: datetime | None = None) -> Sequence[SysOidcSigningKey]: + """ + 查询已发布的 OIDC Signing Key + + :param db: orm对象 + :param now: 当前时间 + :return: 已发布的 OIDC Signing Key 序列 + """ + + current = now or TimezoneUtil.utc_now() + result = await db.execute( + select(SysOidcSigningKey) + .where( + SysOidcSigningKey.status.in_(['pending', 'active', 'retiring']), + SysOidcSigningKey.publish_at <= current, + (SysOidcSigningKey.remove_from_jwks_at.is_(None) | (SysOidcSigningKey.remove_from_jwks_at > current)), + ) + .order_by(SysOidcSigningKey.create_time) + ) + + return result.scalars().all() + + @classmethod + async def get_verifying(cls, db: AsyncSession, kid: str, now: datetime) -> SysOidcSigningKey | None: + """ + 按 kid 查询可用于验签的 OIDC Signing Key + + :param db: orm对象 + :param kid: Signing Key 标识 + :param now: 当前时间 + :return: 可验签的 OIDC Signing Key,不存在时返回 None + """ + + result = await db.execute( + select(SysOidcSigningKey).where( + SysOidcSigningKey.kid == kid, + SysOidcSigningKey.alg == 'RS256', + SysOidcSigningKey.status.in_(['active', 'retiring']), + SysOidcSigningKey.publish_at <= now, + (SysOidcSigningKey.remove_from_jwks_at.is_(None) | (SysOidcSigningKey.remove_from_jwks_at > now)), + ) + ) + + return result.scalars().first() + + @classmethod + async def list_admin( + cls, db: AsyncSession, *, status: str | None = None, offset: int = 0, limit: int = 200 + ) -> Sequence[SysOidcSigningKey]: + """ + 分页查询管理端 OIDC Signing Key + + :param db: orm对象 + :param status: 状态过滤值 + :param offset: 分页偏移量 + :param limit: 分页大小 + :return: 管理端 OIDC Signing Key 序列 + """ + + query = select(SysOidcSigningKey) + if status: + query = query.where(SysOidcSigningKey.status == status) + result = await db.execute( + query.order_by(SysOidcSigningKey.create_time.desc()).offset(max(offset, 0)).limit(min(max(limit, 1), 200)) + ) + + return result.scalars().all() + + @classmethod + async def list_due_pending(cls, db: AsyncSession, now: datetime | None = None) -> Sequence[SysOidcSigningKey]: + """ + 查询待发布时间已到的 OIDC Signing Key + + :param db: orm对象 + :param now: 当前时间 + :return: 待激活 OIDC Signing Key 序列 + """ + + current = now or TimezoneUtil.utc_now() + result = await db.execute( + select(SysOidcSigningKey) + .where( + SysOidcSigningKey.alg == 'RS256', + SysOidcSigningKey.status == 'pending', + SysOidcSigningKey.publish_at <= current, + SysOidcSigningKey.signing_start_at <= current, + ) + .order_by(SysOidcSigningKey.signing_start_at, SysOidcSigningKey.key_pk) + ) + + return result.scalars().all() + + @classmethod + async def count_admin(cls, db: AsyncSession, *, status: str | None = None) -> int: + """ + 统计管理端 OIDC Signing Key 数量 + + :param db: orm对象 + :param status: 状态过滤值 + :return: 管理端 OIDC Signing Key 数量 + """ + + query = select(func.count()).select_from(SysOidcSigningKey) + if status: + query = query.where(SysOidcSigningKey.status == status) + result = await db.execute(query) + + return int(result.scalar_one()) + + @classmethod + async def get_by_kid_for_update(cls, db: AsyncSession, kid: str) -> SysOidcSigningKey | None: + """ + 按 kid 锁定查询 OIDC Signing Key + + :param db: orm对象 + :param kid: Signing Key 标识 + :return: OIDC Signing Key,不存在时返回 None + """ + + result = await db.execute(select(SysOidcSigningKey).where(SysOidcSigningKey.kid == kid).with_for_update()) + + return result.scalars().first() + + @classmethod + async def create(cls, db: AsyncSession, record: SysOidcSigningKey) -> SysOidcSigningKey: + """ + 新增 OIDC Signing Key 并刷新主键 + + :param db: orm对象 + :param record: OIDC Signing Key 对象 + :return: 已写入的 OIDC Signing Key + """ + + db.add(record) + await db.flush() + + return record + + @classmethod + async def set_retiring( + cls, + db: AsyncSession, + kid: str, + remove_from_jwks_at: datetime, + now: datetime, + ) -> bool: + """ + 设置 OIDC Signing Key 退役信息 + + :param db: orm对象 + :param kid: Signing Key 标识 + :param remove_from_jwks_at: 从 JWKS 移除时间 + :param now: 当前时间 + :return: 是否更新成功 + """ + + result = await db.execute( + update(SysOidcSigningKey) + .where(SysOidcSigningKey.kid == kid, SysOidcSigningKey.status.in_(['active', 'retiring'])) + .values(status='retiring', signing_stop_at=now, remove_from_jwks_at=remove_from_jwks_at) + ) + + return bool(result.rowcount) + + @classmethod + async def delete_retired(cls, db: AsyncSession, kid: str) -> bool: + """ + 删除已退役的 OIDC Signing Key + + :param db: orm对象 + :param kid: Signing Key 标识 + :return: 是否删除成功 + """ + + result = await db.execute(delete(SysOidcSigningKey).where(SysOidcSigningKey.kid == kid)) + + return bool(result.rowcount) + + @classmethod + async def lock_algorithm_for_update(cls, db: AsyncSession, alg: str = 'RS256') -> Sequence[SysOidcSigningKey]: + """ + 按签名算法锁定查询 OIDC Signing Key + + :param db: orm对象 + :param alg: 签名算法 + :return: 同一算法的 OIDC Signing Key 序列 + """ + + result = await db.execute( + select(SysOidcSigningKey) + .where(SysOidcSigningKey.alg == alg) + .order_by(SysOidcSigningKey.key_pk) + .with_for_update() + ) + + return result.scalars().all() + + @classmethod + async def activate(cls, db: AsyncSession, kid: str, alg: str = 'RS256', now: datetime | None = None) -> bool: + """ + 激活 OIDC Signing Key + + :param db: orm对象 + :param kid: Signing Key 标识 + :param alg: 签名算法 + :param now: 当前时间 + :return: 是否激活成功 + """ + + current = now or TimezoneUtil.utc_now() + await cls.lock_algorithm_for_update(db, alg=alg) + target_result = await db.execute( + select(SysOidcSigningKey) + .where(SysOidcSigningKey.kid == kid, SysOidcSigningKey.alg == alg) + .with_for_update() + ) + target = target_result.scalars().first() + if target is None or target.status != 'pending': + return False + await db.execute( + update(SysOidcSigningKey) + .where(SysOidcSigningKey.alg == alg, SysOidcSigningKey.status == 'active') + .values(status='retiring', signing_stop_at=current) + ) + result = await db.execute( + update(SysOidcSigningKey) + .where( + SysOidcSigningKey.kid == kid, + SysOidcSigningKey.alg == alg, + SysOidcSigningKey.status == 'pending', + ) + .values(status='active', signing_start_at=current) + ) + + return bool(result.rowcount) + + @classmethod + async def mark_compromised(cls, db: AsyncSession, kid: str) -> bool: + """ + 标记 OIDC Signing Key 已泄露 + + :param db: orm对象 + :param kid: Signing Key 标识 + :return: 是否标记成功 + """ + + result = await db.execute( + update(SysOidcSigningKey) + .where(SysOidcSigningKey.kid == kid, SysOidcSigningKey.status != 'retired') + .values(status='compromised', signing_stop_at=TimezoneUtil.utc_now()) + ) + + return bool(result.rowcount) + + @classmethod + async def retire_due(cls, db: AsyncSession) -> int: + """ + 处理到期退役的 OIDC Signing Key + + :param db: orm对象 + :return: 已退役的 OIDC Signing Key 数量 + """ + + result = await db.execute( + update(SysOidcSigningKey) + .where( + SysOidcSigningKey.status == 'retiring', + SysOidcSigningKey.remove_from_jwks_at.is_not(None), + SysOidcSigningKey.remove_from_jwks_at <= TimezoneUtil.utc_now(), + ) + .values(status='retired') + ) + + return result.rowcount or 0 diff --git a/ruoyi-fastapi-backend/module_identity/dao/sso_session_dao.py b/ruoyi-fastapi-backend/module_identity/dao/sso_session_dao.py new file mode 100644 index 000000000..8482feb6a --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dao/sso_session_dao.py @@ -0,0 +1,514 @@ +from collections.abc import Sequence +from datetime import datetime + +from sqlalchemy import case, func, select, union, update +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthRefreshToken, SysSsoSession, SysSsoSessionClient +from utils.time_util import TimezoneUtil + + +class SsoSessionDao: + """ + SSO Session 数据库操作层 + """ + + @classmethod + async def create(cls, db: AsyncSession, row: SysSsoSession) -> SysSsoSession: + """ + 新增 SSO Session 并刷新主键 + + :param db: orm对象 + :param row: SSO Session 对象 + :return: 已写入的 SSO Session + """ + + db.add(row) + await db.flush() + + return row + + @classmethod + async def get_active( + cls, db: AsyncSession, sid: str, for_update: bool = False, now: datetime | None = None + ) -> SysSsoSession | None: + """ + 按 Session 标识查询活跃 SSO Session + + :param db: orm对象 + :param sid: SSO Session 标识 + :param for_update: 是否锁定查询结果 + :param now: 当前时间 + :return: 活跃 SSO Session,不存在时返回 None + """ + + now = now or TimezoneUtil.utc_now() + query = select(SysSsoSession).where( + SysSsoSession.sid == sid, + SysSsoSession.status == 'active', + SysSsoSession.idle_expires_at > now, + SysSsoSession.absolute_expires_at > now, + ) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def get_by_sid(cls, db: AsyncSession, sid: str, for_update: bool = False) -> SysSsoSession | None: + """ + 按 Session 标识查询 SSO Session + + :param db: orm对象 + :param sid: SSO Session 标识 + :param for_update: 是否锁定查询结果 + :return: SSO Session,不存在时返回 None + """ + + query = select(SysSsoSession).where(SysSsoSession.sid == sid) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().first() + + @classmethod + async def get_for_token( + cls, db: AsyncSession, sid: str, *, now: datetime, allow_offline: bool = False + ) -> SysSsoSession | None: + """ + 获取令牌校验所需的 SSO 会话 + + 自然过期的会话仍可用于离线授权,显式撤销的会话不可继续使用。 + + :param db: orm对象 + :param sid: SSO会话标识 + :param now: 当前UTC时间 + :param allow_offline: 是否允许使用自然过期的离线授权会话 + :return: 可用于令牌校验的会话,校验失败时返回None + """ + + session = await cls.get_active(db, sid, now=now) + if session is not None or not allow_offline: + return session + session = await cls.get_by_sid(db, sid) + if ( + session is not None + and session.status in {'active', 'expired'} + and getattr(session, 'revoked_at', None) is None + ): + return session + return None + + @classmethod + async def touch(cls, db: AsyncSession, sid: str, idle_expires_at: datetime, now: datetime | None = None) -> bool: + """ + 更新 SSO Session 闲置过期时间 + + :param db: orm对象 + :param sid: SSO Session 标识 + :param idle_expires_at: 闲置过期时间 + :param now: 当前时间 + :return: 是否更新成功 + """ + + current = now or TimezoneUtil.utc_now() + result = await db.execute( + update(SysSsoSession) + .where( + SysSsoSession.sid == sid, + SysSsoSession.status == 'active', + SysSsoSession.idle_expires_at > current, + SysSsoSession.absolute_expires_at > current, + idle_expires_at <= SysSsoSession.absolute_expires_at, + ) + .values(last_seen_at=current, idle_expires_at=idle_expires_at) + ) + await db.flush() + + return bool(result.rowcount) + + @classmethod + async def revoke(cls, db: AsyncSession, sid: str, reason: str | None = None, now: datetime | None = None) -> bool: + """ + 撤销 SSO Session + + :param db: orm对象 + :param sid: SSO Session 标识 + :param reason: 撤销或标记原因 + :param now: 当前时间 + :return: 是否撤销成功 + """ + + result = await db.execute( + update(SysSsoSession) + .where(SysSsoSession.sid == sid, SysSsoSession.status.in_(('active', 'expired'))) + .values(status='revoked', revoked_at=now or TimezoneUtil.utc_now(), revoke_reason=reason) + ) + await db.flush() + + return bool(result.rowcount) + + @classmethod + async def revoke_for_user( + cls, db: AsyncSession, user_id: int, reason: str | None = None, now: datetime | None = None + ) -> int: + """ + 撤销用户的 SSO Session + + :param db: orm对象 + :param user_id: 用户编号 + :param reason: 撤销或标记原因 + :param now: 当前时间 + :return: 已撤销的 SSO Session 数量 + """ + + result = await db.execute( + update(SysSsoSession) + .where(SysSsoSession.user_id == user_id, SysSsoSession.status.in_(('active', 'expired'))) + .values(status='revoked', revoked_at=now or TimezoneUtil.utc_now(), revoke_reason=reason) + ) + await db.flush() + + return result.rowcount or 0 + + @classmethod + async def revoke_for_users( + cls, + db: AsyncSession, + user_ids: Sequence[int], + reason: str | None = None, + now: datetime | None = None, + exclude_sid: str | None = None, + ) -> int: + """ + 批量撤销用户的 SSO Session + + :param db: orm对象 + :param user_ids: 用户编号序列 + :param reason: 撤销或标记原因 + :param now: 当前时间 + :param exclude_sid: 排除的 SSO Session 标识 + :return: 已撤销的 SSO Session 数量 + """ + + conditions = [SysSsoSession.user_id.in_(user_ids), SysSsoSession.status.in_(('active', 'expired'))] + if exclude_sid is not None: + conditions.append(SysSsoSession.sid != exclude_sid) + result = await db.execute( + update(SysSsoSession) + .where(*conditions) + .values(status='revoked', revoked_at=now or TimezoneUtil.utc_now(), revoke_reason=reason) + ) + await db.flush() + + return result.rowcount or 0 + + @classmethod + async def expire_due(cls, db: AsyncSession, now: datetime | None = None) -> int: + """ + 处理已到期的 SSO Session + + :param db: orm对象 + :param now: 当前时间 + :return: 已过期的 SSO Session 数量 + """ + + now = now or TimezoneUtil.utc_now() + result = await db.execute( + update(SysSsoSession) + .where( + SysSsoSession.status == 'active', + (SysSsoSession.idle_expires_at <= now) | (SysSsoSession.absolute_expires_at <= now), + ) + .values(status='expired') + ) + await db.flush() + + return result.rowcount or 0 + + @classmethod + async def list_due( + cls, db: AsyncSession, now: datetime | None = None, for_update: bool = False + ) -> Sequence[SysSsoSession]: + """ + 查询待过期的 SSO Session + + :param db: orm对象 + :param now: 当前时间 + :param for_update: 是否锁定查询结果 + :return: 待过期的 SSO Session 序列 + """ + + current = now or TimezoneUtil.utc_now() + query = select(SysSsoSession).where( + SysSsoSession.status == 'active', + (SysSsoSession.idle_expires_at <= current) | (SysSsoSession.absolute_expires_at <= current), + ) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().all() + + @classmethod + async def expire(cls, db: AsyncSession, sid: str, now: datetime | None = None) -> bool: + """ + 使 SSO Session 过期 + + :param db: orm对象 + :param sid: SSO Session 标识 + :param now: 当前时间 + :return: 是否过期成功 + """ + + current = now or TimezoneUtil.utc_now() + result = await db.execute( + update(SysSsoSession) + .where( + SysSsoSession.sid == sid, + SysSsoSession.status == 'active', + (SysSsoSession.idle_expires_at <= current) | (SysSsoSession.absolute_expires_at <= current), + ) + .values(status='expired') + ) + await db.flush() + + return bool(result.rowcount) + + @classmethod + async def rotate_secret( + cls, + db: AsyncSession, + sid: str, + old_secret_hash: str, + new_secret_hash: str, + now: datetime | None = None, + ) -> bool: + """ + 轮换 SSO Session 密钥摘要 + + :param db: orm对象 + :param sid: SSO Session 标识 + :param old_secret_hash: 旧 Session Secret 摘要 + :param new_secret_hash: 新 Session Secret 摘要 + :param now: 当前时间 + :return: 是否轮换成功 + """ + + current = now or TimezoneUtil.utc_now() + result = await db.execute( + update(SysSsoSession) + .where( + SysSsoSession.sid == sid, + SysSsoSession.status == 'active', + SysSsoSession.session_secret_hash == old_secret_hash, + SysSsoSession.idle_expires_at > current, + SysSsoSession.absolute_expires_at > current, + ) + .values(session_secret_hash=new_secret_hash, last_seen_at=current) + ) + await db.flush() + + return bool(result.rowcount) + + @classmethod + async def list_for_user( + cls, + db: AsyncSession, + user_id: int, + active_only: bool = False, + for_update: bool = False, + *, + revocable_only: bool = False, + ) -> Sequence[SysSsoSession]: + """ + 查询用户 SSO Session 列表 + + :param db: orm对象 + :param user_id: 用户编号 + :param active_only: 是否仅查询启用记录 + :param for_update: 是否锁定查询结果 + :param revocable_only: 是否仅查询可撤销的在线或自然过期会话 + :return: 用户 SSO Session 序列 + """ + + query = select(SysSsoSession).where(SysSsoSession.user_id == user_id) + if active_only: + query = query.where(SysSsoSession.status == 'active') + elif revocable_only: + query = query.where(SysSsoSession.status.in_(('active', 'expired'))) + query = query.order_by(SysSsoSession.create_time.desc()) + if for_update: + query = query.with_for_update() + result = await db.execute(query) + + return result.scalars().all() + + @classmethod + async def list_page( + cls, + db: AsyncSession, + *, + user_id: int | None = None, + ip_address: str | None = None, + status: str | None = None, + start_time: datetime | None = None, + end_time: datetime | None = None, + offset: int = 0, + limit: int = 200, + ) -> Sequence[SysSsoSession]: + """ + 分页查询 SSO Session + + :param db: orm对象 + :param user_id: 用户编号 + :param ip_address: 客户端 IP 地址 + :param status: 状态过滤值 + :param start_time: 查询开始时间 + :param end_time: 查询结束时间 + :param offset: 分页偏移量 + :param limit: 分页大小 + :return: SSO Session 序列 + """ + + conditions = [] + if user_id is not None: + conditions.append(SysSsoSession.user_id == user_id) + if ip_address: + conditions.append(SysSsoSession.ip_address == ip_address) + if status: + current = TimezoneUtil.utc_now() + effective_status = case( + ( + (SysSsoSession.status == 'active') + & ((SysSsoSession.idle_expires_at <= current) | (SysSsoSession.absolute_expires_at <= current)), + 'expired', + ), + else_=SysSsoSession.status, + ) + conditions.append(effective_status == status) + if start_time is not None: + conditions.append(SysSsoSession.create_time >= start_time) + if end_time is not None: + conditions.append(SysSsoSession.create_time <= end_time) + query = ( + select(SysSsoSession) + .where(*conditions) + .order_by(SysSsoSession.create_time.desc()) + .offset(max(offset, 0)) + .limit(min(max(limit, 1), 200)) + ) + result = await db.execute(query) + + return result.scalars().all() + + @classmethod + async def count( + cls, + db: AsyncSession, + *, + user_id: int | None = None, + ip_address: str | None = None, + status: str | None = None, + start_time: datetime | None = None, + end_time: datetime | None = None, + ) -> int: + """ + 统计 SSO Session 数量 + + :param db: orm对象 + :param user_id: 用户编号 + :param ip_address: 客户端 IP 地址 + :param status: 状态过滤值 + :param start_time: 查询开始时间 + :param end_time: 查询结束时间 + :return: SSO Session 数量 + """ + + conditions = [] + if user_id is not None: + conditions.append(SysSsoSession.user_id == user_id) + if ip_address: + conditions.append(SysSsoSession.ip_address == ip_address) + if status: + current = TimezoneUtil.utc_now() + effective_status = case( + ( + (SysSsoSession.status == 'active') + & ((SysSsoSession.idle_expires_at <= current) | (SysSsoSession.absolute_expires_at <= current)), + 'expired', + ), + else_=SysSsoSession.status, + ) + conditions.append(effective_status == status) + if start_time is not None: + conditions.append(SysSsoSession.create_time >= start_time) + if end_time is not None: + conditions.append(SysSsoSession.create_time <= end_time) + result = await db.execute(select(func.count()).select_from(SysSsoSession).where(*conditions)) + + return int(result.scalar_one()) + + @classmethod + async def record_client(cls, db: AsyncSession, sid: str, client_pk: int, now: datetime) -> None: + """ + 记录成功兑换授权码的会话与应用;调用方须先锁定 Client + + :param db: orm对象 + :param sid: SSO Session 标识 + :param client_pk: Client 内部主键 + :param now: 本次授权时间 + :return: None + """ + + row = await db.scalar( + select(SysSsoSessionClient).where( + SysSsoSessionClient.sid == sid, SysSsoSessionClient.client_pk == client_pk + ) + ) + if row is None: + db.add(SysSsoSessionClient(sid=sid, client_pk=client_pk, create_time=now, last_used_at=now)) + else: + row.last_used_at = now + await db.flush() + + @classmethod + async def client_ids_for_sid(cls, db: AsyncSession, sid: str) -> Sequence[str]: + """ + 查询精确绑定该会话的参与应用,兼容升级前已有的离线凭据 + + :param db: orm对象 + :param sid: SSO Session 标识 + :return: 去重后的 Client 公开标识序列 + """ + + return (await cls.client_ids_for_sids(db, [sid])).get(sid, []) + + @classmethod + async def client_ids_for_sids(cls, db: AsyncSession, sids: list[str]) -> dict[str, list[str]]: + """ + 批量查询会话关联应用,避免管理列表逐行查询 + + :param db: orm对象 + :param sids: SSO Session 标识列表 + :return: 会话标识与去重客户端标识列表的映射 + """ + + if not sids: + return {} + participants = union( + select(SysSsoSessionClient.sid, SysSsoSessionClient.client_pk).where(SysSsoSessionClient.sid.in_(sids)), + select(SysOAuthRefreshToken.sid, SysOAuthRefreshToken.client_pk).where(SysOAuthRefreshToken.sid.in_(sids)), + ).subquery() + result = await db.execute( + select(participants.c.sid, SysOAuthClient.client_id) + .join(SysOAuthClient, SysOAuthClient.client_pk == participants.c.client_pk) + .order_by(participants.c.sid, SysOAuthClient.client_id) + ) + mapping: dict[str, list[str]] = {} + for sid, client_id in result.all(): + mapping.setdefault(sid, []).append(client_id) + return mapping diff --git a/ruoyi-fastapi-backend/module_identity/dependencies.py b/ruoyi-fastapi-backend/module_identity/dependencies.py new file mode 100644 index 000000000..0e2dca628 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/dependencies.py @@ -0,0 +1,330 @@ +import json +import re +from dataclasses import dataclass +from datetime import datetime +from typing import Any +from urllib.parse import parse_qsl + +import jwt +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey +from fastapi import Depends, HTTPException, Request, params +from jwt.exceptions import PyJWTError +from redis.asyncio import Redis +from redis.exceptions import RedisError +from sqlalchemy.ext.asyncio import AsyncSession + +from common.aspect.db_session import DBSessionDependency +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException, OidcInteractionException +from module_identity.dao.oidc_key_dao import OidcKeyDao +from module_identity.security.client_auth import ClientAuthenticationError +from module_identity.security.jwt_profile import JwtProfileError, decode_access_token +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.interaction_service import InteractionService +from module_identity.service.runtime_service import OidcRuntimeService +from module_identity.service.token_service import TokenService +from utils.time_util import TimezoneUtil + +_MAX_PROTOCOL_FORM_BYTES = 16 * 1024 +_INVALID_PERCENT_ESCAPE = re.compile(r'%(?![0-9A-Fa-f]{2})') + + +@dataclass(frozen=True, slots=True) +class AccessTokenContext: + """ + 已验签且绑定 UserInfo audience 的 Access Token 快照 + """ + + token: str + claims: dict[str, Any] + + +_JWT_DOT_COUNT = 2 +_REMOTE_KEY_HEADERS = frozenset({'jku', 'x5u', 'jwk', 'x5c'}) + + +async def load_access_verification_key( + token: str, + query_db: AsyncSession, + *, + now: datetime | None = None, +) -> RSAPublicKey: + """ + 按 Access Token header kid 加载发布窗口内的本地公钥 + + :param token: 待验证的 JWT;仅用于读取 header kid,不信任其中的公钥材料 + :param query_db: 异步数据库会话 + :param now: 可注入的统一项目当前时间 + :return: 数据库记录对应的 RSA 公钥对象 + :raises JwtProfileError: Header、kid 或本地公钥不可用 + """ + + try: + header = jwt.get_unverified_header(token) + except (PyJWTError, TypeError, ValueError) as exc: + raise JwtProfileError('JWT 头部格式无效') from exc + if _REMOTE_KEY_HEADERS.intersection(header) or header.get('alg') != 'RS256' or header.get('typ') != 'at+jwt': + raise JwtProfileError('JWT 头部包含不支持的参数') + kid = header.get('kid') + if not isinstance(kid, str) or not kid: + raise JwtProfileError('签名密钥标识 kid 不能为空') + current = TimezoneUtil.to_utc(now or TimezoneUtil.utc_now()) + record = await OidcKeyDao.get_verifying(query_db, kid, current) + if record is None or not isinstance(record.public_jwk, dict): + raise JwtProfileError('未知的签名密钥标识 kid') + try: + return jwt.algorithms.RSAAlgorithm.from_jwk(json.dumps(record.public_jwk, separators=(',', ':'))) + except (TypeError, ValueError, KeyError) as exc: + raise JwtProfileError('签名公钥无效') from exc + + +def _disabled() -> None: + """ + 在 OIDC 关闭时返回本地 404,不执行协议认证 + """ + + if not OidcConfig.oidc_enabled: + raise HTTPException( + status_code=404, + detail='请求的资源不存在', + headers={'Cache-Control': 'no-store', 'Pragma': 'no-cache'}, + ) + + +async def require_oidc_protocol_ready( + request: Request, + query_db: AsyncSession = DBSessionDependency(), +) -> None: + """ + 在 OIDC 已启用但签名能力尚未就绪时阻断公共协议端点 + + OIDC 关闭时仍由各控制器维持原有本地 404 行为;启用但缺少有效 + active 密钥时统一返回标准 OAuth 503,不影响后台管理接口。 + + :param request: 当前 HTTP 请求 + :param query_db: 异步数据库会话 + :return: None + :raises OAuthProtocolException: OIDC 协议尚未具备签名能力 + """ + + if not OidcConfig.oidc_enabled: + return + readiness = await OidcRuntimeService.cached_readiness(request.app, query_db) + if not readiness.ready: + raise OAuthProtocolException( + 'temporarily_unavailable', + 'Authorization server is not ready', + 503, + ) + + +def _redis(request: Request) -> Redis: + """ + 读取应用生命周期创建的 Redis 客户端 + + :param request: 当前HTTP请求 + :return: 应用生命周期创建的Redis客户端 + """ + + value = getattr(request.app.state, 'redis', None) + if value is None: + raise HTTPException(status_code=503, detail='认证服务暂不可用') + return value + + +async def read_form(request: Request) -> dict[str, str]: + """ + 严格读取 URL encoded 表单并拒绝重复字段 + + :param request: 当前 HTTP 请求 + :return: 唯一字段映射 + :raises HTTPException: Content-Type 错误或字段重复 + """ + + _disabled() + content_type = request.headers.get('content-type', '').split(';', 1)[0].strip().lower() + if content_type != 'application/x-www-form-urlencoded': + raise HTTPException(status_code=415, detail='不支持当前请求内容类型') + content_length = request.headers.get('content-length') + declared_length: int | None = None + if content_length is not None: + try: + declared_length = int(content_length) + except ValueError as exc: + raise HTTPException(status_code=413, detail='请求体大小超过限制') from exc + if declared_length < 0 or declared_length > _MAX_PROTOCOL_FORM_BYTES: + raise HTTPException(status_code=413, detail='请求体大小超过限制') + chunks: list[bytes] = [] + actual_length = 0 + async for chunk in request.stream(): + actual_length += len(chunk) + if actual_length > _MAX_PROTOCOL_FORM_BYTES: + raise HTTPException(status_code=413, detail='请求体大小超过限制') + chunks.append(chunk) + if declared_length is not None and actual_length != declared_length: + raise HTTPException(status_code=413, detail='请求体大小超过限制') + try: + decoded_body = b''.join(chunks).decode('utf-8') + if _INVALID_PERCENT_ESCAPE.search(decoded_body): + raise ValueError('百分号编码无效') + form_items = parse_qsl( + decoded_body, + keep_blank_values=True, + strict_parsing=True, + encoding='utf-8', + errors='strict', + max_num_fields=64, + ) + except (UnicodeDecodeError, ValueError) as exc: + raise HTTPException(status_code=400, detail='表单格式或参数无效') from exc + result: dict[str, str] = {} + for key, value in form_items: + if key in result: + raise HTTPException(status_code=400, detail='表单格式或参数无效') + result[key] = value + return result + + +async def get_oidc_client( + request: Request, + query_db: AsyncSession = DBSessionDependency(), +) -> OAuthClientPrincipal: + """ + 验证 Basic 或 Public Client 表单认证 + + :param request: 当前 HTTP 请求 + :param query_db: 异步数据库会话 + :return: 已完成认证的 Client Principal + """ + + form = await read_form(request) + authorization = request.headers.get('authorization') + try: + _, principal = await TokenService.authenticate_client( + query_db, + authorization=authorization, + client_id=form.get('client_id'), + client_secret=form.get('client_secret'), + ) + except (ClientAuthenticationError, ValueError, OAuthProtocolException): + raise HTTPException( + status_code=401, + detail='客户端认证失败', + headers={'WWW-Authenticate': 'Basic realm="oauth2/token"'}, + ) from None + return principal + + +async def _load_access_context(request: Request, query_db: AsyncSession) -> AccessTokenContext: + """ + 按 JWT header kid 查询数据库公钥并验证严格 Access Token Profile + + :param request: 当前HTTP请求 + :param query_db: orm对象 + :return: 已校验的访问令牌及声明上下文 + """ + + _disabled() + authorization = request.headers.get('authorization', '') + if authorization[:7].lower() != 'bearer ' or ',' in authorization: + raise _invalid_token() + token = authorization[7:].strip() + if not token or token.count('.') != _JWT_DOT_COUNT: + raise _invalid_token() + try: + key = await load_access_verification_key(token, query_db) + issuer = OidcConfig.oidc_issuer.rstrip('/') + claims = decode_access_token( + token, + verification_key=key, + issuer=issuer, + audience=f'{issuer}/oauth2/userinfo', + clock_skew=OidcConfig.oidc_allowed_clock_skew_seconds, + ) + scope = claims.get('scope', '').split() + if ( + 'openid' not in scope + or not isinstance(claims.get('sid'), str) + or claims.get('sub', '').startswith('client:') + ): + raise JwtProfileError('当前令牌不是用户访问令牌') + if isinstance(claims.get('ver'), bool) or not isinstance(claims.get('ver'), int) or claims['ver'] < 1: + raise JwtProfileError('身份安全版本无效') + return AccessTokenContext(token, claims) + except (JwtProfileError, PyJWTError, TypeError, ValueError, KeyError): + raise _invalid_token() from None + + +async def get_oidc_access_token( + request: Request, + query_db: AsyncSession = DBSessionDependency(), +) -> AccessTokenContext: + """ + 验证 UserInfo 使用的 Bearer Access Token + + :param request: 当前HTTP请求 + :param query_db: orm对象 + :return: UserInfo使用的访问令牌上下文 + """ + + return await _load_access_context(request, query_db) + + +async def get_interaction_csrf( + request: Request, + interaction_id: str, +) -> dict[str, Any]: + """ + 验证 Interaction CSRF Header,并返回内部状态记录 + + :param request: 当前 HTTP 请求 + :param interaction_id: 路径中的 Interaction 标识 + :return: 已验证的内部 Interaction 记录 + """ + + _disabled() + redis = _redis(request) + csrf_token = request.headers.get('x-csrf-token') + if not csrf_token or ',' in csrf_token: + raise HTTPException(status_code=403, detail='CSRF 校验失败,请重新发起认证') + try: + record = await InteractionService.get_record(redis, interaction_id) + except (OidcInteractionException, RedisError, TypeError, ValueError): + raise HTTPException(status_code=404, detail='认证交互不存在') from None + if not InteractionService.verify_csrf(record, csrf_token, pepper=OidcConfig.oidc_token_hash_pepper): + raise HTTPException(status_code=403, detail='CSRF 校验失败,请重新发起认证') + return record + + +def _invalid_token() -> HTTPException: + """ + 构造不泄露验证细节的 Bearer 401 + """ + + return HTTPException( + status_code=401, detail='访问令牌无效', headers={'WWW-Authenticate': 'Bearer error="invalid_token"'} + ) + + +def OidcClientDependency() -> params.Depends: # noqa: N802 + """ + 返回 Client Authentication 依赖 + """ + + return Depends(get_oidc_client) + + +def OidcAccessTokenDependency() -> params.Depends: # noqa: N802 + """ + 返回严格 Access Token 依赖 + """ + + return Depends(get_oidc_access_token) + + +def InteractionCsrfDependency() -> params.Depends: # noqa: N802 + """ + 返回 Interaction CSRF 依赖 + """ + + return Depends(get_interaction_csrf) diff --git a/ruoyi-fastapi-backend/module_identity/entity/do/identity_subject_do.py b/ruoyi-fastapi-backend/module_identity/entity/do/identity_subject_do.py new file mode 100644 index 000000000..c22edcc00 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/do/identity_subject_do.py @@ -0,0 +1,39 @@ +from sqlalchemy import BigInteger, Column, ForeignKey, Index, Integer, String, UniqueConstraint + +from common.types import DbUtcDateTime +from config.database import Base +from utils.time_util import TimezoneUtil + + +class SysIdentitySubject(Base): + """ + 统一认证主体关联表 + """ + + __tablename__ = 'sys_identity_subject' + __table_args__ = ( + UniqueConstraint('user_id', name='uk_identity_subject_user'), + UniqueConstraint('subject_id', name='uk_identity_subject_subject'), + Index('idx_identity_subject_auth_version', 'auth_version'), + {'comment': '统一认证主体关联表'}, + ) + + identity_id = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='内部主键', + ) + user_id = Column( + BigInteger, + ForeignKey('sys_user.user_id', name='fk_identity_subject_user', ondelete='RESTRICT'), + nullable=False, + comment='本地用户ID', + ) + subject_id = Column(String(36), nullable=False, comment='OIDC Subject') + auth_version = Column(BigInteger, nullable=False, server_default='1', comment='认证安全版本') + create_by = Column(String(64), nullable=True, comment='创建者') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + update_by = Column(String(64), nullable=True, comment='更新者') + update_time = Column(DbUtcDateTime(), nullable=True, onupdate=TimezoneUtil.utc_now, comment='更新时间') diff --git a/ruoyi-fastapi-backend/module_identity/entity/do/oauth_audit_do.py b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_audit_do.py new file mode 100644 index 000000000..f654a5b45 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_audit_do.py @@ -0,0 +1,77 @@ +from sqlalchemy import JSON, BigInteger, Column, Index, Integer, String + +from common.types import DbUtcDateTime +from config.database import Base +from utils.time_util import TimezoneUtil + + +class SysOAuthAuditLog(Base): + """ + OAuth/OIDC安全审计日志表 + """ + + __tablename__ = 'sys_oauth_audit_log' + __table_args__ = ( + Index('idx_oauth_audit_time', 'create_time'), + Index('idx_oauth_audit_client', 'client_id', 'create_time'), + Index('idx_oauth_audit_user', 'user_id', 'create_time'), + Index('idx_oauth_audit_event', 'event_type', 'result', 'create_time'), + Index('idx_oauth_audit_risk', 'risk_level', 'create_time'), + {'comment': 'OAuth Audit Log'}, + ) + + event_id = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='事件ID', + ) + trace_id = Column(String(64), nullable=True, comment='链路追踪ID') + event_type = Column(String(64), nullable=False, comment='事件类型') + result = Column(String(16), nullable=False, comment='结果') + risk_level = Column(String(16), nullable=False, server_default='normal', comment='风险等级') + client_id = Column(String(64), nullable=True, comment='Client ID 快照') + resource_id = Column(String(64), nullable=True, comment='Resource ID 快照') + user_id = Column(BigInteger, nullable=True, comment='用户ID快照') + subject_id = Column(String(36), nullable=True, comment='Subject 快照') + sid = Column(String(36), nullable=True, comment='SSO Session ID') + grant_id = Column(String(36), nullable=True, comment='Grant ID') + token_id = Column(String(36), nullable=True, comment='Token ID') + ip_address = Column(String(128), nullable=True, comment='客户端 IP') + user_agent = Column(String(500), nullable=True, comment='脱敏 User-Agent') + failure_code = Column(String(64), nullable=True, comment='失败码') + detail = Column(JSON, nullable=True, comment='脱敏扩展详情') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='事件时间') + + +class SysOAuthAuditArchive(Base): + """ + OAuth/OIDC安全审计归档表 + """ + + __tablename__ = 'sys_oauth_audit_archive' + __table_args__ = ( + Index('idx_oauth_audit_archive_time', 'create_time'), + Index('idx_oauth_audit_archive_event', 'event_type', 'result', 'create_time'), + {'comment': 'OAuth Audit Archive'}, + ) + + event_id = Column(BigInteger, primary_key=True, nullable=False, autoincrement=False, comment='原事件ID') + trace_id = Column(String(64), nullable=True, comment='链路追踪ID') + event_type = Column(String(64), nullable=False, comment='事件类型') + result = Column(String(16), nullable=False, comment='结果') + risk_level = Column(String(16), nullable=False, comment='风险等级') + client_id = Column(String(64), nullable=True, comment='Client ID 快照') + resource_id = Column(String(64), nullable=True, comment='Resource ID 快照') + user_id = Column(BigInteger, nullable=True, comment='用户ID快照') + subject_id = Column(String(36), nullable=True, comment='Subject 快照') + sid = Column(String(36), nullable=True, comment='SSO Session ID') + grant_id = Column(String(36), nullable=True, comment='Grant ID') + token_id = Column(String(36), nullable=True, comment='Token ID') + ip_address = Column(String(128), nullable=True, comment='客户端 IP') + user_agent = Column(String(500), nullable=True, comment='脱敏 User-Agent') + failure_code = Column(String(64), nullable=True, comment='失败码') + detail = Column(JSON, nullable=True, comment='脱敏扩展详情') + create_time = Column(DbUtcDateTime(), nullable=False, comment='事件时间') + archived_at = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='归档时间') diff --git a/ruoyi-fastapi-backend/module_identity/entity/do/oauth_client_do.py b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_client_do.py new file mode 100644 index 000000000..8dc730c7d --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_client_do.py @@ -0,0 +1,145 @@ +from sqlalchemy import ( + CHAR, + JSON, + BigInteger, + Column, + ForeignKey, + Index, + Integer, + SmallInteger, + String, + UniqueConstraint, +) +from sqlalchemy.orm import validates + +from common.types import DbUtcDateTime +from config.database import Base +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + + +class SysOAuthClient(Base): + """ + OAuth客户端表 + """ + + __tablename__ = 'sys_oauth_client' + __table_args__ = ( + UniqueConstraint('client_id', name='uk_oauth_client_client_id'), + Index('idx_oauth_client_status', 'status'), + {'comment': 'OAuth Client'}, + ) + + client_pk = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='内部主键', + ) + client_id = Column(String(64), nullable=False, comment='Client ID') + client_name = Column(String(100), nullable=False, comment='客户端名称') + client_type = Column(String(20), nullable=False, comment='Client 类型') + token_endpoint_auth_method = Column(String(32), nullable=False, comment='Token 端点认证方式') + grant_types = Column(JSON, nullable=False, comment='Grant Type 列表') + response_types = Column(JSON, nullable=False, comment='Response Type 列表') + subject_type = Column(String(16), nullable=False, server_default='public', comment='Subject 类型') + require_pkce = Column(SmallInteger, nullable=False, server_default='1', comment='是否要求 PKCE') + require_consent = Column(SmallInteger, nullable=False, server_default='1', comment='是否要求同意') + trusted_client = Column(SmallInteger, nullable=False, server_default='0', comment='是否受信任 Client') + policy_version = Column(BigInteger, nullable=False, server_default='1', comment='安全策略版本') + id_token_signed_response_alg = Column(String(16), nullable=False, server_default='RS256', comment='ID Token 算法') + access_token_ttl_seconds = Column(Integer, nullable=True, comment='Access Token 有效期') + refresh_token_idle_seconds = Column(Integer, nullable=True, comment='Refresh Token 闲置有效期') + refresh_token_absolute_seconds = Column(Integer, nullable=True, comment='Refresh Token 绝对有效期') + logo_uri = Column(String(500), nullable=True, comment='Logo URI') + policy_uri = Column(String(500), nullable=True, comment='隐私政策 URI') + tos_uri = Column(String(500), nullable=True, comment='服务条款 URI') + backchannel_logout_session_required = Column( + SmallInteger, nullable=False, server_default='1', comment='是否要求 Back-Channel Session' + ) + status = Column(CHAR(1), nullable=False, server_default='0', comment='状态(0正常 1停用)') + create_by = Column(String(64), nullable=False, server_default='', comment='创建者') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + update_by = Column(String(64), nullable=False, server_default='', default='', comment='更新者') + update_time = Column( + DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, onupdate=TimezoneUtil.utc_now, comment='更新时间' + ) + remark = Column(String(500), nullable=True, comment='备注') + + +class SysOAuthClientSecret(Base): + """ + OAuth客户端密钥表,仅保存密钥强哈希 + """ + + __tablename__ = 'sys_oauth_client_secret' + __table_args__ = ( + Index('idx_oauth_client_secret_client', 'client_pk', 'status'), + {'comment': 'OAuth Client Secret'}, + ) + + secret_id = Column(String(36), primary_key=True, nullable=False, comment='Secret ID') + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_client_secret_client', ondelete='RESTRICT'), + nullable=False, + comment='Client 主键', + ) + secret_hash = Column(String(100), nullable=False, comment='Secret 强哈希') + secret_hint = Column(String(12), nullable=False, comment='Secret 提示') + status = Column(String(16), nullable=False, server_default='active', comment='Secret 状态') + not_before = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='生效时间') + expires_at = Column(DbUtcDateTime(), nullable=True, comment='过期时间') + last_used_at = Column(DbUtcDateTime(), nullable=True, comment='最近使用时间') + create_by = Column(String(64), nullable=False, comment='创建者') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + revoked_by = Column(String(64), nullable=True, comment='撤销者') + revoked_at = Column(DbUtcDateTime(), nullable=True, comment='撤销时间') + + +class SysOAuthClientUri(Base): + """ + OAuth客户端注册地址表 + """ + + __tablename__ = 'sys_oauth_client_uri' + __table_args__ = ( + UniqueConstraint('client_pk', 'uri_type', 'uri_hash', name='uk_oauth_client_uri_hash'), + Index('idx_oauth_client_uri_type', 'client_pk', 'uri_type', 'status'), + {'comment': 'OAuth Client URI'}, + ) + + uri_id = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='URI 主键', + ) + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_client_uri_client', ondelete='RESTRICT'), + nullable=False, + comment='Client 主键', + ) + uri_type = Column(String(32), nullable=False, comment='URI 类型') + uri = Column(String(1000), nullable=False, comment='精确 URI') + uri_hash = Column(CHAR(64), nullable=False, comment='URI SHA-256 摘要') + is_default = Column(SmallInteger, nullable=False, server_default='0', comment='是否默认 URI') + status = Column(CHAR(1), nullable=False, server_default='0', comment='状态(0正常 1停用)') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + + @validates('uri') + def derive_uri_hash(self, key: str, value: str) -> str: + """ + 根据 URI 更新摘要 + + :param key: 触发校验的属性名称 + :param value: 待写入的注册地址 + :return: 用于继续写入的原始注册地址 + """ + + self.uri_hash = OidcUtil.sha256_digest(value) + + return value diff --git a/ruoyi-fastapi-backend/module_identity/entity/do/oauth_grant_do.py b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_grant_do.py new file mode 100644 index 000000000..c2edc6ef4 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_grant_do.py @@ -0,0 +1,206 @@ +from sqlalchemy import CHAR, JSON, BigInteger, Column, ForeignKey, Index, SmallInteger, String, UniqueConstraint + +from common.types import DbUtcDateTime +from config.database import Base +from utils.time_util import TimezoneUtil + + +class SysOAuthAccessPolicy(Base): + """ + 用户与 OAuth Client 的访问控制表 + """ + + __tablename__ = 'sys_oauth_access_policy' + __table_args__ = ({'comment': 'OAuth 用户应用访问控制'},) + + user_id = Column( + BigInteger, + ForeignKey('sys_user.user_id', name='fk_oauth_access_user', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='用户ID', + ) + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_access_client', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='Client 主键', + ) + access_status = Column(String(16), nullable=False, server_default='allowed', comment='allowed允许 blocked禁止') + reason = Column(String(200), nullable=True, comment='访问控制原因') + update_by = Column(String(64), nullable=False, comment='操作人') + update_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='操作时间') + + +class SysOAuthGrant(Base): + """ + OAuth授权记录表 + """ + + __tablename__ = 'sys_oauth_grant' + __table_args__ = ( + Index('idx_oauth_grant_user', 'user_id', 'status'), + Index('idx_oauth_grant_client', 'client_pk', 'status'), + Index('idx_oauth_grant_user_client', 'user_id', 'client_pk', 'status'), + {'comment': 'OAuth Grant'}, + ) + + grant_id = Column(String(36), primary_key=True, nullable=False, comment='Grant ID') + user_id = Column( + BigInteger, + ForeignKey('sys_user.user_id', name='fk_oauth_grant_user', ondelete='RESTRICT'), + nullable=False, + comment='用户ID', + ) + subject_id = Column(String(36), nullable=False, comment='Subject 快照') + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_grant_client', ondelete='RESTRICT'), + nullable=False, + comment='Client 主键', + ) + granted_scopes = Column(JSON, nullable=False, comment='已同意 Scope') + granted_resources = Column(JSON, nullable=False, comment='已同意 Resource audience') + remembered_scopes = Column(JSON, nullable=True, comment='后续可免确认的 Scope') + remembered_resources = Column(JSON, nullable=True, comment='后续可免确认的 Resource audience') + client_policy_version = Column(BigInteger, nullable=False, comment='Client Policy Version') + status = Column(String(16), nullable=False, server_default='active', comment='Grant 状态') + consented_at = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='同意时间') + expires_at = Column(DbUtcDateTime(), nullable=True, comment='过期时间') + revoked_at = Column(DbUtcDateTime(), nullable=True, comment='撤销时间') + revoke_reason = Column(String(200), nullable=True, comment='撤销原因') + last_used_at = Column(DbUtcDateTime(), nullable=True, comment='最近使用时间') + + +class SysSsoSession(Base): + """ + OIDC单点登录会话表 + """ + + __tablename__ = 'sys_sso_session' + __table_args__ = ( + Index('idx_sso_session_user', 'user_id', 'status'), + Index('idx_sso_session_idle', 'status', 'idle_expires_at'), + Index('idx_sso_session_absolute', 'status', 'absolute_expires_at'), + {'comment': 'OIDC SSO Session'}, + ) + + sid = Column(String(36), primary_key=True, nullable=False, comment='OIDC Session ID') + session_secret_hash = Column(CHAR(64), nullable=False, comment='SSO Cookie 摘要') + user_id = Column( + BigInteger, + ForeignKey('sys_user.user_id', name='fk_sso_session_user', ondelete='RESTRICT'), + nullable=False, + comment='用户ID', + ) + subject_id = Column(String(36), nullable=False, comment='Subject 快照') + auth_version = Column(BigInteger, nullable=False, comment='认证安全版本') + auth_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='认证时间') + last_seen_at = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='最近活动时间') + idle_expires_at = Column(DbUtcDateTime(), nullable=False, comment='闲置过期时间') + absolute_expires_at = Column(DbUtcDateTime(), nullable=False, comment='绝对过期时间') + acr = Column(String(100), nullable=False, comment='认证上下文') + amr = Column(JSON, nullable=False, comment='认证方式') + remember_me = Column(SmallInteger, nullable=False, server_default='0', comment='是否长期会话') + ip_address = Column(String(128), nullable=True, comment='登录 IP') + user_agent_hash = Column(CHAR(64), nullable=True, comment='User-Agent 摘要') + status = Column(String(16), nullable=False, server_default='active', comment='Session 状态') + revoked_at = Column(DbUtcDateTime(), nullable=True, comment='撤销时间') + revoke_reason = Column(String(200), nullable=True, comment='撤销原因') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + + +class SysSsoSessionClient(Base): + """ + SSO 会话与参与应用关联表 + """ + + __tablename__ = 'sys_sso_session_client' + __table_args__ = (Index('idx_sso_session_client_client', 'client_pk'), {'comment': 'SSO 会话参与应用'}) + + sid = Column( + String(36), + ForeignKey('sys_sso_session.sid', name='fk_sso_session_client_sid', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='SSO Session ID', + ) + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_sso_session_client_client', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='Client 主键', + ) + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='首次授权时间') + last_used_at = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='最近授权时间') + + +class SysOAuthRefreshToken(Base): + """ + OAuth刷新令牌表,仅保存令牌摘要 + """ + + __tablename__ = 'sys_oauth_refresh_token' + __table_args__ = ( + UniqueConstraint('token_hash', name='uk_oauth_refresh_token_hash'), + Index('idx_oauth_refresh_family', 'family_id', 'status'), + Index('idx_oauth_refresh_user', 'user_id', 'status'), + Index('idx_oauth_refresh_client', 'client_pk', 'status'), + Index('idx_oauth_refresh_sid', 'sid', 'status'), + Index('idx_oauth_refresh_expire', 'status', 'absolute_expires_at'), + {'comment': 'OAuth Refresh Token'}, + ) + + token_id = Column(String(36), primary_key=True, nullable=False, comment='Token ID') + token_hash = Column(CHAR(64), nullable=False, comment='Token HMAC 摘要') + family_id = Column(String(36), nullable=False, comment='Token Family ID') + parent_token_id = Column( + String(36), + ForeignKey('sys_oauth_refresh_token.token_id', name='fk_oauth_refresh_parent', ondelete='RESTRICT'), + nullable=True, + comment='父 Token ID', + ) + replaced_by_token_id = Column( + String(36), + ForeignKey('sys_oauth_refresh_token.token_id', name='fk_oauth_refresh_replaced_by', ondelete='RESTRICT'), + nullable=True, + comment='替代 Token ID', + ) + grant_id = Column( + String(36), + ForeignKey('sys_oauth_grant.grant_id', name='fk_oauth_refresh_grant', ondelete='RESTRICT'), + nullable=False, + comment='Grant ID', + ) + user_id = Column( + BigInteger, + ForeignKey('sys_user.user_id', name='fk_oauth_refresh_user', ondelete='RESTRICT'), + nullable=False, + comment='用户ID', + ) + subject_id = Column(String(36), nullable=False, comment='Subject 快照') + auth_version = Column(BigInteger, nullable=False, comment='认证安全版本') + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_refresh_client', ondelete='RESTRICT'), + nullable=False, + comment='Client 主键', + ) + sid = Column( + String(36), + ForeignKey('sys_sso_session.sid', name='fk_oauth_refresh_sid', ondelete='RESTRICT'), + nullable=False, + comment='SSO Session ID', + ) + scopes = Column(JSON, nullable=False, comment='绑定 Scope') + resources = Column(JSON, nullable=False, comment='绑定 Resource audience') + status = Column(String(24), nullable=False, server_default='active', comment='Token 状态') + issued_at = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='签发时间') + last_used_at = Column(DbUtcDateTime(), nullable=True, comment='最近使用时间') + idle_expires_at = Column(DbUtcDateTime(), nullable=False, comment='闲置过期时间') + absolute_expires_at = Column(DbUtcDateTime(), nullable=False, comment='绝对过期时间') + revoked_at = Column(DbUtcDateTime(), nullable=True, comment='撤销时间') + revoke_reason = Column(String(200), nullable=True, comment='撤销原因') + reuse_detected_at = Column(DbUtcDateTime(), nullable=True, comment='重放检测时间') diff --git a/ruoyi-fastapi-backend/module_identity/entity/do/oauth_resource_do.py b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_resource_do.py new file mode 100644 index 000000000..da8655a14 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/do/oauth_resource_do.py @@ -0,0 +1,164 @@ +from sqlalchemy import ( + CHAR, + JSON, + BigInteger, + Column, + ForeignKey, + Index, + Integer, + PrimaryKeyConstraint, + SmallInteger, + String, + UniqueConstraint, +) + +from common.types import DbUtcDateTime +from config.database import Base +from utils.time_util import TimezoneUtil + + +class SysOAuthResource(Base): + """ + OAuth资源服务器表 + """ + + __tablename__ = 'sys_oauth_resource' + __table_args__ = ( + UniqueConstraint('resource_id', name='uk_oauth_resource_resource_id'), + UniqueConstraint('audience', name='uk_oauth_resource_audience'), + Index('idx_oauth_resource_status', 'status'), + {'comment': 'OAuth Resource'}, + ) + + resource_pk = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='内部主键', + ) + resource_id = Column(String(64), nullable=False, comment='Resource ID') + resource_name = Column(String(100), nullable=False, comment='Resource 名称') + audience = Column(String(500), nullable=False, comment='Access Token audience') + token_format = Column(String(16), nullable=False, server_default='jwt', comment='Token 格式') + signing_alg = Column(String(16), nullable=False, server_default='RS256', comment='签名算法') + access_token_ttl_seconds = Column(Integer, nullable=True, comment='Access Token 有效期') + introspection_client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_resource_introspection_client', ondelete='RESTRICT'), + nullable=True, + comment='Introspection Client 主键', + ) + allowed_claims = Column(JSON, nullable=False, comment='允许的 Claims') + status = Column(CHAR(1), nullable=False, server_default='0', comment='状态(0正常 1停用)') + create_by = Column(String(64), nullable=False, comment='创建者') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + update_by = Column(String(64), nullable=False, comment='更新者') + update_time = Column( + DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, onupdate=TimezoneUtil.utc_now, comment='更新时间' + ) + remark = Column(String(500), nullable=True, comment='备注') + + +class SysOAuthScope(Base): + """ + OAuth权限表 + """ + + __tablename__ = 'sys_oauth_scope' + __table_args__ = ( + UniqueConstraint('scope_code', name='uk_oauth_scope_code'), + Index('idx_oauth_scope_status', 'status'), + Index('idx_oauth_scope_resource', 'resource_pk'), + {'comment': 'OAuth Scope'}, + ) + + scope_pk = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='内部主键', + ) + scope_code = Column(String(100), nullable=False, comment='Scope 编码') + scope_name = Column(String(100), nullable=False, comment='Scope 名称') + scope_type = Column(String(16), nullable=False, comment='Scope 类型') + resource_pk = Column( + BigInteger, + ForeignKey('sys_oauth_resource.resource_pk', name='fk_oauth_scope_resource', ondelete='RESTRICT'), + nullable=True, + comment='Resource 主键', + ) + claims = Column(JSON, nullable=False, comment='Claims 列表') + consent_required = Column(SmallInteger, nullable=False, server_default='1', comment='是否需要同意') + sensitive = Column(SmallInteger, nullable=False, server_default='0', comment='是否敏感') + status = Column(CHAR(1), nullable=False, server_default='0', comment='状态(0正常 1停用)') + create_by = Column(String(64), nullable=False, comment='创建者') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + update_by = Column(String(64), nullable=False, comment='更新者') + update_time = Column( + DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, onupdate=TimezoneUtil.utc_now, comment='更新时间' + ) + remark = Column(String(500), nullable=True, comment='备注') + + +class SysOAuthClientScope(Base): + """ + OAuth客户端与权限关联表 + """ + + __tablename__ = 'sys_oauth_client_scope' + __table_args__ = ( + PrimaryKeyConstraint('client_pk', 'scope_pk', name='pk_oauth_client_scope'), + Index('idx_oauth_client_scope_scope', 'scope_pk'), + {'comment': 'OAuth Client Scope'}, + ) + + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_client_scope_client', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='Client 主键', + ) + scope_pk = Column( + BigInteger, + ForeignKey('sys_oauth_scope.scope_pk', name='fk_oauth_client_scope_scope', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='Scope 主键', + ) + is_default = Column(SmallInteger, nullable=False, server_default='0', comment='是否默认 Scope') + pre_authorized = Column(SmallInteger, nullable=False, server_default='0', comment='是否预授权') + claim_filter = Column(JSON, nullable=True, comment='Client Claim 过滤策略') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + + +class SysOAuthClientResource(Base): + """ + OAuth客户端与资源关联表 + """ + + __tablename__ = 'sys_oauth_client_resource' + __table_args__ = ( + PrimaryKeyConstraint('client_pk', 'resource_pk', name='pk_oauth_client_resource'), + Index('idx_oauth_client_resource_resource', 'resource_pk'), + {'comment': 'OAuth Client Resource'}, + ) + + client_pk = Column( + BigInteger, + ForeignKey('sys_oauth_client.client_pk', name='fk_oauth_client_resource_client', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='Client 主键', + ) + resource_pk = Column( + BigInteger, + ForeignKey('sys_oauth_resource.resource_pk', name='fk_oauth_client_resource_resource', ondelete='RESTRICT'), + primary_key=True, + nullable=False, + comment='Resource 主键', + ) + is_default = Column(SmallInteger, nullable=False, server_default='0', comment='是否默认 Resource') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') diff --git a/ruoyi-fastapi-backend/module_identity/entity/do/oidc_key_do.py b/ruoyi-fastapi-backend/module_identity/entity/do/oidc_key_do.py new file mode 100644 index 000000000..398c30d13 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/do/oidc_key_do.py @@ -0,0 +1,46 @@ +from sqlalchemy import JSON, BigInteger, CheckConstraint, Column, Index, Integer, String, Text, UniqueConstraint + +from common.types import DbUtcDateTime +from config.database import Base +from utils.time_util import TimezoneUtil + + +class SysOidcSigningKey(Base): + """ + OIDC签名密钥表 + """ + + __tablename__ = 'sys_oidc_signing_key' + __table_args__ = ( + UniqueConstraint('kid', name='uk_oidc_signing_key_kid'), + CheckConstraint( + '((CASE WHEN private_key_ref IS NULL THEN 0 ELSE 1 END) + ' + '(CASE WHEN private_key_ciphertext IS NULL THEN 0 ELSE 1 END)) = 1', + name='ck_oidc_signing_key_private_material', + ), + Index('idx_oidc_signing_key_status_publish', 'status', 'publish_at'), + Index('idx_oidc_signing_key_jwks_remove', 'status', 'remove_from_jwks_at'), + {'comment': 'OIDC Signing Key'}, + ) + + key_pk = Column( + BigInteger().with_variant(Integer, 'sqlite'), + primary_key=True, + nullable=False, + autoincrement=True, + comment='内部主键', + ) + kid = Column(String(100), nullable=False, comment='JWKS Key ID') + key_use = Column(String(16), nullable=False, server_default='sig', comment='JWK 用途') + alg = Column(String(16), nullable=False, server_default='RS256', comment='签名算法') + public_jwk = Column(JSON, nullable=False, comment='公开 JWK') + private_key_ref = Column(String(1000), nullable=True, comment='KMS/HSM/文件引用') + private_key_ciphertext = Column(Text, nullable=True, comment='加密私钥材料') + status = Column(String(16), nullable=False, comment='密钥状态') + publish_at = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='发布时间') + signing_start_at = Column(DbUtcDateTime(), nullable=True, comment='开始签名时间') + signing_stop_at = Column(DbUtcDateTime(), nullable=True, comment='停止签名时间') + remove_from_jwks_at = Column(DbUtcDateTime(), nullable=True, comment='移出 JWKS 时间') + create_by = Column(String(64), nullable=False, comment='创建者') + create_time = Column(DbUtcDateTime(), nullable=False, default=TimezoneUtil.utc_now, comment='创建时间') + remark = Column(String(500), nullable=True, comment='备注') diff --git a/ruoyi-fastapi-backend/module_identity/entity/vo/interaction_vo.py b/ruoyi-fastapi-backend/module_identity/entity/vo/interaction_vo.py new file mode 100644 index 000000000..d06b4e30b --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/vo/interaction_vo.py @@ -0,0 +1,133 @@ +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic.alias_generators import to_camel + +_MAX_INTERACTION_SCOPE_LENGTH = 100 + + +class InteractionModel(BaseModel): + """ + 认证交互 API 模型基类 + + 交互字段使用 camelCase,并拒绝未定义字段。 + """ + + model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, extra='forbid') + + +class InteractionLoginModel(InteractionModel): + """ + 认证中心登录提交模型 + """ + + user_name: str = Field(min_length=1, max_length=64, description='用户账号') + password: str = Field(min_length=1, max_length=256, description='登录密码') + code: str | None = Field(default=None, max_length=32, description='验证码') + uuid: str | None = Field(default=None, max_length=128, description='验证码唯一标识') + remember_me: bool = Field(default=False, description='是否保持长期登录') + + +class InteractionConsentModel(InteractionModel): + """ + 认证中心授权同意提交模型 + """ + + approved: bool = Field(description='是否同意本次授权') + scopes: list[str] = Field(default_factory=list, max_length=100, description='本次同意的权限范围') + remember_consent: bool = Field(default=False, description='是否记住本次选择以便后续免确认') + + @model_validator(mode='after') + def validate_scopes(self) -> 'InteractionConsentModel': + """ + 拒绝空 Scope、超长 Scope 和重复提交 + """ + + if any( + not isinstance(scope, str) or not scope or len(scope) > _MAX_INTERACTION_SCOPE_LENGTH + for scope in self.scopes + ): + raise ValueError('权限范围标识不能为空,且每项不得超过 100 个字符') + if len(set(self.scopes)) != len(self.scopes): + raise ValueError('权限范围不得重复') + return self + + +class ChangePasswordModel(InteractionModel): + """ + 初始密码或过期密码修改模型 + """ + + old_password: str = Field(min_length=1, max_length=256, description='当前密码') + new_password: str = Field(min_length=1, max_length=256, description='新密码') + confirm_password: str = Field(min_length=1, max_length=256, description='确认新密码') + + @model_validator(mode='after') + def validate_confirmation(self) -> 'ChangePasswordModel': + """ + 确保新密码与确认密码完全一致 + + :return: 当前已校验的改密请求 + """ + + if self.new_password != self.confirm_password: + raise ValueError('新密码与确认密码必须一致') + return self + + +class CaptchaResponseModel(InteractionModel): + """ + 认证交互验证码响应模型 + """ + + captcha_enabled: bool = Field(description='是否启用验证码') + uuid: str | None = Field(default=None, description='验证码唯一标识') + img: str | None = Field(default=None, description='验证码图片Base64内容') + + +class InteractionClientModel(InteractionModel): + """ + 授权交互页面展示的 Client 摘要 + """ + + client_id: str = Field(description='客户端标识') + client_name: str = Field(description='客户端名称') + logo_uri: str | None = Field(default=None, description='客户端图标地址') + policy_uri: str | None = Field(default=None, description='隐私政策地址') + tos_uri: str | None = Field(default=None, description='服务条款地址') + + +class RequestedScopeModel(InteractionModel): + """ + 授权交互页面展示的 Scope 摘要 + """ + + scope: str = Field(description='权限标识') + name: str = Field(description='权限名称') + description: str | None = Field(default=None, description='权限说明') + sensitive: bool = Field(default=False, description='是否为敏感权限') + required: bool = Field(default=False, description='是否为必选权限') + + +class InteractionViewModel(InteractionModel): + """ + Interaction 页面视图模型 + """ + + interaction_id: str = Field(description='认证交互标识') + client: InteractionClientModel = Field(description='发起认证的客户端摘要') + requested_scopes: list[RequestedScopeModel] = Field(default_factory=list, description='客户端请求的权限列表') + next_action: Literal['login', 'consent', 'redirect', 'changePassword'] = Field(description='认证交互的下一步动作') + captcha_enabled: bool = Field(default=False, description='是否启用验证码') + expires_in: int = Field(description='剩余有效时间,单位为秒') + + +class InteractionResultModel(InteractionModel): + """ + 登录、同意或改密后的下一步动作模型 + """ + + next_action: Literal['login', 'consent', 'redirect', 'changePassword'] = Field(description='认证交互的下一步动作') + interaction_id: str | None = Field(default=None, description='认证交互标识') + redirect_url: str | None = Field(default=None, description='服务端返回的下一步跳转地址') + reason: str | None = Field(default=None, description='下一步动作的原因') diff --git a/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_client_vo.py b/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_client_vo.py new file mode 100644 index 000000000..f1b7ff3dd --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_client_vo.py @@ -0,0 +1,221 @@ +from typing import Any, Literal + +from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic.alias_generators import to_camel + +from common.types import ApiUtcDateTime +from utils.oidc_util import OidcUtil + +_MAX_ROLE_KEY_LENGTH = 100 + + +class ClientModel(BaseModel): + """ + OAuth Client 管理模型基类 + """ + + model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, from_attributes=True, extra='forbid') + + +class ClientCreateModel(ClientModel): + """ + 创建 OAuth Client 请求 + """ + + client_name: str = Field(min_length=1, max_length=100, description='客户端名称') + client_type: Literal['public', 'confidential'] = Field( + description='客户端类型(public公开客户端 confidential机密客户端)' + ) + token_endpoint_auth_method: Literal['none', 'client_secret_basic'] = Field(description='令牌端点的客户端认证方式') + grant_types: list[str] = Field(default_factory=lambda: ['authorization_code'], description='允许使用的授权类型') + response_types: list[Literal['code']] = Field( + default_factory=lambda: ['code'], description='允许使用的授权响应类型' + ) + require_pkce: bool = Field(default=True, description='是否强制使用PKCE') + require_consent: bool = Field(default=True, description='是否要求用户确认授权') + trusted_client: bool = Field(default=False, description='是否为受信任客户端') + scope_codes: list[str] = Field(default_factory=list, description='客户端允许申请的权限标识列表') + allowed_role_keys: list[str] = Field( + default_factory=list, max_length=100, description='允许向客户端发布的角色权限字符,空列表表示不发布角色' + ) + pre_authorized_scope_codes: list[str] = Field(default_factory=list, description='预先授权的权限标识列表') + resource_ids: list[str] = Field(default_factory=list, description='允许访问的资源标识列表') + redirect_uris: list[str] = Field(default_factory=list, description='授权完成后允许跳转的回调地址列表') + post_logout_redirect_uris: list[str] = Field(default_factory=list, description='退出完成后允许跳转的回调地址列表') + backchannel_logout_uris: list[str] = Field(default_factory=list, description='后端退出通知地址列表') + cors_origins: list[str] = Field(default_factory=list, description='允许跨域访问的源地址列表') + access_token_ttl_seconds: int | None = Field(default=None, gt=0, description='访问令牌有效期,单位为秒') + refresh_token_idle_seconds: int | None = Field(default=None, gt=0, description='刷新令牌闲置有效期,单位为秒') + refresh_token_absolute_seconds: int | None = Field(default=None, gt=0, description='刷新令牌绝对有效期,单位为秒') + logo_uri: str | None = Field(default=None, description='客户端图标地址') + policy_uri: str | None = Field(default=None, description='隐私政策地址') + tos_uri: str | None = Field(default=None, description='服务条款地址') + remark: str | None = Field(default=None, max_length=500, description='备注') + + @model_validator(mode='before') + @classmethod + def default_non_code_response_types(cls, value: Any) -> Any: + """ + 为不含 authorization_code 的 Client 默认空 response_types + + :param value: 管理端提交的原始字段映射 + :return: 应用默认值后的原始字段映射 + """ + + if isinstance(value, dict) and 'response_types' not in value and 'responseTypes' not in value: + grant_key = 'grant_types' if 'grant_types' in value else 'grantTypes' + response_key = 'response_types' if 'grant_types' in value else 'responseTypes' + grants = value.get(grant_key, ['authorization_code']) + if isinstance(grants, list) and 'authorization_code' not in grants: + value = {**value, response_key: []} + return value + + @model_validator(mode='after') + def validate_auth_policy(self) -> 'ClientCreateModel': # noqa: PLR0912 + """ + 校验 Client Authentication、Grant、PKCE 与 URI 注册策略 + + :return: 当前已校验 Client 创建模型 + :raises ValueError: 策略、Grant 或注册 URI 不符合规范 + """ + + if len(set(self.allowed_role_keys)) != len(self.allowed_role_keys) or any( + not role or len(role) > _MAX_ROLE_KEY_LENGTH or '*' in role or any(char.isspace() for char in role) + for role in self.allowed_role_keys + ): + raise ValueError('允许发布的角色标识不得重复,也不得包含通配符或空白字符') + expected = 'none' if self.client_type == 'public' else 'client_secret_basic' + if self.token_endpoint_auth_method != expected: + raise ValueError('令牌端点认证方式与客户端类型不匹配') + if self.client_type == 'public' and not self.require_pkce: + raise ValueError('公开客户端必须启用 PKCE') + allowed_grants = {'authorization_code', 'refresh_token', 'client_credentials'} + if not self.grant_types or any(grant not in allowed_grants for grant in self.grant_types): + raise ValueError('grant_types 包含不支持的授权类型') + if len(set(self.grant_types)) != len(self.grant_types): + raise ValueError('grant_types 不得包含重复的授权类型') + if 'refresh_token' in self.grant_types and 'authorization_code' not in self.grant_types: + raise ValueError('启用刷新令牌必须同时启用授权码模式') + if self.client_type == 'public' and 'client_credentials' in self.grant_types: + raise ValueError('公开客户端不能使用 client_credentials 授权模式') + if 'authorization_code' in self.grant_types and not self.require_pkce: + raise ValueError('授权码客户端必须启用 PKCE') + if 'authorization_code' in self.grant_types: + if self.response_types != ['code']: + raise ValueError('授权码客户端的 response_types 必须为 [code]') + if not self.redirect_uris: + raise ValueError('授权码客户端必须配置登录回调地址') + elif self.response_types: + raise ValueError('未启用授权码模式的客户端不得配置响应类型') + for uri_type, uris in ( + ('redirect', self.redirect_uris), + ('post_logout', self.post_logout_redirect_uris), + ('backchannel_logout', self.backchannel_logout_uris), + ('cors_origin', self.cors_origins), + ): + for uri in uris: + OidcUtil.validate_registered_uri(uri_type, uri) + return self + + +class ClientUpdateModel(ClientCreateModel): + """ + 更新 OAuth Client 的管理请求 + + 更新请求必须携带外部 client_id。 + """ + + client_id: str = Field(min_length=1, max_length=64, description='客户端标识') + + +class ClientViewModel(ClientCreateModel): + """ + OAuth Client 安全字段脱敏后的详情模型 + + 模型不包含任何明文或哈希 Client Secret。 + """ + + client_id: str = Field(description='客户端标识') + status: Literal['0', '1'] = Field(default='0', description='状态(0正常 1停用)') + policy_version: int = Field(default=1, description='客户端策略版本') + create_time: ApiUtcDateTime | None = Field(default=None, description='创建时间') + update_time: ApiUtcDateTime | None = Field(default=None, description='更新时间') + + +class ClientSecretResponseModel(ClientModel): + """ + Client Secret 一次性展示响应 + + 调用方必须在本次响应后立即安全保存明文 Secret。 + """ + + client_id: str = Field(description='客户端标识') + secret_id: str = Field(description='客户端密钥标识') + client_secret: str = Field(description='仅在创建或轮换时返回的客户端密钥明文') + secret_hint: str = Field(description='客户端密钥展示提示') + not_before: ApiUtcDateTime = Field(description='密钥生效时间') + expires_at: ApiUtcDateTime | None = Field(default=None, description='过期时间') + + +class SecretRotationModel(ClientModel): + """ + Client Secret 轮换的时间策略 + """ + + not_before: ApiUtcDateTime | None = Field(default=None, description='密钥生效时间') + expires_at: ApiUtcDateTime | None = Field(default=None, description='过期时间') + retirement_seconds: int | None = Field(default=None, gt=0, description='旧客户端密钥的过渡期,单位为秒') + + +class ClientUriModel(ClientModel): + """ + 单个注册 URI 的管理模型 + + URI 在 DTO 边界执行精确匹配所需的安全校验。 + """ + + uri_type: Literal['redirect', 'post_logout', 'backchannel_logout', 'cors_origin'] = Field( + description='注册地址类型' + ) + uri: str = Field(min_length=1, max_length=1000, description='完整注册地址') + is_default: bool = Field(default=False, description='是否为默认地址') + status: Literal['0', '1'] = Field(default='0', description='状态(0正常 1停用)') + + @model_validator(mode='after') + def validate_uri(self) -> 'ClientUriModel': + """ + 校验并规范化单个注册 URI + + :return: 当前已校验 URI 模型 + """ + + self.uri = OidcUtil.validate_registered_uri(self.uri_type, self.uri) + + return self + + +class ClientPageQueryModel(ClientModel): + """ + Client 分页查询参数 + + page_num 与 page_size 使用管理端统一分页约束。 + """ + + client_name: str | None = Field(default=None, description='客户端名称') + client_type: Literal['public', 'confidential'] | None = Field( + default=None, description='客户端类型(public公开客户端 confidential机密客户端)' + ) + status: Literal['0', '1'] | None = Field(default=None, description='状态(0正常 1停用)') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + +class ClientStatusModel(ClientModel): + """ + Client 启停状态变更请求 + + status 只允许项目约定的 0(正常)或 1(停用)。 + """ + + client_id: str = Field(description='客户端标识') + status: Literal['0', '1'] = Field(description='状态(0正常 1停用)') diff --git a/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_resource_vo.py b/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_resource_vo.py new file mode 100644 index 000000000..5bf41b746 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_resource_vo.py @@ -0,0 +1,144 @@ +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator +from pydantic.alias_generators import to_camel + +from utils.oidc_util import OidcUtil + + +class ResourceModel(BaseModel): + """ + Resource 与 Scope 管理模型基类 + """ + + model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, from_attributes=True, extra='forbid') + + +class ScopeModel(ResourceModel): + """ + OAuth Scope 管理模型 + """ + + scope_code: str = Field(min_length=1, max_length=100, description='权限标识') + scope_name: str = Field(min_length=1, max_length=100, description='权限名称') + scope_type: Literal['identity', 'resource'] = Field(description='权限类型(identity身份权限 resource资源权限)') + resource_id: str | None = Field(default=None, description='资源标识') + claims: list[str] = Field(default_factory=list, description='权限关联的声明字段列表') + consent_required: bool = Field(default=True, description='是否要求用户确认该权限') + sensitive: bool = Field(default=False, description='是否为敏感权限') + status: Literal['0', '1'] = Field(default='0', description='状态(0正常 1停用)') + remark: str | None = Field(default=None, max_length=500, description='备注') + + _validate_code = field_validator('scope_code')(lambda value: OidcUtil.validate_path_identifier(value, 'scope_code')) + + @model_validator(mode='after') + def validate_scope_resource(self) -> 'ScopeModel': + """ + 校验 Resource Scope 与 Resource ID 的绑定关系 + + :return: 当前已校验 Scope + """ + + if self.scope_type == 'resource' and not self.resource_id: + raise ValueError('资源权限必须提供 resource_id') + if self.scope_type == 'identity' and self.resource_id: + raise ValueError('身份权限范围不能绑定业务资源') + return self + + +class ResourceCreateModel(ResourceModel): + """ + Resource Server 创建模型 + """ + + resource_id: str = Field(min_length=1, max_length=64, description='资源标识') + resource_name: str = Field(min_length=1, max_length=100, description='资源名称') + audience: str = Field(min_length=1, max_length=500, description='资源服务器的令牌受众标识') + token_format: Literal['jwt'] = Field(default='jwt', description='访问令牌格式') + signing_alg: Literal['RS256'] = Field(default='RS256', description='令牌签名算法') + access_token_ttl_seconds: int | None = Field(default=None, gt=0, description='访问令牌有效期,单位为秒') + introspection_client_id: str | None = Field(default=None, description='允许调用令牌内省的客户端标识') + allowed_claims: list[str] = Field(default_factory=list, description='允许发布的声明字段列表') + remark: str | None = Field(default=None, max_length=500, description='备注') + + _validate_resource_id = field_validator('resource_id')( + lambda value: OidcUtil.validate_path_identifier(value, 'resource_id') + ) + + +class ResourceUpdateModel(ResourceCreateModel): + """ + Resource Server 更新模型 + """ + + status: Literal['0', '1'] = Field(default='0', description='状态(0正常 1停用)') + + +class ResourceViewModel(ResourceCreateModel): + """ + Resource Server 管理详情模型 + """ + + status: Literal['0', '1'] = Field(default='0', description='状态(0正常 1停用)') + + +class ClaimPolicyModel(ResourceModel): + """ + Client/Resource Claim 允许列表模型 + """ + + client_id: str | None = Field(default=None, description='客户端标识') + resource_id: str | None = Field(default=None, description='资源标识') + allowed_claims: list[str] = Field(default_factory=list, description='允许发布的声明字段列表') + required_scopes: list[str] = Field(default_factory=list, description='发布声明所需的权限列表') + + +class ResourcePageQueryModel(ResourceModel): + """ + Resource Server 分页查询模型 + """ + + resource_name: str | None = Field(default=None, description='资源名称') + status: Literal['0', '1'] | None = Field(default=None, description='状态(0正常 1停用)') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + +class ScopePageQueryModel(ResourceModel): + """ + Scope 分页查询模型 + """ + + scope_name: str | None = Field(default=None, description='权限名称') + scope_type: Literal['identity', 'resource'] | None = Field( + default=None, description='权限类型(identity身份权限 resource资源权限)' + ) + status: Literal['0', '1'] | None = Field(default=None, description='状态(0正常 1停用)') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + +class ResourceStatusModel(ResourceModel): + """ + Resource Server 启停状态模型 + """ + + resource_id: str = Field(description='资源标识') + status: Literal['0', '1'] = Field(description='状态(0正常 1停用)') + + _validate_resource_id = field_validator('resource_id')( + lambda value: OidcUtil.validate_path_identifier(value, 'resource_id') + ) + + +class ScopeStatusModel(ResourceModel): + """ + Scope 启停状态模型 + """ + + scope_code: str = Field(description='权限标识') + status: Literal['0', '1'] = Field(description='状态(0正常 1停用)') + + _validate_scope_code = field_validator('scope_code')( + lambda value: OidcUtil.validate_path_identifier(value, 'scope_code') + ) diff --git a/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_session_vo.py b/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_session_vo.py new file mode 100644 index 000000000..ab6a2786f --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/vo/oauth_session_vo.py @@ -0,0 +1,209 @@ +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic.alias_generators import to_camel + +from common.types import ApiUtcDateTime + + +class SessionModel(BaseModel): + """ + Session、Grant 和 Audit 管理模型基类 + + 管理端字段接受 camelCase 别名并拒绝额外字段。 + """ + + model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, from_attributes=True, extra='forbid') + + +class GrantModel(SessionModel): + """ + 用户对 OAuth Client 的授权详情模型 + """ + + grant_id: str = Field(description='授权记录标识') + user_id: int | None = Field(default=None, description='用户ID') + subject_id: str | None = Field(default=None, description='用户稳定主体标识') + client_id: str = Field(description='客户端标识') + client_name: str | None = Field(default=None, description='客户端名称') + granted_scopes: list[str] = Field(default_factory=list, description='已同意的权限列表') + granted_resources: list[str] = Field(default_factory=list, description='已同意的资源受众列表') + remembered_scopes: list[str] = Field(default_factory=list, description='后续可免确认的权限列表') + remembered_resources: list[str] = Field(default_factory=list, description='后续可免确认的资源列表') + access_status: Literal['allowed', 'blocked'] = Field(default='allowed', description='用户对应用的访问策略') + access_reason: str | None = Field(default=None, description='访问策略的操作原因') + status: Literal['active', 'revoked', 'expired'] = Field( + default='active', description='状态(active有效 revoked已撤销 expired已过期)' + ) + client_policy_version: int | None = Field(default=None, description='授权时的客户端策略版本') + consented_at: ApiUtcDateTime | None = Field(default=None, description='用户同意授权的时间') + last_used_at: ApiUtcDateTime | None = Field(default=None, description='授权最近使用时间') + expires_at: ApiUtcDateTime | None = Field(default=None, description='过期时间') + revoke_reason: str | None = Field(default=None, description='撤销原因') + + +class SsoSessionModel(SessionModel): + """ + 认证中心 SSO Session 脱敏详情模型 + """ + + sid: str = Field(description='SSO会话标识') + user_id: int | None = Field(default=None, description='用户ID') + subject_id: str | None = Field(default=None, description='用户稳定主体标识') + auth_version: int | None = Field(default=None, description='用户认证安全版本') + auth_time: ApiUtcDateTime | None = Field(default=None, description='用户认证时间') + last_seen_at: ApiUtcDateTime | None = Field(default=None, description='会话最近活动时间') + idle_expires_at: ApiUtcDateTime | None = Field(default=None, description='闲置过期时间') + absolute_expires_at: ApiUtcDateTime | None = Field(default=None, description='绝对过期时间') + acr: str | None = Field(default=None, description='认证上下文') + amr: list[str] = Field(default_factory=list, description='认证方式列表') + remember_me: bool = Field(default=False, description='是否保持长期登录') + ip_address: str | None = Field(default=None, description='IP地址') + status: Literal['active', 'revoked', 'expired'] = Field( + default='active', description='状态(active有效 revoked已撤销 expired已过期)' + ) + revoked_at: ApiUtcDateTime | None = Field(default=None, description='撤销时间') + revoke_reason: str | None = Field(default=None, description='撤销原因') + create_time: ApiUtcDateTime | None = Field(default=None, description='创建时间') + client_ids: list[str] = Field(default_factory=list, description='会话关联的客户端标识列表') + + +class SessionRevokeModel(SessionModel): + """ + SSO Session 撤销请求模型 + """ + + reason: str = Field(min_length=1, max_length=200, description='会话或授权的撤销原因') + + +class GrantAccessModel(SessionRevokeModel): + """ + 用户对应用的访问控制请求模型 + """ + + blocked: bool = Field(strict=True, description='是否禁止用户访问应用') + + +class SessionPageQueryModel(SessionModel): + """ + SSO Session 分页查询模型 + """ + + user_id: int | None = Field(default=None, description='用户ID') + ip_address: str | None = Field(default=None, description='IP地址') + status: Literal['active', 'revoked', 'expired'] | None = Field( + default=None, description='状态(active有效 revoked已撤销 expired已过期)' + ) + start_time: ApiUtcDateTime | None = Field(default=None, description='查询开始时间') + end_time: ApiUtcDateTime | None = Field(default=None, description='查询结束时间') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + @model_validator(mode='after') + def validate_time_range(self) -> 'SessionPageQueryModel': + """ + 校验 Session 查询时间范围 + """ + + if self.start_time and self.end_time and self.end_time < self.start_time: + raise ValueError('结束时间不得早于开始时间') + return self + + +class GrantPageQueryModel(SessionModel): + """ + OAuth Grant 分页查询模型 + """ + + user_id: int | None = Field(default=None, description='用户ID') + client_id: str | None = Field(default=None, description='客户端标识') + status: Literal['active', 'revoked', 'expired'] | None = Field( + default=None, description='状态(active有效 revoked已撤销 expired已过期)' + ) + access_status: Literal['allowed', 'blocked'] | None = Field(default=None, description='用户对应用的访问策略') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + +class AccessPolicyModel(SessionModel): + """ + 用户应用访问策略管理模型 + """ + + user_id: int = Field(description='用户ID') + user_name: str = Field(description='用户名称') + client_id: str = Field(description='客户端标识') + client_name: str = Field(description='客户端名称') + access_status: Literal['allowed', 'blocked'] = Field(description='用户对应用的访问策略') + reason: str | None = Field(default=None, description='操作原因') + update_by: str | None = Field(default=None, description='操作人') + update_time: ApiUtcDateTime = Field(description='最近操作时间') + + +class AccessPolicyPageQueryModel(SessionModel): + """ + 用户应用访问策略分页查询模型 + """ + + user_id: int | None = Field(default=None, gt=0, description='用户ID') + client_id: str | None = Field(default=None, max_length=128, description='客户端标识') + access_status: Literal['allowed', 'blocked'] | None = Field(default=None, description='访问策略') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + +class AuditPageQueryModel(SessionModel): + """ + OAuth 审计分页查询模型 + + 支持按 Client、用户、事件结果、风险和时间范围过滤。 + """ + + event_type: str | None = Field(default=None, description='审计事件类型') + client_id: str | None = Field(default=None, description='客户端标识') + user_id: int | None = Field(default=None, description='用户ID') + result: Literal['success', 'failure'] | None = Field( + default=None, description='处理结果(success成功 failure失败)' + ) + risk_level: Literal['normal', 'medium', 'high', 'critical'] | None = Field( + default=None, description='风险等级(normal普通 medium中等 high高 critical严重)' + ) + start_time: ApiUtcDateTime | None = Field(default=None, description='查询开始时间') + end_time: ApiUtcDateTime | None = Field(default=None, description='查询结束时间') + page_num: int = Field(default=1, ge=1, description='当前页码') + page_size: int = Field(default=10, ge=1, le=200, description='每页记录数') + + @model_validator(mode='after') + def validate_time_range(self) -> 'AuditPageQueryModel': + """ + 校验审计查询的起止时间顺序 + + :return: 当前已校验查询模型 + :raises ValueError: 结束时间早于开始时间 + """ + + if self.start_time and self.end_time and self.end_time < self.start_time: + raise ValueError('结束时间不得早于开始时间') + return self + + +class AuditModel(SessionModel): + """ + 脱敏 OAuth 审计事件详情模型 + """ + + audit_id: int = Field(description='审计记录ID') + event_type: str = Field(description='审计事件类型') + result: Literal['success', 'failure'] = Field(description='处理结果(success成功 failure失败)') + risk_level: Literal['normal', 'medium', 'high', 'critical'] = Field( + default='normal', description='风险等级(normal普通 medium中等 high高 critical严重)' + ) + trace_id: str | None = Field(default=None, description='请求跟踪标识') + client_id: str | None = Field(default=None, description='客户端标识') + resource_id: str | None = Field(default=None, description='资源标识') + user_id: int | None = Field(default=None, description='用户ID') + subject_id: str | None = Field(default=None, description='用户稳定主体标识') + sid: str | None = Field(default=None, description='SSO会话标识') + ip_address: str | None = Field(default=None, description='IP地址') + failure_code: str | None = Field(default=None, description='失败原因代码') + create_time: ApiUtcDateTime = Field(description='创建时间') diff --git a/ruoyi-fastapi-backend/module_identity/entity/vo/oidc_key_vo.py b/ruoyi-fastapi-backend/module_identity/entity/vo/oidc_key_vo.py new file mode 100644 index 000000000..a29260c55 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/vo/oidc_key_vo.py @@ -0,0 +1,59 @@ +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, field_validator +from pydantic.alias_generators import to_camel + +from common.types import ApiUtcDateTime +from module_identity.entity.vo.protocol_vo import Jwk +from utils.oidc_util import OidcUtil + + +class OidcKeyModel(BaseModel): + """ + OIDC 签名密钥管理模型基类 + """ + + model_config = ConfigDict(alias_generator=to_camel, populate_by_name=True, from_attributes=True, extra='forbid') + + +class OidcKeyRotateModel(OidcKeyModel): + """ + OIDC 签名密钥轮换请求模型 + """ + + alg: Literal['RS256'] = Field(default='RS256', description='签名算法') + kid: str = Field(min_length=1, max_length=100, description='签名密钥标识') + publish_at: ApiUtcDateTime = Field(description='计划发布到JWKS的时间') + activate_at: ApiUtcDateTime | None = Field(default=None, description='计划启用签名的时间,为空时需手动激活') + remark: str | None = Field(default=None, max_length=500, description='备注') + + @field_validator('kid') + @classmethod + def validate_kid(cls, value: str) -> str: + """ + 拒绝无法安全放入密钥路径的 kid + + :param value: 待校验的签名密钥标识 + :return: 校验通过的签名密钥标识 + """ + + return OidcUtil.validate_path_identifier(value, '签名密钥标识 kid') + + +class OidcKeyViewModel(OidcKeyModel): + """ + 仅包含公开 JWK 的签名密钥详情模型 + """ + + kid: str = Field(description='签名密钥标识') + key_use: Literal['sig'] = Field(default='sig', description='密钥用途') + alg: Literal['RS256'] = Field(default='RS256', description='签名算法') + public_jwk: Jwk = Field(description='签名公钥的JWK表示') + status: Literal['pending', 'active', 'retiring', 'retired', 'compromised'] = Field( + description='密钥状态(pending待启用 active签名中 retiring退役中 retired已退役 compromised已泄露)' + ) + publish_at: ApiUtcDateTime = Field(description='计划发布到JWKS的时间') + signing_start_at: ApiUtcDateTime | None = Field(default=None, description='开始签名的时间') + signing_stop_at: ApiUtcDateTime | None = Field(default=None, description='停止签名的时间') + remove_from_jwks_at: ApiUtcDateTime | None = Field(default=None, description='从JWKS移除的时间') + create_time: ApiUtcDateTime | None = Field(default=None, description='创建时间') diff --git a/ruoyi-fastapi-backend/module_identity/entity/vo/protocol_vo.py b/ruoyi-fastapi-backend/module_identity/entity/vo/protocol_vo.py new file mode 100644 index 000000000..649f50484 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/entity/vo/protocol_vo.py @@ -0,0 +1,329 @@ +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field, NonNegativeInt, field_validator, model_validator + +from utils.oidc_util import OidcUtil + + +class ProtocolModel(BaseModel): + """ + 标准协议模型基类 + + 标准端点字段保持 snake_case,并拒绝未定义参数。 + """ + + model_config = ConfigDict(extra='forbid', populate_by_name=True, str_strip_whitespace=True) + + +class AuthorizeRequest(ProtocolModel): + """ + Authorization Endpoint 请求参数 + + OIDC 请求必须携带 nonce,并固定使用 PKCE S256。 + """ + + response_type: str = Field(min_length=1, max_length=100, description='授权响应类型') + response_mode: Literal['query'] = Field(default='query', description='授权响应参数传递方式') + client_id: str = Field(min_length=1, max_length=64, description='客户端标识') + redirect_uri: str = Field(min_length=1, max_length=1000, description='已注册的授权回调地址') + scope: str = Field(min_length=1, max_length=2000, description='空格分隔的权限范围') + resource: str | None = Field(default=None, min_length=1, max_length=500, description='请求访问的资源受众') + state: str | None = Field(default=None, max_length=1024, description='客户端用于关联请求和响应的状态值') + nonce: str | None = Field(default=None, max_length=1024, description='绑定认证请求与ID Token的随机数') + code_challenge: str = Field(min_length=43, max_length=128, description='PKCE校验挑战值') + code_challenge_method: str = Field(min_length=1, max_length=20, description='PKCE挑战方法') + prompt: str | None = Field(default=None, max_length=100, description='空格分隔的认证交互要求') + max_age: NonNegativeInt | None = Field(default=None, description='允许的最长认证间隔,单位为秒') + login_hint: str | None = Field(default=None, max_length=512, description='客户端提供的登录账号提示') + + @field_validator('code_challenge') + @classmethod + def validate_code_challenge(cls, value: str) -> str: + """ + 校验未填充的 base64url SHA-256 challenge + + :param value: 客户端提交的 code_challenge + :return: 已校验 challenge + """ + + if not OidcUtil.is_s256_challenge(value): + raise ValueError('code_challenge 必须为不带填充的 Base64URL 编码 SHA-256 摘要') + return value + + @field_validator('prompt') + @classmethod + def validate_prompt(cls, value: str | None) -> str | None: + """ + 只接受规范支持的 prompt 值,且 none 不能和其他值混用 + + :param value: 空格分隔的 prompt 列表 + :return: 规范化后的 prompt 或 None + """ + + return OidcUtil.normalize_prompt(value) + + @model_validator(mode='after') + def validate_oidc_nonce(self) -> 'AuthorizeRequest': + """ + 校验 OIDC 授权请求必须绑定 nonce + + :return: 当前已校验请求 + """ + + if 'openid' in self.scope.split() and not self.nonce: + raise ValueError('申请 openid 权限时必须提供 nonce') + return self + + +class TokenRequest(ProtocolModel): + """ + Token Endpoint 表单参数 + + 不同 grant_type 只允许携带其对应凭据。 + """ + + grant_type: str = Field(min_length=1, max_length=50, description='授权类型') + code: str | None = Field(default=None, min_length=1, max_length=4096, description='授权码') + redirect_uri: str | None = Field(default=None, min_length=1, max_length=1000, description='已注册的授权回调地址') + client_id: str | None = Field(default=None, min_length=1, max_length=64, description='客户端标识') + code_verifier: str | None = Field(default=None, min_length=43, max_length=128, description='PKCE校验原文') + refresh_token: str | None = Field(default=None, min_length=1, max_length=4096, description='刷新令牌') + scope: str | None = Field(default=None, min_length=1, max_length=2000, description='空格分隔的权限范围') + resource: str | None = Field(default=None, min_length=1, max_length=1000, description='请求访问的资源受众') + + @model_validator(mode='after') + def check_grant_fields(self) -> 'TokenRequest': + """ + 拒绝与 grant type 不相符的凭据,避免跨流程混用 + + :return: 当前已校验请求 + """ + + if self.grant_type == 'authorization_code': + if not self.code or not self.code_verifier or not self.redirect_uri: + raise ValueError('授权码兑换必须提供 code、redirect_uri 和 code_verifier') + if self.refresh_token or self.scope is not None or self.resource is not None: + raise ValueError('授权码兑换请求包含不适用的字段') + elif self.grant_type == 'refresh_token': + if not self.refresh_token: + raise ValueError('刷新令牌请求必须提供 refresh_token') + if self.code or self.code_verifier or self.redirect_uri: + raise ValueError('刷新令牌请求包含不适用的字段') + elif self.grant_type == 'client_credentials' and any( + (self.code, self.code_verifier, self.refresh_token, self.redirect_uri) + ): + raise ValueError('客户端凭据授权请求包含不适用的字段') + return self + + +class TokenResponse(ProtocolModel): + """ + Token Endpoint 成功响应 + + 字段直接对应 OAuth 2.0 标准 JSON 响应。 + """ + + access_token: str = Field(description='访问令牌') + token_type: Literal['Bearer'] = Field(default='Bearer', description='令牌类型') + expires_in: NonNegativeInt = Field(description='剩余有效时间,单位为秒') + refresh_token: str | None = Field(default=None, description='刷新令牌') + scope: str | None = Field(default=None, description='空格分隔的权限范围') + id_token: str | None = Field(default=None, description='身份令牌') + + +class ErrorResponse(ProtocolModel): + """ + OAuth/OIDC 标准错误响应 + + 不包含项目业务响应包装字段。 + """ + + error: str = Field(description='标准协议错误代码') + error_description: str | None = Field(default=None, description='标准协议错误说明') + error_uri: str | None = Field(default=None, description='错误说明页面地址') + + def as_dict(self) -> dict[str, str]: + """ + 返回可直接作为标准 JSON 响应的非空字段 + + :return: 排除 None 字段的标准错误字典 + """ + + return self.model_dump(exclude_none=True) + + +class Jwk(ProtocolModel): + """ + 公开 RSA JWK + + 模型严格拒绝 d、p、q 等私钥参数。 + """ + + kty: Literal['RSA'] = Field(default='RSA', description='公钥类型') + use: Literal['sig'] = Field(default='sig', description='公钥用途') + kid: str = Field(min_length=1, max_length=100, description='签名密钥标识') + alg: Literal['RS256'] = Field(default='RS256', description='签名算法') + n: str = Field(min_length=1, description='RSA公钥模数的Base64URL编码') + e: str = Field(min_length=1, description='RSA公钥指数的Base64URL编码') + + +class JwksResponse(ProtocolModel): + """ + JWKS 响应 + + keys 仅包含公开签名 JWK。 + """ + + keys: list[Jwk] = Field(description='公开签名密钥列表') + + +class DiscoveryResponse(ProtocolModel): + """ + OIDC Discovery 响应 + + 元数据只宣称当前实现确实支持的端点和算法。 + """ + + issuer: str = Field(description='认证中心签发方地址') + authorization_endpoint: str = Field(description='授权端点地址') + token_endpoint: str = Field(description='令牌端点地址') + userinfo_endpoint: str = Field(description='用户信息端点地址') + jwks_uri: str = Field(description='公开签名密钥集地址') + revocation_endpoint: str = Field(description='令牌撤销端点地址') + introspection_endpoint: str = Field(description='令牌内省端点地址') + end_session_endpoint: str = Field(description='退出端点地址') + scopes_supported: list[str] = Field(default_factory=list, description='支持的权限列表') + response_types_supported: list[Literal['code']] = Field( + default_factory=lambda: ['code'], description='支持的授权响应类型列表' + ) + response_modes_supported: list[Literal['query']] = Field( + default_factory=lambda: ['query'], description='支持的授权响应模式列表' + ) + grant_types_supported: list[str] = Field( + default_factory=lambda: ['authorization_code', 'refresh_token'], description='支持的授权类型列表' + ) + subject_types_supported: list[Literal['public']] = Field( + default_factory=lambda: ['public'], description='支持的主体标识类型列表' + ) + id_token_signing_alg_values_supported: list[Literal['RS256']] = Field( + default_factory=lambda: ['RS256'], description='支持的ID Token签名算法列表' + ) + token_endpoint_auth_methods_supported: list[Literal['none', 'client_secret_basic']] = Field( + default_factory=lambda: ['none', 'client_secret_basic'], description='支持的令牌端点认证方式列表' + ) + code_challenge_methods_supported: list[Literal['S256']] = Field( + default_factory=lambda: ['S256'], description='支持的PKCE挑战方法列表' + ) + claims_supported: list[str] = Field(default_factory=list, description='支持的声明字段列表') + authorization_response_iss_parameter_supported: bool = Field( + default=True, description='是否支持在授权响应中返回签发方' + ) + backchannel_logout_supported: bool = Field(default=False, description='是否支持后端退出通知') + backchannel_logout_session_supported: bool = Field(default=False, description='后端退出通知是否支持会话标识') + + +class OAuthServerMetadata(ProtocolModel): + """ + RFC 8414 OAuth Authorization Server Metadata 子集 + + 不继承要求 OIDC 专属 userinfo_endpoint 和 end_session_endpoint 的模型。 + """ + + issuer: str = Field(description='认证中心签发方地址') + authorization_endpoint: str = Field(description='授权端点地址') + token_endpoint: str = Field(description='令牌端点地址') + jwks_uri: str = Field(description='公开签名密钥集地址') + revocation_endpoint: str | None = Field(default=None, description='令牌撤销端点地址') + introspection_endpoint: str | None = Field(default=None, description='令牌内省端点地址') + scopes_supported: list[str] = Field(default_factory=list, description='支持的权限列表') + response_types_supported: list[Literal['code']] = Field( + default_factory=lambda: ['code'], description='支持的授权响应类型列表' + ) + grant_types_supported: list[str] = Field( + default_factory=lambda: ['authorization_code', 'refresh_token'], description='支持的授权类型列表' + ) + token_endpoint_auth_methods_supported: list[Literal['none', 'client_secret_basic']] = Field( + default_factory=lambda: ['none', 'client_secret_basic'], description='支持的令牌端点认证方式列表' + ) + code_challenge_methods_supported: list[Literal['S256']] = Field( + default_factory=lambda: ['S256'], description='支持的PKCE挑战方法列表' + ) + + +class UserInfoResponse(ProtocolModel): + """ + UserInfo 基础响应 + + 允许的扩展字段对应已授权的标准身份声明。 + """ + + sub: str = Field(description='用户稳定主体标识') + name: str | None = Field(default=None, description='显示名称') + preferred_username: str | None = Field(default=None, description='用户首选账号名称') + picture: str | None = Field(default=None, description='用户头像地址') + updated_at: int | None = Field(default=None, description='用户资料更新时间,Unix时间戳,单位为秒') + email: str | None = Field(default=None, description='电子邮箱') + email_verified: bool | None = Field(default=None, description='电子邮箱是否已验证') + phone_number: str | None = Field(default=None, description='手机号码') + phone_number_verified: bool | None = Field(default=None, description='手机号码是否已验证') + dept_id: str | int | None = Field(default=None, description='部门ID') + dept_name: str | None = Field(default=None, description='部门名称') + roles: list[str] | None = Field(default=None, description='当前授权允许发布的角色权限字符列表') + + +class RevocationRequest(ProtocolModel): + """ + Token Revocation 请求 + + token_type_hint 只用于减少服务端查找歧义,不作为安全依据。 + """ + + token: str = Field(min_length=1, description='待撤销或内省的令牌') + token_type_hint: Literal['access_token', 'refresh_token'] | None = Field(default=None, description='令牌类型提示') + + +class IntrospectionRequest(ProtocolModel): + """ + Token Introspection 请求 + + 请求方权限和 audience 绑定由服务层继续校验。 + """ + + token: str = Field(min_length=1, description='待撤销或内省的令牌') + token_type_hint: Literal['access_token', 'refresh_token'] | None = Field(default=None, description='令牌类型提示') + + +class IntrospectionResponse(ProtocolModel): + """ + Token Introspection 响应 + + inactive Token 只应返回 active=false。 + """ + + active: bool = Field(description='令牌是否有效') + scope: str | None = Field(default=None, description='空格分隔的权限范围') + client_id: str | None = Field(default=None, description='客户端标识') + username: str | None = Field(default=None, description='用户账号') + token_type: str | None = Field(default=None, description='令牌类型') + exp: int | None = Field(default=None, description='过期时间,Unix时间戳,单位为秒') + iat: int | None = Field(default=None, description='签发时间,Unix时间戳,单位为秒') + nbf: int | None = Field(default=None, description='最早生效时间,Unix时间戳,单位为秒') + sub: str | None = Field(default=None, description='用户稳定主体标识') + aud: str | list[str] | None = Field(default=None, description='令牌接收方') + iss: str | None = Field(default=None, description='令牌签发方') + jti: str | None = Field(default=None, description='令牌唯一标识') + sid: str | None = Field(default=None, description='SSO会话标识') + + +class LogoutRequest(ProtocolModel): + """ + RP-Initiated Logout 请求 + + post_logout_redirect_uri 必须由服务端按 Client 注册值精确匹配。 + """ + + id_token_hint: str | None = Field(default=None, description='用于关联退出会话的ID Token提示') + logout_hint: str | None = Field(default=None, description='退出会话提示') + client_id: str | None = Field(default=None, description='客户端标识') + post_logout_redirect_uri: str | None = Field(default=None, description='已注册的退出回调地址') + state: str | None = Field(default=None, description='客户端用于关联请求和响应的状态值') diff --git a/ruoyi-fastapi-backend/module_identity/redis_keys.py b/ruoyi-fastapi-backend/module_identity/redis_keys.py new file mode 100644 index 000000000..61e8274f8 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/redis_keys.py @@ -0,0 +1,218 @@ +from typing import Final + +from utils.oidc_util import OidcUtil + + +class OidcRedisKey: + """ + 统一认证中心 Redis Key Builder + """ + + PREFIX: Final[str] = 'oidc' + SESSION_REVOKED_CHANNEL: Final[str] = 'oidc:event:session_revoked' + USER_SECURITY_CHANGED_CHANNEL: Final[str] = 'oidc:event:user_security_changed' + + @classmethod + def interaction(cls, interaction_id: str) -> str: + """ + 生成 Interaction 状态键 + + :param interaction_id: Interaction 标识 + :return: Redis Key + """ + + return f'{cls.PREFIX}:interaction:{OidcUtil.redis_key_component(interaction_id, name="interaction_id")}' + + @classmethod + def authorization_code(cls, code_id: str) -> str: + """ + 生成授权码状态键 + + :param code_id: 授权码内部标识 + :return: Redis Key + """ + + return f'{cls.PREFIX}:authorization_code:{OidcUtil.redis_key_component(code_id, name="code_id")}' + + @classmethod + def authorization_code_consumed(cls, code_id: str) -> str: + """ + 生成授权码消费短期 tombstone 键 + + :param code_id: 授权码公开标识部分 + :return: 不含授权码 Secret 的消费状态键 + """ + + return f'{cls.PREFIX}:authorization_code_consumed:{OidcUtil.redis_key_component(code_id, name="code_id")}' + + @classmethod + def authorization_code_consumed_payload(cls, code_id: str) -> str: + """ + 生成授权码消费绑定载荷短期键 + + :param code_id: 授权码公开标识部分 + :return: 不含授权码 Secret 的消费绑定载荷键 + """ + + return ( + f'{cls.PREFIX}:authorization_code_consumed_payload:{OidcUtil.redis_key_component(code_id, name="code_id")}' + ) + + @classmethod + def sso_session(cls, sid: str) -> str: + """ + 生成 SSO 会话热缓存键 + + :param sid: SSO 会话标识 + :return: Redis Key + """ + + return f'{cls.PREFIX}:sso_session:{OidcUtil.redis_key_component(sid, name="sid")}' + + @classmethod + def user_sessions(cls, user_id: str | int) -> str: + """ + 生成用户 SSO 会话索引键 + + :param user_id: 内部用户标识 + :return: Redis Key + """ + + return f'{cls.PREFIX}:user_sessions:{OidcUtil.redis_key_component(user_id, name="user_id")}' + + @classmethod + def sso_cookie(cls, session_hash: str) -> str: + """ + 生成 SSO Cookie 摘要索引键 + + :param session_hash: Cookie Secret 的 SHA-256 摘要 + :return: Redis Key + """ + + return f'{cls.PREFIX}:sso_cookie:{OidcUtil.sha256_hex_digest(session_hash)}' + + @classmethod + def revoked_jti(cls, jti: str) -> str: + """ + 生成 Access Token JTI 撤销键 + + :param jti: Access Token 的 JTI + :return: Redis Key + """ + + return f'{cls.PREFIX}:revoked_jti:{OidcUtil.redis_key_component(jti, name="jti")}' + + @classmethod + def backchannel_retry_queue(cls) -> str: + """ + 返回 Back-Channel Logout 有界重试队列键 + """ + + return f'{cls.PREFIX}:backchannel:retry' + + @classmethod + def signing_key_rotation_lock(cls) -> str: + """ + 返回跨实例签名密钥轮换锁键 + + :return: OIDC 签名密钥轮换专用 Redis 锁键 + """ + + return f'{cls.PREFIX}:signing_key:rotation_lock' + + @classmethod + def authorize_ip_rate_limit(cls, ip_hash: str) -> str: + """ + 生成按 IP 限流键 + + :param ip_hash: 使用独立 Pepper 生成的 IP 摘要 + :return: Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:authorize:ip:{OidcUtil.sha256_hex_digest(ip_hash)}' + + @classmethod + def login_user_rate_limit(cls, user_name_hash: str) -> str: + """ + 生成按用户限流键 + + :param user_name_hash: 使用独立 Pepper 生成的用户名摘要 + :return: Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:login:user:{OidcUtil.sha256_hex_digest(user_name_hash)}' + + @classmethod + def interaction_captcha_rate_limit(cls, subject_hash: str) -> str: + """ + 生成认证中心验证码专用限流键 + + :param subject_hash: Interaction 与来源地址组合的 HMAC 摘要 + :return: 验证码端点独立 Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:interaction:captcha:{OidcUtil.sha256_hex_digest(subject_hash)}' + + @classmethod + def token_client_rate_limit(cls, client_id: str) -> str: + """ + 生成按 Client 限流键 + + :param client_id: 外部 Client ID + :return: Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:token:client:{OidcUtil.redis_key_component(client_id, name="client_id")}' + + @classmethod + def introspect_client_rate_limit(cls, client_id: str) -> str: + """ + 生成按 introspection Client 限流键 + + :param client_id: introspection Client ID + :return: Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:introspect:client:{OidcUtil.redis_key_component(client_id, name="client_id")}' + + @classmethod + def revoke_client_rate_limit(cls, client_id: str) -> str: + """ + 返回独立于 Token Endpoint 的撤销端点限流键 + + :param client_id: 撤销调用方的外部 Client ID 或匿名占位标识 + :return: 撤销端点专用 Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:revoke:client:{OidcUtil.redis_key_component(client_id, name="client_id")}' + + @classmethod + def logout_rate_limit(cls, scope: str = 'anonymous') -> str: + """ + 返回按安全摘要隔离的 RP-Initiated Logout 限流键 + + :param scope: 已由调用方 HMAC 的 IP 或安全 fallback 摘要 + :return: Logout Endpoint 专用 Redis Key + """ + + return f'{cls.PREFIX}:rate_limit:logout:{OidcUtil.redis_key_component(scope, name="logout_scope")}' + + @classmethod + def event_session_revoked(cls) -> str: + """ + 返回 SSO 会话撤销事件频道 + + :return: Redis Pub/Sub 频道名 + """ + + return cls.SESSION_REVOKED_CHANNEL + + @classmethod + def event_user_security_changed(cls) -> str: + """ + 返回用户安全变化事件频道 + + :return: Redis Pub/Sub 频道名 + """ + + return cls.USER_SECURITY_CHANGED_CHANNEL diff --git a/ruoyi-fastapi-backend/module_identity/security/backchannel_transport.py b/ruoyi-fastapi-backend/module_identity/security/backchannel_transport.py new file mode 100644 index 000000000..7a4fd2eaf --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/backchannel_transport.py @@ -0,0 +1,227 @@ +from collections.abc import Awaitable, Callable +from typing import Any +from urllib.parse import urlsplit + +import httpcore +import httpx +from httpcore._backends.auto import AutoBackend + +_MAX_BACKCHANNEL_RESPONSE_BYTES = 64 * 1024 +_HTTP_CLIENT_ERROR_MIN = 400 +_HTTP_SERVER_ERROR_MIN = 500 +_HTTP_REQUEST_TIMEOUT = 408 +_HTTP_TOO_MANY_REQUESTS = 429 + + +class PermanentBackchannelError(ValueError): + """ + 表示 Back-Channel 请求永久失败且不应继续重试 + """ + + def __init__(self, failure_code: str, message: str | None = None) -> None: + """分别保存稳定的审计错误码和中文异常说明。""" + self.failure_code = failure_code + self.message = message or '后端退出通知永久失败' + super().__init__(self.message) + + +class _PinnedNetworkBackend(httpcore.AsyncNetworkBackend): + """ + 保持注册域名的 TLS 和 SNI,同时将 TCP 连接固定到已验证 IP + """ + + def __init__(self, hostname: str, addresses: set[str]) -> None: + """ + 初始化固定目标地址的网络后端 + + :param hostname: 已验证的注册主机名 + :param addresses: 已验证的公网 IP 集合 + :return: None + """ + + self._hostname = hostname + self._addresses = addresses + self._backend = AutoBackend() + + async def connect_tcp(self, host: str, port: int, **kwargs: Any) -> httpcore.AsyncNetworkStream: + """ + 连接已验证主机对应的固定公网地址 + + :param host: HTTP Core 请求的原始主机名 + :param port: 目标 TCP 端口 + :param kwargs: HTTP Core 连接参数 + :return: 已建立的异步网络流 + :raises OSError: 主机未验证、地址为空或全部地址连接失败 + """ + + if host != self._hostname or not self._addresses: + raise OSError('网络目标未通过安全校验') + # 逐个尝试已验证地址,网络库仍以原始 origin 处理 TLS SNI 和 Host + last_error: OSError | None = None + for address in sorted(self._addresses): + try: + return await self._backend.connect_tcp(address, port, **kwargs) + except OSError as exc: # noqa: PERF203 + last_error = exc + raise last_error or OSError('没有通过安全校验的网络目标') + + +class PinnedHttpxTransport(httpx.AsyncBaseTransport): + """ + 单次请求使用的 IP 固定 HTTPX Transport + """ + + def __init__(self, hostname: str, addresses: set[str]) -> None: + """ + 初始化使用固定网络后端的连接池 + + :param hostname: 已验证的注册主机名 + :param addresses: 已验证的公网 IP 集合 + :return: None + """ + + self._pool = httpcore.AsyncConnectionPool(network_backend=_PinnedNetworkBackend(hostname, addresses)) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + """ + 通过固定地址连接池发送异步 HTTP 请求 + + :param request: HTTPX 异步请求 + :return: 限制响应体大小的 HTTPX 响应 + """ + + request_stream = httpcore.Request( + method=request.method, + url=httpcore.URL( + scheme=request.url.raw_scheme, + host=request.url.raw_host, + port=request.url.port, + target=request.url.raw_path, + ), + headers=request.headers.raw, + content=request.stream, + extensions=request.extensions, + ) + response = await self._pool.handle_async_request(request_stream) + + class ResponseStream(httpx.AsyncByteStream): + """ + 限制 Back-Channel 响应体大小的异步字节流 + """ + + def __init__(self) -> None: + """ + 初始化已接收字节计数器 + + :return: None + """ + + self._received = 0 + + async def __aiter__(self) -> Any: + """ + 逐块读取响应并执行大小限制 + + :return: HTTP 响应字节块异步迭代器 + :raises OSError: 响应体超过大小限制 + """ + + async for chunk in response.stream: + self._received += len(chunk) + if self._received > _MAX_BACKCHANNEL_RESPONSE_BYTES: + await response.aclose() + raise OSError('后端退出通知响应大小超过限制') + yield chunk + + async def aclose(self) -> None: + """ + 关闭底层 HTTP Core 响应流 + + :return: None + """ + + await response.aclose() + + return httpx.Response( + response.status, + headers=response.headers, + stream=ResponseStream(), + extensions=response.extensions, + request=request, + ) + + async def aclose(self) -> None: + """ + 关闭固定地址连接池 + + :return: None + """ + + await self._pool.aclose() + + +BackchannelNotifier = Callable[[str, str], Awaitable[None]] + + +async def send_backchannel_once( + uri: str, + token: str, + notifier: BackchannelNotifier | None, + timeout_seconds: float, + *, + addresses: set[str] | None = None, +) -> None: + """ + 发送单次 Back-Channel 请求并分类永久失败 + + :param uri: 已验证的 Back-Channel URI + :param token: 待发送的 Logout Token + :param notifier: 可选的外部通知器 + :param timeout_seconds: 请求超时秒数 + :param addresses: 已验证的公网 IP 集合 + :return: None + :raises PermanentBackchannelError: 收到不可重试的 HTTP 客户端错误 + :raises OSError: 未提供安全地址或网络传输失败 + """ + + if notifier is not None: + try: + await notifier(uri, token) + except httpx.HTTPStatusError as exc: + _raise_permanent_http_error(exc) + return + + parsed = urlsplit(uri) + hostname = parsed.hostname or '' + if not addresses: + raise OSError('后端退出通知目标不是公网地址') + transport = PinnedHttpxTransport(hostname, addresses) + async with httpx.AsyncClient( + timeout=timeout_seconds, + follow_redirects=False, + transport=transport, + ) as client: + response = await client.post(uri, data={'logout_token': token}) + try: + response.raise_for_status() + except httpx.HTTPStatusError as exc: + _raise_permanent_http_error(exc) + + +def _raise_permanent_http_error(exc: httpx.HTTPStatusError) -> None: + """ + 将不可重试的 HTTP 客户端错误转换为永久失败 + + :param exc: HTTPX 状态码异常 + :return: None + :raises PermanentBackchannelError: 状态码属于不可重试的客户端错误 + :raises httpx.HTTPStatusError: 状态码仍允许调用方重试 + """ + + status = exc.response.status_code + if _HTTP_CLIENT_ERROR_MIN <= status < _HTTP_SERVER_ERROR_MIN and status not in { + _HTTP_REQUEST_TIMEOUT, + _HTTP_TOO_MANY_REQUESTS, + }: + raise PermanentBackchannelError(f'http_{status}', f'后端退出通知被接收方拒绝(HTTP {status})') from None + raise exc diff --git a/ruoyi-fastapi-backend/module_identity/security/client_auth.py b/ruoyi-fastapi-backend/module_identity/security/client_auth.py new file mode 100644 index 000000000..90844ab59 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/client_auth.py @@ -0,0 +1,126 @@ +import secrets +from collections.abc import Callable, Iterable +from typing import Any + +from module_identity.security.principal import OAuthClientPrincipal +from utils.oidc_util import OidcUtil +from utils.pwd_util import PwdUtil + +_MISSING_STATUS = object() + + +class ClientAuthenticationError(ValueError): + """ + 客户端认证失败 + + 协议层应将此异常统一映射为 invalid_client,不能泄漏 Secret 细节 + """ + + +def hash_client_secret(client_secret: str) -> str: + """ + 使用项目密码哈希器保存 Client Secret + + :param client_secret: 仅在创建/轮换时可见的明文 Secret + :return: 可持久化的强哈希 + """ + + if not isinstance(client_secret, str) or not client_secret: + raise ValueError('客户端密钥不能为空') + return PwdUtil.get_password_hash(client_secret) + + +def verify_client_secret(client_secret: str, secret_hash: str) -> bool: + """ + 验证 Client Secret,错误输入统一返回 False + + :param client_secret: 请求携带的明文 Secret + :param secret_hash: 已存储的 Secret 哈希 + :return: Secret 是否匹配 + """ + + if not isinstance(client_secret, str) or not isinstance(secret_hash, str) or not secret_hash: + return False + try: + return bool(PwdUtil.verify_password(client_secret, secret_hash)) + except (ValueError, TypeError, OSError): + return False + + +def authenticate_client( # noqa: PLR0912 + client: Any, + authorization: str | None = None, + *, + client_id: str | None = None, + client_secret: str | None = None, + secret_hashes: Iterable[str] | None = None, + client_lookup: Callable[[str], Any] | None = None, + secret_match_callback: Callable[[str], None] | None = None, +) -> OAuthClientPrincipal: + """ + 按已注册 Client 策略认证 + + 机密 Client 仅接受 Basic;公共 Client 仅接受 form 的 client_id 且不得 + 携带 Secret。``client_lookup`` 仅用于在 Header client_id 与传入对象 + 分离时查找注册记录 + + :param client: 已注册 Client 或支持属性访问的对象 + :param authorization: RFC 7617 Basic Header + :param client_id: 公共 Client 的表单 client_id + :param client_secret: 仅用于拒绝请求体 Secret 的参数 + :param secret_hashes: 可注入的有效 Secret 哈希集合 + :param client_lookup: 按 client_id 查找注册 Client 的回调 + :param secret_match_callback: Secret 匹配时接收其哈希的同步回调 + :return: 已认证的 OAuth Client 主体 + :raises ClientAuthenticationError: Client 策略或凭据不匹配 + """ + + supplied_id = client_id + supplied_secret = client_secret + auth_method = 'none' + if authorization is not None: + if client_secret is not None: + raise ClientAuthenticationError('客户端认证信息无效') + try: + basic_id, basic_secret = OidcUtil.parse_basic_credentials(authorization) + except ValueError as exc: + raise ClientAuthenticationError(str(exc)) from None + if supplied_id is not None and supplied_id != basic_id: + raise ClientAuthenticationError('客户端认证信息无效') + supplied_id, supplied_secret, auth_method = basic_id, basic_secret, 'client_secret_basic' + if supplied_id is None: + supplied_id = OidcUtil.read_field(client, 'client_id') + if client_lookup is not None and supplied_id: + looked_up = client_lookup(supplied_id) + if looked_up is not None: + client = looked_up + registered_id = OidcUtil.read_field(client, 'client_id') + if not registered_id or not supplied_id or not secrets.compare_digest(str(registered_id), str(supplied_id)): + raise ClientAuthenticationError('客户端认证信息无效') + status = OidcUtil.read_field(client, 'status', _MISSING_STATUS) + if status not in ('0', 0, 'active'): + raise ClientAuthenticationError('客户端认证信息无效') + client_type = OidcUtil.read_field(client, 'client_type', 'confidential') + registered_method = OidcUtil.read_field(client, 'token_endpoint_auth_method', 'client_secret_basic') + if client_type == 'public': + if registered_method != 'none' or authorization is not None or supplied_secret is not None: + raise ClientAuthenticationError('客户端认证信息无效') + return OAuthClientPrincipal(str(registered_id), 'public', 'none') + if ( + client_type != 'confidential' + or registered_method != 'client_secret_basic' + or auth_method != 'client_secret_basic' + ): + raise ClientAuthenticationError('客户端认证信息无效') + matched = False + matched_hash: str | None = None + for item in OidcUtil.client_secret_hashes(client, secret_hashes): + current_match = verify_client_secret(supplied_secret or '', item) + if current_match and matched_hash is None: + matched_hash = item + matched = current_match or matched + if not supplied_secret or not matched: + raise ClientAuthenticationError('客户端认证信息无效') + if matched_hash is not None and secret_match_callback is not None: + secret_match_callback(matched_hash) + return OAuthClientPrincipal(str(registered_id), 'confidential', 'client_secret_basic') diff --git a/ruoyi-fastapi-backend/module_identity/security/jwt_profile.py b/ruoyi-fastapi-backend/module_identity/security/jwt_profile.py new file mode 100644 index 000000000..2cb037def --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/jwt_profile.py @@ -0,0 +1,602 @@ +import math +from collections.abc import Mapping, Sequence +from datetime import datetime, timezone +from typing import Any + +import jwt +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey +from jwt.exceptions import PyJWTError + +ACCESS_TOKEN_ALGORITHM = 'RS256' +ACCESS_TOKEN_TYPE = 'at+jwt' +ID_TOKEN_TYPE = 'JWT' +LOGOUT_TOKEN_TYPE = 'logout+jwt' +BACKCHANNEL_LOGOUT_EVENT = 'http://schemas.openid.net/event/backchannel-logout' +ALLOWED_SIGNING_ALGORITHMS = frozenset({ACCESS_TOKEN_ALGORITHM}) + + +class JwtProfileError(ValueError): + """ + JWT Profile 校验失败 + """ + + +def _key_for_kid(verification_keys: object, kid: str) -> object: + """ + 按 Key ID 选择本地验签密钥 + + :param verification_keys: 验签密钥、按 Key ID 索引的映射或查找回调 + :param kid: JWT Header 中的 Key ID + :return: 匹配的验签密钥 + :raises JwtProfileError: 验签密钥缺失或 Key ID 未知 + """ + + if verification_keys is None: + raise JwtProfileError('签名密钥不能为空') + if isinstance(verification_keys, Mapping): + try: + return verification_keys[kid] + except KeyError: + raise JwtProfileError('签名密钥不存在或不可用') from None + if callable(verification_keys): + value = verification_keys(kid) + if value is None: + raise JwtProfileError('签名密钥不存在或不可用') + return value + return verification_keys + + +def _header(token: str) -> dict[str, Any]: + """ + 读取并校验 JWT Header 安全字段 + + :param token: 待读取的 JWT + :return: 通过固定算法与 Header 白名单校验的字段映射 + :raises JwtProfileError: JWT 编码、算法或 Header 字段不合法 + """ + + try: + value = jwt.get_unverified_header(token) + except (PyJWTError, TypeError, ValueError) as exc: + raise JwtProfileError('JWT 编码格式无效') from exc + if value.get('alg') not in ALLOWED_SIGNING_ALGORITHMS: + raise JwtProfileError('不支持当前 JWT 签名算法') + if any(name in value for name in ('crit', 'jku', 'jwk', 'x5u', 'x5c')): + raise JwtProfileError('不支持当前 JWT 头部参数') + if not isinstance(value.get('kid'), str) or not value['kid'].strip(): + raise JwtProfileError('JWT 缺少签名密钥标识 kid') + return value + + +def _numeric_claims(payload: Mapping[str, Any], *, now: float, skew: float, verify_exp: bool = True) -> None: + """ + 校验 JWT NumericDate 类型和时效 + + :param payload: 已验签的 JWT Claims + :param now: 当前 UTC 时间戳 + :param skew: 允许的时钟偏差秒数 + :param verify_exp: 是否校验过期时间,仅退出提示流程可关闭 + :return: None + :raises JwtProfileError: NumericDate 类型或时效不合法 + """ + + for name in ('iat', 'exp', 'nbf', 'auth_time'): + if name not in payload: + if name == 'iat': + raise JwtProfileError('JWT 缺少签发时间 iat') + continue + value = payload[name] + if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value): + raise JwtProfileError(f'JWT 声明 {name} 无效') + if name == 'iat' and value > now + skew: + raise JwtProfileError('JWT 签发时间晚于当前时间') + if name == 'auth_time' and value > now + skew: + raise JwtProfileError('JWT 认证时间 auth_time 晚于当前时间') + if name == 'exp' and verify_exp and now - skew >= value: + raise JwtProfileError('JWT 已过期') + if name == 'nbf' and now + skew < value: + raise JwtProfileError('JWT 尚未生效') + + +def _validate_numeric_types(payload: Mapping[str, Any]) -> None: + """ + 仅校验 NumericDate 类型,不在签发时用当前时间判断有效期 + + :param payload: 待签发的 JWT Claims + :return: None + :raises JwtProfileError: NumericDate 类型不合法 + """ + + for name in ('iat', 'exp', 'nbf', 'auth_time'): + if name in payload and ( + isinstance(payload[name], bool) + or not isinstance(payload[name], (int, float)) + or not math.isfinite(payload[name]) + ): + raise JwtProfileError(f'JWT 声明 {name} 无效') + + +def _audience_matches(actual: object, expected: str | Sequence[str] | None) -> bool: + """ + 判断实际 Audience 是否匹配允许值 + + :param actual: JWT 中的 Audience + :param expected: 允许的单个或多个 Audience + :return: Audience 是否匹配 + """ + + if expected is None: + return True + values = [actual] if isinstance(actual, str) else actual + if not isinstance(values, (list, tuple, set)) or any(not isinstance(item, str) for item in values): + return False + wanted = [expected] if isinstance(expected, str) else list(expected) + + return any(item in wanted for item in values) + + +def _validate_string_claim(payload: Mapping[str, Any], name: str, required: bool = True) -> None: + """ + 校验 JWT 字符串 Claim + + :param payload: JWT Claims + :param name: Claim 名称 + :param required: 是否要求 Claim 必须存在 + :return: None + :raises JwtProfileError: Claim 缺失或不是非空字符串 + """ + + value = payload.get(name) + if value is None and not required: + return + if not isinstance(value, str) or not value.strip(): + raise JwtProfileError(f'JWT 声明 {name} 无效') + + +def _validate_string_list(payload: Mapping[str, Any], name: str, required: bool = True) -> None: + """ + 校验 JWT 字符串列表 Claim + + :param payload: JWT Claims + :param name: Claim 名称 + :param required: 是否要求 Claim 必须存在 + :return: None + :raises JwtProfileError: Claim 缺失或不是非空字符串列表 + """ + + value = payload.get(name) + if value is None and not required: + return + if not isinstance(value, list) or any(not isinstance(item, str) or not item for item in value): + raise JwtProfileError(f'JWT 声明 {name} 无效') + + +def _validate_audience_claim(payload: Mapping[str, Any]) -> None: + """ + 校验 JWT Audience Claim + + :param payload: JWT Claims + :return: None + :raises JwtProfileError: Audience 结构、内容或唯一性不合法 + """ + + aud = payload.get('aud') + if not isinstance(aud, (str, list)): + raise JwtProfileError('JWT 受众声明 aud 无效') + if isinstance(aud, list) and (not aud or any(not isinstance(item, str) or not item for item in aud)): + raise JwtProfileError('JWT 受众声明 aud 无效') + if isinstance(aud, list) and len(set(aud)) != len(aud): + raise JwtProfileError('JWT 受众声明 aud 无效') + if isinstance(aud, str) and not aud: + raise JwtProfileError('JWT 受众声明 aud 无效') + + +def _validate_access_claims(payload: Mapping[str, Any]) -> None: + """ + 校验 Access Token Profile Claims + + :param payload: Access Token Claims + :return: None + :raises JwtProfileError: 必需 Claim 或身份绑定不合法 + """ + + for name in ('iss', 'sub', 'client_id', 'scope'): + _validate_string_claim(payload, name) + _validate_audience_claim(payload) + _validate_string_claim(payload, 'gty') + if payload.get('gty') not in {'authorization_code', 'refresh_token', 'client_credentials'}: + raise JwtProfileError('访问令牌的授权类型无效') + for name in ('exp', 'iat', 'nbf', 'jti'): + if name not in payload: + raise JwtProfileError(f'访问令牌缺少必需声明:{name}') + _validate_string_claim(payload, 'jti') + _validate_access_authorization_context(payload) + if payload['gty'] == 'client_credentials': + if payload.get('sub') != f'client:{payload["client_id"]}' or any( + name in payload for name in ('sid', 'ver', 'auth_time', 'acr', 'amr') + ): + raise JwtProfileError('机器访问令牌的客户端身份绑定无效') + return + for name in ('sid', 'acr'): + _validate_string_claim(payload, name) + _validate_string_list(payload, 'amr') + if ( + 'ver' not in payload + or isinstance(payload['ver'], bool) + or not isinstance(payload['ver'], int) + or payload['ver'] < 1 + ): + raise JwtProfileError('访问令牌的安全版本无效') + if 'auth_time' not in payload: + raise JwtProfileError('访问令牌缺少认证时间 auth_time') + + +def _validate_access_authorization_context(payload: Mapping[str, Any]) -> None: + """ + 校验 Access Token 的可选授权来源 Claims + + :param payload: Access Token Claims,存量令牌可不包含授权来源字段 + :return: None + :raises JwtProfileError: Client 策略版本或 Grant 绑定不合法 + """ + + if 'client_policy_version' not in payload: + if 'grant_id' in payload: + raise JwtProfileError('访问令牌缺少客户端策略版本') + return + version = payload['client_policy_version'] + if isinstance(version, bool) or not isinstance(version, int) or version < 1: + raise JwtProfileError('访问令牌的客户端策略版本无效') + if payload['gty'] == 'client_credentials': + if 'grant_id' in payload: + raise JwtProfileError('机器访问令牌不得携带用户授权记录') + return + if 'grant_id' not in payload: + raise JwtProfileError('访问令牌缺少授权记录绑定') + grant_id = payload['grant_id'] + if grant_id is None: + if payload['gty'] != 'authorization_code' or 'offline_access' in payload['scope'].split(): + raise JwtProfileError('离线访问令牌必须绑定已持久化的授权记录') + elif not isinstance(grant_id, str) or not grant_id.strip(): + raise JwtProfileError('访问令牌的授权记录绑定无效') + + +def _validate_id_claims(payload: Mapping[str, Any]) -> None: + """ + 校验 ID Token Profile Claims + + :param payload: ID Token Claims + :return: None + :raises JwtProfileError: 必需 Claim 不合法 + """ + + for name in ('iss', 'sub', 'sid', 'acr', 'nonce'): + _validate_string_claim(payload, name) + _validate_audience_claim(payload) + _validate_string_list(payload, 'amr') + for name in ('exp', 'iat', 'auth_time'): + if name not in payload: + raise JwtProfileError(f'身份令牌缺少必需声明:{name}') + + +def _validate_logout_claims(payload: Mapping[str, Any]) -> None: + """ + 校验 Logout Token Profile Claims + + :param payload: Logout Token Claims + :return: None + :raises JwtProfileError: 必需 Claim、事件或主体绑定不合法 + """ + + for name in ('iss', 'jti'): + _validate_string_claim(payload, name) + _validate_audience_claim(payload) + for name in ('sid', 'sub'): + _validate_string_claim(payload, name, required=False) + for name in ('iat', 'exp'): + if name not in payload: + raise JwtProfileError(f'退出通知令牌缺少必需声明:{name}') + if not payload.get('sid') and not payload.get('sub'): + raise JwtProfileError('退出通知令牌必须包含 sid 或 sub') + events = payload.get('events') + if not isinstance(events, dict) or BACKCHANNEL_LOGOUT_EVENT not in events or events[BACKCHANNEL_LOGOUT_EVENT] != {}: + raise JwtProfileError('退出通知令牌的事件声明无效') + if 'nonce' in payload: + raise JwtProfileError('退出通知令牌不得包含 nonce') + + +def _validate_profile_claims(profile: str, payload: Mapping[str, Any]) -> None: + """ + 按 JWT Profile 分派 Claims 校验 + + :param profile: access、id 或 logout Profile + :param payload: JWT Claims + :return: None + :raises JwtProfileError: Profile Claims 不合法 + """ + + if profile == 'access': + _validate_access_claims(payload) + elif profile == 'id': + _validate_id_claims(payload) + else: + _validate_logout_claims(payload) + + +def _decode( + token: str, + verification_keys: object, + issuer: str, + audience: str | Sequence[str] | None, + *, + profile: str, + expected_type: str, + required_claims: set[str], + clock_skew: float = 60, + verification_key: object | None = None, + allow_expired_hint: bool = False, +) -> dict[str, Any]: + """ + 按固定 Profile 验签并校验 JWT + + :param token: 待验证的 JWT + :param verification_keys: 验签密钥、映射或查找回调 + :param issuer: 必须精确匹配的 Issuer + :param audience: 允许的 Audience + :param profile: access、id 或 logout Profile + :param expected_type: 预期的 JWT typ + :param required_claims: PyJWT 必须检查的 Claim 集合 + :param clock_skew: 允许的时钟偏差秒数 + :param verification_key: 可选的单一验签密钥 + :param allow_expired_hint: 是否允许退出提示使用过期ID Token + :return: 已验签并通过 Profile 校验的 Claims + :raises JwtProfileError: Header、签名、Issuer、Audience 或 Claims 不合法 + """ + + if clock_skew < 0: + raise JwtProfileError('时钟容差不得为负数') + headers = _header(token) + if headers.get('typ') != expected_type: + raise JwtProfileError('JWT 用途类型不匹配') + key_source = verification_key if verification_key is not None else verification_keys + key = _key_for_kid(key_source, str(headers['kid'])) + try: + payload = jwt.decode( + token, + key, + algorithms=[ACCESS_TOKEN_ALGORITHM], + leeway=clock_skew, + options={ + 'require': sorted(required_claims), + 'verify_aud': False, + 'verify_iss': False, + 'verify_exp': not allow_expired_hint, + }, + ) + except PyJWTError as exc: + raise JwtProfileError('JWT 签名或标准声明校验失败') from exc + if payload.get('iss') != issuer or not _audience_matches(payload.get('aud'), audience): + raise JwtProfileError('JWT 签发者或受众不匹配') + _validate_profile_claims(profile, payload) + now = datetime.now(timezone.utc).timestamp() + _numeric_claims(payload, now=now, skew=clock_skew, verify_exp=not allow_expired_hint) + + return dict(payload) + + +def _encode(claims: Mapping[str, Any], signing_key: RSAPrivateKey, kid: str, expected_type: str, profile: str) -> str: + """ + 按固定 Profile Header 签发 JWT + + :param claims: 待签发且已包含协议必需字段的 Claims + :param signing_key: RSA 私钥 + :param kid: 签名密钥标识 + :param expected_type: 固定 JWT typ + :param profile: access、id 或 logout + :return: 签名后的 JWT 字符串 + :raises JwtProfileError: Key ID、Claims 或编码过程不合法 + """ + + payload = dict(claims) + if not isinstance(kid, str) or not kid.strip(): + raise JwtProfileError('JWT 缺少签名密钥标识 kid') + _validate_profile_claims(profile, payload) + _validate_numeric_types(payload) + headers = {'alg': ACCESS_TOKEN_ALGORITHM, 'kid': kid, 'typ': expected_type} + try: + return jwt.encode(payload, signing_key, algorithm=ACCESS_TOKEN_ALGORITHM, headers=headers) + except PyJWTError as exc: + raise JwtProfileError('JWT 签发编码失败') from exc + + +def encode_access_token(claims: Mapping[str, Any], signing_key: RSAPrivateKey, kid: str) -> str: + """ + 签发 JWT Access Token + + :param claims: Access Token Claims + :param signing_key: RSA 私钥 + :param kid: 签名密钥标识 + :return: typ 为 at+jwt 的 JWT + :raises JwtProfileError: Claims 或签名参数不合法 + """ + + return _encode(claims, signing_key, kid, ACCESS_TOKEN_TYPE, 'access') + + +def decode_access_token( + token: str, + verification_keys: object | None = None, + issuer: str | None = None, + audience: str | Sequence[str] | None = None, + *, + clock_skew: float = 60, + verification_key: RSAPublicKey | None = None, +) -> dict[str, Any]: + """ + 验证 JWT Access Token Profile + + :param token: 待验证 JWT + :param verification_keys: 按 kid 索引的公钥、单一公钥或查找回调 + :param issuer: 必须精确匹配的 issuer + :param audience: 允许的 audience + :param clock_skew: 允许的时钟偏差秒数 + :param verification_key: 可选的单一 RSA 公钥 + :return: 已验签的 Access Token Claims + :raises JwtProfileError: Access Token 不符合固定 Profile + """ + + if issuer is None: + raise JwtProfileError('签发者地址 issuer 不能为空') + return _decode( + token, + verification_keys, + issuer, + audience, + profile='access', + expected_type=ACCESS_TOKEN_TYPE, + required_claims={'iss', 'sub', 'aud', 'exp', 'iat', 'nbf', 'jti', 'client_id', 'scope', 'gty'}, + clock_skew=clock_skew, + verification_key=verification_key, + ) + + +def encode_id_token(claims: Mapping[str, Any], signing_key: RSAPrivateKey, kid: str) -> str: + """ + 签发 OIDC ID Token + + :param claims: ID Token Claims + :param signing_key: RSA 私钥 + :param kid: 签名密钥标识 + :return: typ 为 JWT 的 ID Token + :raises JwtProfileError: Claims 或签名参数不合法 + """ + + return _encode(claims, signing_key, kid, ID_TOKEN_TYPE, 'id') + + +def decode_id_token( + token: str, + verification_keys: object | None = None, + issuer: str | None = None, + audience: str | Sequence[str] | None = None, + *, + nonce: str | None = None, + clock_skew: float = 60, + verification_key: RSAPublicKey | None = None, +) -> dict[str, Any]: + """ + 验证 OIDC ID Token 并可校验 nonce + + :param token: 待验证 ID Token + :param verification_keys: 按 kid 索引的公钥、单一公钥或查找回调 + :param issuer: 必须精确匹配的 issuer + :param audience: OIDC Client ID + :param nonce: 可选的原始授权 nonce + :param clock_skew: 允许的时钟偏差秒数 + :param verification_key: 可选的单一 RSA 公钥 + :return: 已验签的 ID Token Claims + :raises JwtProfileError: ID Token 不符合固定 Profile 或 nonce 不匹配 + """ + + if issuer is None: + raise JwtProfileError('签发者地址 issuer 不能为空') + claims = _decode( + token, + verification_keys, + issuer, + audience, + profile='id', + expected_type=ID_TOKEN_TYPE, + required_claims={'iss', 'sub', 'aud', 'exp', 'iat', 'auth_time', 'nonce', 'sid', 'acr', 'amr'}, + clock_skew=clock_skew, + verification_key=verification_key, + ) + if nonce is not None and claims.get('nonce') != nonce: + raise JwtProfileError('身份令牌的 nonce 与认证请求不匹配') + return claims + + +def decode_id_token_hint( + token: str, + *, + verification_key: RSAPublicKey, + issuer: str, + audience: str, + clock_skew: float = 60, +) -> dict[str, Any]: + """ + 校验退出请求中的ID Token提示 + + 仅退出流程允许使用过期ID Token,调用方仍需将其绑定到已知会话。 + + :param token: ID Token字符串 + :param verification_key: 签名验证公钥 + :param issuer: 预期签发方 + :param audience: 预期接收方 + :param clock_skew: 允许的时钟偏差,单位为秒 + :return: 已校验的ID Token声明 + """ + + return _decode( + token, + None, + issuer, + audience, + profile='id', + expected_type=ID_TOKEN_TYPE, + required_claims={'iss', 'sub', 'aud', 'exp', 'iat', 'auth_time', 'nonce', 'sid', 'acr', 'amr'}, + verification_key=verification_key, + clock_skew=clock_skew, + allow_expired_hint=True, + ) + + +def encode_logout_token(claims: Mapping[str, Any], signing_key: RSAPrivateKey, kid: str) -> str: + """ + 签发 Back-Channel Logout Token + + :param claims: Logout Token Claims + :param signing_key: RSA 私钥 + :param kid: 签名密钥标识 + :return: typ 为 logout+jwt 的 JWT + :raises JwtProfileError: Claims 或签名参数不合法 + """ + + return _encode(claims, signing_key, kid, LOGOUT_TOKEN_TYPE, 'logout') + + +def decode_logout_token( + token: str, + verification_keys: object | None = None, + issuer: str | None = None, + audience: str | Sequence[str] | None = None, + *, + clock_skew: float = 60, + verification_key: RSAPublicKey | None = None, +) -> dict[str, Any]: + """ + 验证 Back-Channel Logout Token + + :param token: 待验证 Logout Token + :param verification_keys: 按 kid 索引的公钥、单一公钥或查找回调 + :param issuer: 必须精确匹配的 issuer + :param audience: 目标 Client ID + :param clock_skew: 允许的时钟偏差秒数 + :param verification_key: 可选的单一 RSA 公钥 + :return: 已验签的 Logout Token Claims + :raises JwtProfileError: Logout Token 不符合固定 Profile + """ + + if issuer is None: + raise JwtProfileError('签发者地址 issuer 不能为空') + return _decode( + token, + verification_keys, + issuer, + audience, + profile='logout', + expected_type=LOGOUT_TOKEN_TYPE, + required_claims={'iss', 'aud', 'iat', 'exp', 'jti', 'events'}, + clock_skew=clock_skew, + verification_key=verification_key, + ) diff --git a/ruoyi-fastapi-backend/module_identity/security/opaque_token.py b/ruoyi-fastapi-backend/module_identity/security/opaque_token.py new file mode 100644 index 000000000..6b854a3d8 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/opaque_token.py @@ -0,0 +1,174 @@ +import base64 +import hmac +import re +import secrets +import uuid +from dataclasses import dataclass + +from utils.oidc_util import OidcUtil + + +class OpaqueTokenError(ValueError): + """ + 不透明令牌结构、类型或 Secret 不符合安全约束 + """ + + +@dataclass(frozen=True, slots=True) +class ParsedOpaqueToken: + """ + 已解析但尚未证明有效的不透明令牌 + + 解析结果只描述格式,数据库摘要校验成功后才具有凭据语义 + """ + + prefix: str + token_id: str + secret: str + + @property + def value(self) -> str: + """ + 返回完整的不透明令牌文本 + + :return: 带类型前缀的完整令牌 + """ + + return f'{self.prefix}.{self.token_id}.{self.secret}' + + +_PREFIXES = frozenset({'ac1', 'rt1', 'ss1'}) +_TOKEN_ID_RE = re.compile(r'^[A-Za-z0-9_-]{8,128}$') +_SECRET_RE = re.compile(r'^[A-Za-z0-9_-]{32,256}$') +_SECRET_BYTES = 32 +_MIN_PEPPER_BYTES = 32 +_OPAQUE_PART_COUNT = 3 + + +def generate_opaque_token(prefix: str, token_id: str | None = None) -> str: + """ + 生成带类型前缀和 256 bit 随机 Secret 的不透明令牌 + + :param prefix: ac1、rt1 或 ss1 + :param token_id: 可选的外部 ID + :return: 完整不透明令牌 + """ + + if prefix not in _PREFIXES: + raise OpaqueTokenError(f'不支持的不透明令牌前缀:{prefix}') + identifier = token_id or str(uuid.uuid4()) + if not _TOKEN_ID_RE.fullmatch(identifier): + raise OpaqueTokenError('不透明令牌标识无效') + return f'{prefix}.{identifier}.{OidcUtil.base64url_encode(secrets.token_bytes(_SECRET_BYTES))}' + + +def generate_authorization_code(code_id: str | None = None) -> str: + """ + 生成授权码 + + :param code_id: 可选的授权码 ID + :return: ac1 前缀的不透明授权码 + """ + + return generate_opaque_token('ac1', code_id) + + +def generate_refresh_token(token_id: str | None = None) -> str: + """ + 生成 Refresh Token + + :param token_id: 可选的 Refresh Token ID + :return: rt1 前缀的不透明 Refresh Token + """ + + return generate_opaque_token('rt1', token_id) + + +def generate_sso_cookie(sid: str | None = None) -> str: + """ + 生成 SSO Cookie Secret + + :param sid: 可选的 SSO Session ID + :return: ss1 前缀的 Cookie 值 + """ + + return generate_opaque_token('ss1', sid) + + +def parse_opaque_token(token: str, expected_prefix: str | None = None) -> ParsedOpaqueToken: + """ + 解析令牌结构,不把 token_id 当作 Secret 证明 + + :param token: 待解析的不透明令牌 + :param expected_prefix: 可选的预期类型前缀 + :return: 结构化令牌 + :raises OpaqueTokenError: 令牌结构、类型或 Secret 不合法 + """ + + if not isinstance(token, str): + raise OpaqueTokenError('不透明令牌必须为字符串') + parts = token.split('.') + if len(parts) != _OPAQUE_PART_COUNT: + raise OpaqueTokenError('不透明令牌格式无效') + prefix, token_id, secret = parts + if prefix not in _PREFIXES or (expected_prefix is not None and prefix != expected_prefix): + raise OpaqueTokenError('不透明令牌类型不匹配') + if not _TOKEN_ID_RE.fullmatch(token_id): + raise OpaqueTokenError('不透明令牌标识无效') + if not _SECRET_RE.fullmatch(secret): + raise OpaqueTokenError('不透明令牌密钥无效') + try: + decoded = base64.urlsafe_b64decode(secret + '=' * (-len(secret) % 4)) + except (ValueError, UnicodeError): + raise OpaqueTokenError('不透明令牌密钥无效') from None + if len(decoded) != _SECRET_BYTES: + raise OpaqueTokenError('不透明令牌密钥必须为 256 位') + return ParsedOpaqueToken(prefix, token_id, secret) + + +def _pepper_bytes(pepper: str | bytes) -> bytes: + """ + 校验并转换令牌摘要 Pepper + + :param pepper: 字符串或字节形式的摘要 Pepper + :return: 长度满足安全约束的 Pepper 字节 + :raises ValueError: Pepper 类型或长度不符合安全约束 + """ + + value = pepper.encode('utf-8') if isinstance(pepper, str) else pepper + if not isinstance(value, bytes) or len(value) < _MIN_PEPPER_BYTES: + raise ValueError('身份令牌摘要密钥至少需要 256 位') + return value + + +def token_digest(token: str, pepper: str | bytes) -> str: + """ + 返回完整不透明令牌的 HMAC-SHA256 十六进制摘要 + + :param token: 完整 typed opaque token + :param pepper: 独立的摘要 Pepper + :return: 64 位十六进制摘要 + """ + + return OidcUtil.hmac_sha256(token.encode('utf-8'), _pepper_bytes(pepper)) + + +def verify_token_digest(token: str, expected_digest: str, pepper: str | bytes) -> bool: + """ + 恒定时间校验摘要,畸形 token 或 digest 时拒绝验证 + + :param token: 完整 typed opaque token + :param expected_digest: 数据库存储的摘要 + :param pepper: 独立的摘要 Pepper + :return: 摘要是否匹配完整令牌 + """ + + if not isinstance(expected_digest, str) or not re.fullmatch(r'[0-9a-fA-F]{64}', expected_digest): + return False + try: + # 先验证完整 typed token;token_id 本身不能成为可摘要认证凭据 + parse_opaque_token(token) + actual = token_digest(token, pepper) + except (TypeError, ValueError, UnicodeError): + return False + return hmac.compare_digest(actual, expected_digest.lower()) diff --git a/ruoyi-fastapi-backend/module_identity/security/pkce.py b/ruoyi-fastapi-backend/module_identity/security/pkce.py new file mode 100644 index 000000000..185781aff --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/pkce.py @@ -0,0 +1,84 @@ +import hashlib +import hmac +import re +import secrets + +from utils.oidc_util import OidcUtil + + +class PkceError(ValueError): + """ + PKCE 参数格式错误或使用了不支持的算法 + """ + + +_MIN_VERIFIER_LENGTH = 43 +_MAX_VERIFIER_LENGTH = 128 +_VERIFIER_RE = re.compile(r'^[A-Za-z0-9._~-]+$') + + +def validate_code_verifier(code_verifier: str) -> str: + """ + 校验并返回 RFC 7636 code_verifier + + :param code_verifier: 客户端生成的 PKCE verifier + :return: 原样返回已校验 verifier + :raises PkceError: 长度、字符集或类型不合法 + """ + + if not isinstance(code_verifier, str): + raise PkceError('code_verifier 必须为字符串') + if not _MIN_VERIFIER_LENGTH <= len(code_verifier) <= _MAX_VERIFIER_LENGTH: + raise PkceError('code_verifier 必须包含 43 至 128 个字符') + if not _VERIFIER_RE.fullmatch(code_verifier): + raise PkceError('code_verifier 包含 RFC 7636 不允许的字符') + return code_verifier + + +def generate_code_verifier(length: int = 64) -> str: + """ + 生成至少 256 bit 熵的 code_verifier + + :param length: verifier 长度,范围为 43 至 128 + :return: 随机 code_verifier + """ + + if not _MIN_VERIFIER_LENGTH <= length <= _MAX_VERIFIER_LENGTH: + raise ValueError('code_verifier 长度必须为 43 至 128 个字符') + # token_urlsafe 的结果可能略长;截断仍保留足够熵且满足 RFC 字符集 + return validate_code_verifier(secrets.token_urlsafe(length)[:length]) + + +def generate_code_challenge(code_verifier: str) -> str: + """ + 计算未填充的 base64url(SHA-256(verifier)) + + :param code_verifier: 已校验的 PKCE verifier + :return: S256 code_challenge + """ + + verifier = validate_code_verifier(code_verifier) + + return OidcUtil.base64url_encode(hashlib.sha256(verifier.encode('ascii')).digest()) + + +def verify_code_challenge(code_verifier: str, code_challenge: str, method: str = 'S256') -> bool: + """ + 以恒定时间比较 challenge,畸形值统一返回 False + + :param code_verifier: Token Endpoint 提交的 verifier + :param code_challenge: Authorization Endpoint 保存的 challenge + :param method: PKCE 方法,仅支持 S256 + :return: verifier 是否匹配 challenge + :raises PkceError: 使用不支持的 PKCE 方法 + """ + + if method != 'S256': + raise PkceError('仅支持 PKCE S256 算法') + try: + if not OidcUtil.is_s256_challenge(code_challenge): + return False + expected = generate_code_challenge(code_verifier) + except (PkceError, TypeError, UnicodeError): + return False + return hmac.compare_digest(expected, code_challenge) diff --git a/ruoyi-fastapi-backend/module_identity/security/principal.py b/ruoyi-fastapi-backend/module_identity/security/principal.py new file mode 100644 index 000000000..0cdaa267f --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/principal.py @@ -0,0 +1,58 @@ +from dataclasses import dataclass, field + + +@dataclass(frozen=True, slots=True) +class OidcUserPrincipal: + """ + OIDC 用户主体,不承载可变数据库 user_id + """ + + subject_id: str + client_id: str + scopes: frozenset[str] = field(default_factory=frozenset) + audience: tuple[str, ...] = () + sid: str | None = None + auth_version: int | None = None + + @property + def sub(self) -> str: + """ + 返回稳定的 OIDC Subject + + :return: OIDC Subject + """ + + return self.subject_id + + +@dataclass(frozen=True, slots=True) +class OAuthClientPrincipal: + """ + 已完成 Client Authentication 的外部应用主体 + """ + + client_id: str + client_type: str + auth_method: str | None = None + + +@dataclass(frozen=True, slots=True) +class MachinePrincipal: + """ + 明确的 client_credentials 身份,绝不伪装为用户 + """ + + client_id: str + subject: str | None = None + scopes: frozenset[str] = field(default_factory=frozenset) + audience: tuple[str, ...] = () + + @property + def is_machine(self) -> bool: + """ + 标识该主体来自 client_credentials 机器身份 + + :return: 固定返回 True + """ + + return True diff --git a/ruoyi-fastapi-backend/module_identity/security/uri_validator.py b/ruoyi-fastapi-backend/module_identity/security/uri_validator.py new file mode 100644 index 000000000..0290b3a7f --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/security/uri_validator.py @@ -0,0 +1,54 @@ +import asyncio +import socket + +from utils.oidc_util import OidcUtil + + +async def public_dns_only(hostname: str, port: int) -> bool: + """ + 解析目标主机全部 A/AAAA 记录并拒绝任一受限地址 + + :param hostname: 已从 URI 解析出的主机名 + :param port: 目标 TCP 端口 + :return: 解析成功且全部地址为公网地址时返回 ``True`` + """ + + return bool(await public_dns_addresses(hostname, port)) + + +async def public_dns_addresses(hostname: str, port: int) -> set[str]: + """ + 解析并返回本次连接允许固定使用的公网地址集合 + + :param hostname: 已解析的主机名 + :param port: 目标 TCP 端口 + :return: 全部解析结果均安全时的地址集合,失败返回空集合 + """ + + loop = asyncio.get_running_loop() + try: + records = await loop.run_in_executor( + None, + lambda: socket.getaddrinfo(hostname, port, type=socket.SOCK_STREAM), + ) + except (OSError, socket.gaierror, ValueError): + return set() + addresses = {item[4][0] for item in records if item[4]} + + return addresses if addresses and all(OidcUtil.is_public_ip(address) for address in addresses) else set() + + +async def is_safe_backchannel_uri(uri: object) -> bool: + """ + 校验 Back-Channel URI 结构并执行 DNS 公网重解析 + + :param uri: 待校验的注册或发送 URI + :return: 仅 HTTPS、无用户信息、查询、片段且 DNS 目标全部公网时返回 ``True`` + """ + + parsed = OidcUtil.parse_backchannel_uri(uri) + if parsed is None: + return False + hostname, port = parsed + + return bool(await public_dns_addresses(hostname, port)) diff --git a/ruoyi-fastapi-backend/module_identity/service/audit_service.py b/ruoyi-fastapi-backend/module_identity/service/audit_service.py new file mode 100644 index 000000000..ddb8e7ca5 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/audit_service.py @@ -0,0 +1,321 @@ +from collections.abc import Awaitable, Callable, Mapping +from typing import Any + +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker + +from common.constant import OidcAuditEvent +from config.database import DataSourceRegistry +from exceptions.exception import OidcInteractionException +from module_identity.dao.oauth_audit_dao import OAuthAuditDao +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.vo.oauth_session_vo import AuditModel, AuditPageQueryModel +from utils.common_util import export_list2excel +from utils.oidc_util import OidcUtil + + +class AuditService: + """ + 认证审计模块服务层 + """ + + EVENT_FIELDS = frozenset( + { + 'trace_id', + 'event_type', + 'result', + 'risk_level', + 'client_id', + 'resource_id', + 'user_id', + 'subject_id', + 'sid', + 'grant_id', + 'token_id', + 'ip_address', + 'user_agent', + 'failure_code', + 'detail', + 'create_time', + } + ) + HIGH_RISK_EVENTS = frozenset( + { + OidcAuditEvent.IDENTITY_SUBJECT_MISSING, + OidcAuditEvent.AUTHORIZATION_CODE_REUSED, + OidcAuditEvent.REFRESH_REUSE_DETECTED, + OidcAuditEvent.CLIENT_SECRET_ROTATED, + OidcAuditEvent.SIGNING_KEY_ROTATED, + OidcAuditEvent.SESSION_REVOKED, + OidcAuditEvent.TOKEN_REVOKED, + OidcAuditEvent.TOKEN_FAILED, + OidcAuditEvent.SECURITY_VERSION_CHANGED, + OidcAuditEvent.INVALID_CLIENT, + OidcAuditEvent.RESOURCE_POLICY_CHANGED, + OidcAuditEvent.SCOPE_POLICY_CHANGED, + } + ) + _EVENT_TYPE_MAX_LENGTH = 64 + + @classmethod + def _risk_level(cls, event_type: str, risk_level: str | None) -> str: + """ + 计算事件风险等级 + + :param event_type: 审计事件类型 + :param risk_level: 风险等级 + :return: 审计风险等级 + """ + + if event_type in cls.HIGH_RISK_EVENTS: + return 'high' + return risk_level if risk_level in {'normal', 'medium', 'high', 'critical'} else 'normal' + + @classmethod + def build_event(cls, **fields: Any) -> SysOAuthAuditLog: + """ + 按审计模型字段白名单构建安全事件对象 + + :param fields: 审计字段;白名单以外的字段会被丢弃 + :return: 未提交的审计日志实体 + :raises ValueError: 缺少事件类型或结果时抛出 + """ + + event_type = fields.get('event_type') + result = fields.get('result') + if ( + not isinstance(event_type, str) + or not event_type.strip() + or len(event_type.strip()) > cls._EVENT_TYPE_MAX_LENGTH + ): + raise ValueError('审计事件类型长度必须为 1 至 64 个字符') + if result not in {'success', 'failure'}: + raise ValueError('审计处理结果不能为空') + values = {key: value for key, value in fields.items() if key in cls.EVENT_FIELDS} + values['event_type'] = event_type.strip() + values['result'] = result.strip() + values['risk_level'] = cls._risk_level(values['event_type'], values.get('risk_level')) + if 'detail' in values: + values['detail'] = OidcUtil.sanitize_audit_detail(values['detail']) + for field, limit in ( + ('trace_id', 64), + ('client_id', 64), + ('resource_id', 64), + ('subject_id', 36), + ('sid', 36), + ('grant_id', 36), + ('token_id', 36), + ('ip_address', 128), + ('user_agent', 500), + ('failure_code', 64), + ): + if isinstance(values.get(field), str): + safe_value = OidcUtil.sanitize_audit_detail(values[field]) if field == 'user_agent' else values[field] + values[field] = safe_value[:limit] if isinstance(safe_value, str) else None + return SysOAuthAuditLog(**values) + + @classmethod + async def record( + cls, + db: AsyncSession, + event_type: str, + result: str, + *, + risk_level: str = 'normal', + **fields: Any, + ) -> SysOAuthAuditLog: + """ + 脱敏并追加一条审计事件,不提交调用方事务 + + :param db: 异步数据库会话 + :param event_type: 事件类型 + :param result: 事件结果 + :param risk_level: 调用方建议的风险等级 + :param fields: 其余白名单审计字段 + :return: 已刷新但尚未提交的审计日志实体 + """ + + event = cls.build_event(event_type=event_type, result=result, risk_level=risk_level, **fields) + + return await OAuthAuditDao.append(db, event) + + @classmethod + async def record_independent( + cls, + db: AsyncSession, + event_type: str, + result: str, + *, + risk_level: str = 'normal', + **fields: Any, + ) -> SysOAuthAuditLog: + """ + 在不依赖业务事务的独立会话中提交审计事件 + + :param db: 当前业务会话,用于测试环境复用其绑定引擎 + :param event_type: 事件类型 + :param result: 结果 + :param risk_level: 风险等级 + :param fields: 经过白名单过滤的安全字段 + :return: 已提交的审计日志实体 + """ + + engine = db.info.get('service_engine') or getattr(db, 'bind', None) + if engine is not None: + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as audit_db: + event = await cls.record( + audit_db, + event_type, + result, + risk_level=risk_level, + **fields, + ) + await audit_db.commit() + return event + async with DataSourceRegistry.session() as audit_db: + event = await cls.record( + audit_db, + event_type, + result, + risk_level=risk_level, + **fields, + ) + await audit_db.commit() + return event + + @classmethod + async def append(cls, db: AsyncSession, event: SysOAuthAuditLog | Mapping[str, Any]) -> SysOAuthAuditLog: + """ + 追加事件对象或字段映射,并保留调用方提交边界 + + :param db: 异步数据库会话 + :param event: 审计实体或待构建字段映射 + :return: 已刷新但尚未提交的审计日志实体 + """ + + row = event if isinstance(event, SysOAuthAuditLog) else cls.build_event(**dict(event)) + + return await OAuthAuditDao.append(db, row) + + @staticmethod + def _admin_filters(query: AuditPageQueryModel) -> dict[str, object]: + """ + 构造管理端审计查询条件 + + :param query: 管理端审计分页查询条件 + :return: 传给审计 DAO 的筛选条件 + """ + + return { + 'client_id': query.client_id, + 'user_id': query.user_id, + 'event_type': query.event_type, + 'result': query.result, + 'risk_level': query.risk_level, + 'start_time': query.start_time, + 'end_time': query.end_time, + } + + @staticmethod + def _admin_model(row: SysOAuthAuditLog) -> dict[str, object]: + """ + 将审计实体投影为脱敏管理端模型 + + :param row: 审计日志实体 + :return: 不含详情敏感字段的管理端模型字典 + """ + + return AuditModel( + audit_id=row.event_id, + event_type=row.event_type, + result=row.result, + risk_level=row.risk_level, + trace_id=getattr(row, 'trace_id', None), + client_id=getattr(row, 'client_id', None), + resource_id=getattr(row, 'resource_id', None), + user_id=getattr(row, 'user_id', None), + subject_id=getattr(row, 'subject_id', None), + sid=getattr(row, 'sid', None), + ip_address=getattr(row, 'ip_address', None), + failure_code=getattr(row, 'failure_code', None), + create_time=row.create_time, + ).model_dump(by_alias=True) + + @classmethod + async def list_admin_page(cls, db: AsyncSession, query: AuditPageQueryModel) -> tuple[list[dict[str, object]], int]: + """ + 查询管理端审计分页并返回脱敏结果 + + :param db: 异步数据库会话 + :param query: 管理端审计分页查询条件 + :return: 脱敏审计行列表及总数 + """ + + params = cls._admin_filters(query) + rows = await OAuthAuditDao.list_admin_page( + db, offset=(query.page_num - 1) * query.page_size, limit=query.page_size, **params + ) + total = await OAuthAuditDao.count_admin(db, **params) + + return [cls._admin_model(row) for row in rows], total + + @classmethod + async def export_admin(cls, db: AsyncSession, query: AuditPageQueryModel) -> bytes: + """ + 导出管理端审计分页范围内的脱敏数据 + + :param db: 异步数据库会话 + :param query: 管理端审计筛选条件 + :return: 脱敏审计 Excel 文件字节 + """ + + rows = await OAuthAuditDao.list_admin_page(db, offset=0, limit=5000, **cls._admin_filters(query)) + + return export_list2excel([cls._admin_model(row) for row in rows]) + + @classmethod + async def record_interaction_failure(cls, db: AsyncSession, event_type: str, **fields: Any) -> None: + """ + 回滚交互事务并独立提交高风险失败审计 + + :param db: 认证交互使用的异步数据库会话 + :param event_type: 失败审计事件类型 + :param fields: 经过白名单过滤的审计字段 + :return: None + :raises OidcInteractionException: 独立审计提交失败时抛出 + """ + + await db.rollback() + try: + await cls.record_independent(db, event_type, 'failure', risk_level='high', **fields) + except Exception as exc: + await db.rollback() + raise OidcInteractionException(error='server_error', status_code=503, message='认证审计服务不可用') from exc + + @staticmethod + def interaction_subject_writer(db: AsyncSession) -> Callable[[int], Awaitable[object]]: + """ + 构造缺失身份主体的独立审计写入器 + + :param db: 认证交互使用的异步数据库会话 + :return: 接收用户 ID 并独立写入审计的异步回调 + """ + + async def write(user_id: int) -> object: + """ + 写入审计记录 + + :param user_id: 本地用户 ID + :return: 已提交的审计日志实体 + """ + + return await AuditService.record_independent( + db, + OidcAuditEvent.IDENTITY_SUBJECT_MISSING, + 'failure', + risk_level='high', + user_id=user_id if isinstance(user_id, int) else None, + failure_code='identity_integrity', + ) + + return write diff --git a/ruoyi-fastapi-backend/module_identity/service/authorization_service.py b/ruoyi-fastapi-backend/module_identity/service/authorization_service.py new file mode 100644 index 000000000..4f187f010 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/authorization_service.py @@ -0,0 +1,1402 @@ +import json +import re +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Any +from urllib.parse import urlsplit + +from pydantic import ValidationError +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException, OidcInteractionException +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysSsoSession +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.protocol_vo import AuthorizeRequest +from module_identity.redis_keys import OidcRedisKey +from module_identity.security.opaque_token import ( + OpaqueTokenError, + generate_authorization_code, + parse_opaque_token, + token_digest, +) +from module_identity.service.audit_service import AuditService +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.interaction_service import InteractionFlowService, InteractionService +from module_identity.service.session_service import SsoSessionError, SsoSessionService +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + +_S256_CHALLENGE = re.compile(r'^[A-Za-z0-9_-]{43}$') +_MAX_STATE_LENGTH = 1024 + + +@dataclass(frozen=True, slots=True) +class ClientSnapshot: + """ + 授权流程使用的 Client 只读标量快照 + """ + + client_pk: int + client_id: str + policy_version: int + grant_types: tuple[str, ...] + response_types: tuple[str, ...] + require_pkce: bool + require_consent: bool + trusted_client: bool + + def __post_init__(self) -> None: + """ + 拒绝可变容器,确保冻结快照不会被原地修改 + + :return: None + :raises TypeError: Grant 或 Response 类型不是不可变元组时抛出 + """ + + if not isinstance(self.grant_types, tuple) or not isinstance(self.response_types, tuple): + raise TypeError('客户端快照的 grant_types 和 response_types 必须为元组') + + +@dataclass(frozen=True, slots=True) +class ScopeSnapshot: + """ + 授权流程使用的 Scope 只读标量快照 + """ + + scope_pk: int + scope_code: str + scope_type: str + resource_pk: int | None + consent_required: bool + + +@dataclass(frozen=True, slots=True) +class ResourceSnapshot: + """ + 授权流程使用的 Resource 只读标量快照 + """ + + resource_pk: int + audience: str + + +@dataclass(frozen=True, slots=True) +class AuthorizationContext: + """ + 已完成协议校验的不可变授权上下文 + + 上下文只保存后续交互和签码所需的白名单字段,不携带未定义查询参数 + """ + + client: ClientSnapshot + redirect_uri: str + scopes: tuple[str, ...] + scope_models: tuple[ScopeSnapshot, ...] + pre_authorized_scopes: frozenset[str] + resource: ResourceSnapshot | None + state: str | None + nonce: str | None + code_challenge: str + code_challenge_method: str + prompt: str | None + max_age: int | None + required_scopes: frozenset[str] = frozenset({'openid'}) + + @property + def requires_consent(self) -> bool: + """ + 判断当前请求是否存在尚未预授权的用户 Scope + + :return: 需要显示同意页时为 True + """ + + consentable = set(self.scopes) - set(self.required_scopes) + if 'consent' in frozenset((self.prompt or '').split()): + return bool(consentable) + if not self.client.require_consent: + return False + return any(scope not in self.pre_authorized_scopes for scope in consentable) + + @property + def resources(self) -> tuple[str, ...]: + """ + 返回授权上下文中的 Resource audience + + :return: 最多一个 audience 的元组 + """ + + return (self.resource.audience,) if self.resource is not None else () + + def to_internal_payload(self, interaction_id: str) -> dict[str, Any]: + """ + 构建仅供服务端 Redis Interaction 使用的完整绑定载荷 + + :param interaction_id: 服务端生成的 Interaction ID + :return: 包含后续签码所需绑定字段且不含 ORM 对象的内部载荷 + """ + + return { + 'interactionId': interaction_id, + 'clientPk': self.client.client_pk, + 'clientId': self.client.client_id, + 'redirectUri': self.redirect_uri, + 'responseType': 'code', + 'scopes': list(self.scopes), + 'resources': list(self.resources), + 'codeChallenge': self.code_challenge, + 'codeChallengeMethod': self.code_challenge_method, + 'maxAge': self.max_age, + 'state': self.state, + 'nonce': self.nonce, + 'prompt': self.prompt, + 'consentRequired': self.requires_consent, + } + + def to_interaction_payload(self, interaction_id: str) -> dict[str, Any]: + """ + 构建交互页面可见的非敏感载荷 + + :param interaction_id: 服务端生成的 Interaction ID + :return: 不包含 state、nonce 和 PKCE challenge 的页面载荷 + """ + + return { + 'interactionId': interaction_id, + 'clientId': self.client.client_id, + 'responseType': 'code', + 'scopes': list(self.scopes), + 'resources': list(self.resources), + 'consentRequired': self.requires_consent, + } + + +@dataclass(frozen=True, slots=True) +class AuthorizationResult: + """ + 授权流程完成后交给 HTTP Controller 的纯领域结果 + """ + + location: str + status_code: int = 303 + + +class AuthorizationService: + """ + 授权模块服务层 + + 校验顺序固定为 Client、精确 Redirect、响应类型、PKCE、Scope 和 Resource, + 只有 Redirect 已验证后才允许协议错误携带 redirect_uri 和 state + """ + + @classmethod + async def _parse_request(cls, db: AsyncSession, raw: dict[str, str]) -> tuple[AuthorizeRequest, str]: + """ + 先验证 Client 和完整 Redirect,再解析其余协议字段 + + :param db: 异步数据库会话 + :param raw: 已拒绝重复键的 Query 参数 + :return: 解析后的授权请求及已验证 Redirect URI + :raises OAuthProtocolException: Client、Redirect 或协议字段无效 + """ + + client_id = raw.get('client_id') + redirect_uri = raw.get('redirect_uri') + if not isinstance(client_id, str) or not isinstance(redirect_uri, str): + raise OAuthProtocolException('invalid_request', 'client_id and redirect_uri are required') + verified_redirect = await cls.verified_redirect(db, client_id, redirect_uri) + if raw.get('response_mode', 'query') != 'query': + state = raw.get('state') if len(raw.get('state', '')) <= _MAX_STATE_LENGTH else None + raise cls._redirect_error( + 'unsupported_response_mode', 'Response mode is not supported', verified_redirect, state + ) + try: + return AuthorizeRequest.model_validate(raw), verified_redirect + except ValidationError as exc: + state = raw.get('state') if len(raw.get('state', '')) <= _MAX_STATE_LENGTH else None + raise OAuthProtocolException( + 'invalid_request', + 'Invalid authorization request', + 400, + redirect_uri=verified_redirect, + state=state, + redirect_uri_verified=True, + issuer=OidcConfig.oidc_issuer, + ) from exc + + @classmethod + async def process_authorization_request( + cls, + db: AsyncSession, + redis: Any, + raw: dict[str, str], + *, + sso_cookie: str | None = None, + ) -> AuthorizationResult: + """ + 执行授权请求业务流程并在服务层完成数据库事务 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: 认证中心 Redis 客户端 + :param raw: Controller 已解析且去重的 Query 参数 + :param sso_cookie: Controller 提取的 SSO Cookie 原文 + :return: 交互页面或 Client 回调的 URL 结果 + :raises OAuthProtocolException: 协议错误,必要时携带已验证 Redirect + """ + + request, verified_redirect = await cls._parse_request(db, raw) + if request.resource and await cls.has_disabled_resource(db, request): + raise OAuthProtocolException('not_found', 'Resource is unavailable', 404) + context = await cls.validate_request(db, request) + coordinator = AfterCommitCoordinator() + try: + session = await cls._load_sso_session(db, redis, sso_cookie, coordinator) + if ( + session is not None + and context.max_age is not None + and cls._requires_reauthentication(session.auth_time, context.max_age) + ): + session = None + grant = None + if session is not None: + if await OAuthAccessPolicyDao.is_blocked(db, session.user_id, context.client.client_pk): + raise OidcInteractionException('pending', '当前用户已被禁止访问此应用', error='access_denied') + grant = await cls.valid_grant(db, session.user_id, context.client.client_pk) + consent_required = context.requires_consent and not cls.consent_is_satisfied(context, grant) + payload = context.to_internal_payload('pending') + if session is not None: + payload.update( + { + 'authenticatedSid': session.sid, + 'userId': session.user_id, + 'subjectId': session.subject_id, + 'authVersion': session.auth_version, + } + ) + payload['consentRequired'] = consent_required + payload['grantId'] = grant.grant_id if grant is not None and not consent_required else None + created = await InteractionService.create(redis, payload, pepper=OidcConfig.oidc_token_hash_pepper) + await AuditService.record( + db, + OidcAuditEvent.AUTHORIZE_REQUESTED, + 'success', + client_id=request.client_id, + detail={'response_type': request.response_type}, + ) + await coordinator.commit(db) + except OidcInteractionException as exc: + await coordinator.rollback(db) + await cls._record_failure_audit( + db, OidcAuditEvent.AUTHORIZE_DENIED, client_id=request.client_id, failure_code=exc.error + ) + raise cls._with_verified_redirect(exc, verified_redirect, request.state) from exc + except Exception as exc: + await coordinator.rollback(db) + raise OAuthProtocolException( + 'server_error', + 'Authorization service is unavailable', + 500, + redirect_uri=verified_redirect, + state=request.state, + redirect_uri_verified=True, + issuer=OidcConfig.oidc_issuer, + ) from exc + if created.initial_status == 'completed': + return await cls._complete_authorization(db, redis, created.interaction_id) + base = ( + OidcConfig.oidc_interaction_consent_url + if created.initial_status == 'awaiting_consent' + else OidcConfig.oidc_interaction_login_url + ) + + return AuthorizationResult(OidcUtil.interaction_url(base, created.interaction_id, created.csrf_token)) + + @staticmethod + async def _load_sso_session( + db: AsyncSession, + redis: Any, + cookie: str | None, + coordinator: AfterCommitCoordinator, + ) -> SysSsoSession | None: + """ + 校验认证中心 SSO Cookie 并更新 Session 空闲期限 + + 仅接受认证中心 SSO Cookie,拒绝 Legacy Token 作为登录态 + + :param db: 异步数据库会话 + :param redis: Redis 客户端 + :param cookie: SSO Cookie 值 + :param coordinator: 提交后副作用协调器 + :return: SSO Session 或 None + """ + + if cookie is None: + return None + try: + return await SsoSessionService.touch( + db, + redis, + cookie, + pepper=OidcConfig.oidc_token_hash_pepper, + now=TimezoneUtil.utc_now(), + coordinator=coordinator, + ) + except SsoSessionError: + return None + + @classmethod + async def _complete_authorization(cls, db: AsyncSession, redis: Any, interaction_id: str) -> AuthorizationResult: + """ + 消费已完成 Interaction 并生成一次性授权码回调结果 + + :param db: 异步数据库会话 + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :return: 授权码重定向结果 + """ + + record = await InteractionService.get_record(redis, interaction_id) + if record.get('status') != 'completed': + raise OidcInteractionException(interaction_id, '认证交互尚未完成', error='interaction_required') + marker = await cls._reserve_completion(redis, interaction_id) + if marker is None: + raise OidcInteractionException( + interaction_id, '认证交互已完成,请勿重复提交', error='invalid_request', status_code=409 + ) + code: str | None = None + try: + client_pk = record['clientPk'] + redirect_uri = record['redirectUri'] + registered_uri = await cls.verified_redirect_for_client(db, client_pk, redirect_uri) + if registered_uri is None: + raise OAuthProtocolException('server_error', 'Validated redirect URI is unavailable', 500) + session = await cls.active_session(db, record.get('authenticatedSid', ''), now=TimezoneUtil.utc_now()) + if session is None: + raise OAuthProtocolException('login_required', 'A current login is required', 400) + grant = await cls._completion_grant(db, record, session) + payload = { + 'clientPk': client_pk, + 'redirectUri': redirect_uri, + 'userId': session.user_id, + 'subjectId': session.subject_id, + 'authVersion': session.auth_version, + 'sid': session.sid, + 'grantId': grant.grant_id, + 'scopes': record['scopes'], + 'resources': record['resources'], + 'nonce': record['nonce'], + 'codeChallenge': record['codeChallenge'], + 'codeChallengeMethod': record['codeChallengeMethod'], + 'authTime': (TimezoneUtil.to_optional_utc(session.auth_time) or TimezoneUtil.utc_now()).isoformat(), + } + code = await AuthorizationCodeService.issue(redis, payload, pepper=OidcConfig.oidc_token_hash_pepper) + await AuditService.record( + db, + OidcAuditEvent.AUTHORIZE_SUCCEEDED, + 'success', + client_id=record.get('clientId'), + user_id=session.user_id, + subject_id=str(session.subject_id), + sid=session.sid, + grant_id=grant.grant_id, + ) + await db.commit() + return AuthorizationResult(cls._success_url(registered_uri, code, record.get('state'))) + except Exception as exc: + if code is not None: + try: + await AuthorizationCodeService.invalidate(redis, code) + except Exception: + pass + await cls._best_effort_delete(redis, marker) + await db.rollback() + if isinstance(exc, OAuthProtocolException): + raise + raise OAuthProtocolException('server_error', 'Authorization code could not be completed', 500) from exc + + @classmethod + async def _completion_grant(cls, db: AsyncSession, record: dict[str, Any], session: SysSsoSession) -> SysOAuthGrant: + """ + 重新检查访问策略并为免确认流程补齐授权记录 + + :param db: 异步数据库会话 + :param record: 已完成的交互记录 + :param session: 当前有效 SSO Session + :return: 本次授权码绑定的有效 Grant + """ + + client_pk = record['clientPk'] + await OAuthAccessPolicyDao.lock_client(db, client_pk) + if await OAuthAccessPolicyDao.is_blocked(db, session.user_id, client_pk, for_update=True): + raise OAuthProtocolException('access_denied', 'Access to this application is blocked') + context = await cls.validate_request( + db, + AuthorizeRequest( + response_type='code', + client_id=record['clientId'], + redirect_uri=record['redirectUri'], + scope=' '.join(record['scopes']), + resource=record['resources'][0] if record['resources'] else None, + nonce=record['nonce'], + code_challenge=record['codeChallenge'], + code_challenge_method=record['codeChallengeMethod'], + ), + ) + grant_id = record.get('grantId') + if grant_id is not None: + grant = await OAuthGrantDao.get_by_grant_id_for_update(db, grant_id, refresh=True) + if ( + grant is None + or grant.status != 'active' + or grant.user_id != session.user_id + or grant.subject_id != session.subject_id + or grant.client_pk != client_pk + or grant.client_policy_version != context.client.policy_version + or ( + grant.expires_at is not None + and (TimezoneUtil.to_optional_utc(grant.expires_at) or TimezoneUtil.utc_now()) + <= TimezoneUtil.utc_now() + ) + or not set(context.scopes).issubset(set(grant.granted_scopes or [])) + or not set(context.resources).issubset(set(grant.granted_resources or [])) + ): + raise OAuthProtocolException('access_denied', 'The authorization grant is no longer valid') + return grant + if record.get('consentRequired') or context.requires_consent: + raise OAuthProtocolException('consent_required', 'A new authorization decision is required') + return await OAuthGrantDao.merge_active_grant( + db, + session.user_id, + session.subject_id, + client_pk, + list(context.scopes), + list(context.resources), + context.client.policy_version, + remember_consent=False, + ) + + @staticmethod + async def _reserve_completion(redis: Any, interaction_id: str) -> str | None: + """ + 以 Interaction 剩余 TTL 保留一次终态完成权 + + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :return: 完成标记 Key 或 None + """ + + ttl = await redis.ttl(OidcRedisKey.interaction(interaction_id)) + if not isinstance(ttl, int) or ttl <= 0: + return None + marker = OidcRedisKey.interaction(f'{interaction_id}-completion') + reserved = await redis.set(marker, 'reserved', ex=ttl, nx=True) + + return marker if reserved else None + + @staticmethod + async def _best_effort_delete(redis: Any, marker: str) -> None: + """ + 尽力删除授权完成 marker + + :param redis: Redis 客户端 + :param marker: 完成标记 Key + :return: None + """ + + try: + await redis.delete(marker) + except Exception: + return + + @staticmethod + def _requires_reauthentication(auth_time: datetime | None, max_age: int) -> bool: + """ + 判断 SSO Session 是否超过授权请求的 max_age + + :param auth_time: 认证时间 + :param max_age: 最大认证时效 + :return: 是否需要重新认证 + """ + + if auth_time is None: + return True + current = TimezoneUtil.to_utc(auth_time) + + return (TimezoneUtil.utc_now() - current).total_seconds() > max_age + + @staticmethod + def _success_url(redirect_uri: str, code: str, state: str | None) -> str: + """ + 构建已验证 Client Redirect 的授权码回调地址 + + :param redirect_uri: 已验证的重定向 URI + :param code: Authorization Code + :param state: 协议 state 值 + :return: 授权码回调 URL + """ + + parsed = urlsplit(redirect_uri) + if parsed.fragment or not parsed.scheme or not parsed.netloc: + raise OAuthProtocolException('server_error', 'Validated redirect URI is invalid', 500) + fields = [('code', code)] + if state is not None: + fields.append(('state', state)) + fields.append(('iss', OidcConfig.oidc_issuer)) + return OidcUtil.replace_query_parameters( + redirect_uri, fields, {'code', 'error', 'error_description', 'error_uri', 'iss', 'state'}, fragment='' + ) + + @staticmethod + def _with_verified_redirect( + exc: OidcInteractionException, redirect_uri: str, state: str | None + ) -> OAuthProtocolException: + """ + 把 Interaction 错误绑定到已验证 Redirect + + :param exc: 协议异常 + :param redirect_uri: 待核验的请求重定向 URI + :param state: 协议 state 值 + :return: 绑定重定向的协议异常 + """ + + return OAuthProtocolException( + exc.error, + exc.message, + exc.status_code, + redirect_uri=redirect_uri, + state=state, + redirect_uri_verified=True, + issuer=OidcConfig.oidc_issuer, + ) + + @staticmethod + async def _record_failure_audit(db: AsyncSession, event_type: str, **fields: Any) -> None: + """ + 尝试记录授权失败审计,审计异常不改变原协议错误响应 + + :param db: 异步数据库会话 + :param event_type: 审计事件类型 + :param fields: 失败审计字段映射 + :return: None + """ + + try: + await AuditService.record_independent(db, event_type, 'failure', risk_level='high', **fields) + except Exception as exc: + raise OAuthProtocolException('server_error', 'Authorization audit service is unavailable', 503) from exc + + @staticmethod + async def verified_redirect(db: AsyncSession, client_id: str, redirect_uri: str) -> str: + """ + 返回已注册的精确 Redirect URI + + :param db: 异步数据库会话 + :param client_id: OAuth Client 标识 + :param redirect_uri: 已验证的重定向 URI + :return: 已注册的精确重定向 URI + """ + + client = await OAuthClientDao.get_by_client_id(db, client_id, active_only=True) + if client is None: + raise OAuthProtocolException('unauthorized_client', 'Client is not registered') + registered = await OAuthClientDao.find_exact_uri(db, client.client_pk, 'redirect', redirect_uri) + if registered is None: + raise OAuthProtocolException('invalid_request', 'redirect_uri is not registered') + return registered.uri + + @staticmethod + async def verified_redirect_for_client(db: AsyncSession, client_pk: int, redirect_uri: str) -> str | None: + """ + 按内部 Client 主键确认已注册的 Redirect URI + + :param db: 异步数据库会话 + :param client_pk: OAuth Client 主键 + :param redirect_uri: 已验证的重定向 URI + :return: 匹配的重定向 URI 或 None + """ + + registered = await OAuthClientDao.find_exact_uri(db, client_pk, 'redirect', redirect_uri) + + return registered.uri if registered is not None else None + + @staticmethod + async def valid_grant(db: AsyncSession, user_id: int, client_pk: int) -> SysOAuthGrant | None: + """ + 读取用户对 Client 的有效 Grant + + :param db: 异步数据库会话 + :param user_id: 本地用户 ID + :param client_pk: OAuth Client 主键 + :return: 有效 Grant 或 None + """ + + return await OAuthGrantDao.get_valid_for_user_client(db, user_id, client_pk) + + @staticmethod + async def has_disabled_resource(db: AsyncSession, request: AuthorizeRequest) -> bool: + """ + 判断请求 Resource 是否为该 Client 已绑定的停用资源 + + :param db: 异步数据库会话 + :param request: 当前 HTTP 请求 + :return: 是否存在禁用 Resource + """ + + return bool(request.resource) and await OAuthClientDao.has_disabled_bound_resource( + db, request.client_id, request.resource + ) + + @staticmethod + async def active_session(db: AsyncSession, sid: str, now: datetime | None = None) -> SysSsoSession | None: + """ + 读取当前有效 SSO Session + + :param db: 异步数据库会话 + :param sid: SSO Session 标识 + :param now: 当前时间 + :return: 活动 SSO Session 或 None + """ + + return await SsoSessionDao.get_active( + db, sid, now=TimezoneUtil.to_utc(now) if now is not None else TimezoneUtil.utc_now() + ) + + @classmethod + async def validate_request(cls, db: AsyncSession, request: AuthorizeRequest) -> AuthorizationContext: # noqa: PLR0912, PLR0915 + """ + 校验授权请求并构建不可变授权上下文 + + :param db: 异步数据库会话 + :param request: 已完成基础字段校验的 Authorization 请求 + :return: 可交给 Interaction 和 Authorization Code 服务的上下文 + :raises OAuthProtocolException: 请求不符合 OAuth/OIDC 协议时抛出 + """ + + if not OidcConfig.oidc_enabled: + raise OAuthProtocolException('temporarily_unavailable', 'OIDC provider is disabled', 503) + + client = await OAuthClientDao.get_by_client_id(db, request.client_id, active_only=True) + if client is None: + raise OAuthProtocolException('unauthorized_client', 'Client is not registered') + + registered_uri = await OAuthClientDao.find_exact_uri(db, client.client_pk, 'redirect', request.redirect_uri) + if registered_uri is None: + raise OAuthProtocolException('invalid_request', 'redirect_uri is not registered') + + redirect_uri = registered_uri.uri + state = request.state + if 'authorization_code' not in (client.grant_types or []) or 'code' not in (client.response_types or []): + raise cls._redirect_error( + 'unauthorized_client', 'Authorization code is not allowed for this client', redirect_uri, state + ) + if request.response_type != 'code': + raise cls._redirect_error( + 'unsupported_response_type', 'Only authorization code is supported', redirect_uri, state + ) + if request.code_challenge_method != 'S256': + raise cls._redirect_error('invalid_request', 'Only PKCE S256 is supported', redirect_uri, state) + if not _S256_CHALLENGE.fullmatch(request.code_challenge): + raise cls._redirect_error( + 'invalid_request', 'PKCE S256 code_challenge must contain 43 characters', redirect_uri, state + ) + if 'S256' not in OidcConfig.pkce_method_list or ( + OidcConfig.oidc_require_pkce and not bool(client.require_pkce) + ): + raise cls._redirect_error( + 'invalid_request', 'Client PKCE policy is weaker than provider policy', redirect_uri, state + ) + try: + OidcUtil.normalize_prompt(request.prompt) + except ValueError: + raise cls._redirect_error( + 'invalid_request', 'prompt contains an unsupported combination', redirect_uri, state + ) from None + if request.max_age is not None and request.max_age < 0: + raise cls._redirect_error('invalid_request', 'max_age must be non-negative', redirect_uri, state) + + requested_scopes = tuple(dict.fromkeys(request.scope.split())) + if 'openid' not in requested_scopes: + raise cls._redirect_error('invalid_scope', 'openid scope is required', redirect_uri, state) + if not request.nonce: + raise cls._redirect_error( + 'invalid_request', 'nonce is required for OpenID Connect authorization', redirect_uri, state + ) + + bindings = await OAuthClientDao.list_scope_bindings(db, client.client_pk) + scope_models = await cls._load_scope_models(db, bindings) + scope_by_code = {scope.scope_code: scope for scope in scope_models if scope.status == '0'} + unknown_scopes = [scope for scope in requested_scopes if scope not in scope_by_code] + if unknown_scopes: + raise cls._redirect_error( + 'invalid_scope', 'Requested scope is not allowed for this client', redirect_uri, state + ) + allowed_scope_codes = {scope.scope_pk for scope in scope_models} + if any(binding.scope_pk not in allowed_scope_codes for binding in bindings): + raise cls._redirect_error( + 'invalid_scope', 'Requested scope is not allowed for this client', redirect_uri, state + ) + bound_scope_codes = { + scope.scope_code + for scope in scope_models + if any(binding.scope_pk == scope.scope_pk for binding in bindings) + } + if any(scope not in bound_scope_codes for scope in requested_scopes): + raise cls._redirect_error( + 'invalid_scope', 'Requested scope is not allowed for this client', redirect_uri, state + ) + + resources = await OAuthClientDao.list_resources(db, client.client_pk) + resource_bindings = await OAuthClientDao.list_resource_bindings(db, client.client_pk) + resource = cls._resolve_resource( + request.resource, + resources, + resource_bindings, + [scope_by_code[code] for code in requested_scopes], + redirect_uri, + state, + ) + selected_resource_pk = resource.resource_pk if resource is not None else None + if any( + scope.scope_type == 'resource' + and (selected_resource_pk is None or scope.resource_pk != selected_resource_pk) + for scope in (scope_by_code[code] for code in requested_scopes) + ): + raise cls._redirect_error( + 'invalid_target', 'Requested scope does not belong to the selected resource', redirect_uri, state + ) + + prompt = request.prompt + requested_scope_models = [scope_by_code[code] for code in requested_scopes] + pre_authorized = frozenset( + scope.scope_code + for scope in requested_scope_models + for binding in bindings + if binding.scope_pk == scope.scope_pk and bool(binding.pre_authorized) + ) + required_scopes = frozenset( + {'openid'} | {scope.scope_code for scope in requested_scope_models if not bool(scope.consent_required)} + ) + + return AuthorizationContext( + client=ClientSnapshot( + client_pk=client.client_pk, + client_id=client.client_id, + policy_version=client.policy_version, + grant_types=tuple(client.grant_types or ()), + response_types=tuple(client.response_types or ()), + require_pkce=bool(client.require_pkce), + require_consent=bool(client.require_consent), + trusted_client=bool(client.trusted_client), + ), + redirect_uri=redirect_uri, + scopes=requested_scopes, + scope_models=tuple( + ScopeSnapshot( + scope_pk=scope_by_code[code].scope_pk, + scope_code=scope_by_code[code].scope_code, + scope_type=scope_by_code[code].scope_type, + resource_pk=scope_by_code[code].resource_pk, + consent_required=bool(scope_by_code[code].consent_required), + ) + for code in requested_scopes + ), + pre_authorized_scopes=pre_authorized, + resource=ResourceSnapshot(resource.resource_pk, resource.audience) if resource is not None else None, + state=state, + nonce=request.nonce, + code_challenge=request.code_challenge, + code_challenge_method=request.code_challenge_method, + prompt=prompt, + max_age=request.max_age, + required_scopes=required_scopes, + ) + + @staticmethod + async def _load_scope_models( + db: AsyncSession, bindings: list[SysOAuthClientScope] | tuple[SysOAuthClientScope, ...] + ) -> list[SysOAuthScope]: + """ + 加载 Client Scope 绑定对应的有效 Scope 定义 + + :param db: 异步数据库会话 + :param bindings: Client-Scope 绑定集合 + :return: Scope 定义列表 + """ + + scope_ids = {binding.scope_pk for binding in bindings} + if not scope_ids: + return [] + definitions = await OAuthClientDao.list_scope_definitions(db, active_only=False) + + return [scope for scope in definitions if scope.scope_pk in scope_ids] + + @staticmethod + async def load_scope_models( + db: AsyncSession, bindings: list[SysOAuthClientScope] | tuple[SysOAuthClientScope, ...] + ) -> list[SysOAuthScope]: + """ + 加载 Client Scope 绑定对应的 Scope 定义 + + :param db: 异步数据库会话 + :param bindings: Client-Scope 绑定集合 + :return: Scope 定义列表 + """ + + return await AuthorizationService._load_scope_models(db, bindings) + + @staticmethod + def _resolve_resource( + requested_resource: str | None, + resources: list[SysOAuthResource] | tuple[SysOAuthResource, ...], + resource_bindings: list[SysOAuthClientResource] | tuple[SysOAuthClientResource, ...], + requested_scopes: list[SysOAuthScope], + redirect_uri: str, + state: str | None, + ) -> SysOAuthResource | None: + """ + 解析并校验最多一个 Resource audience + + :param requested_resource: 请求中的 Resource audience + :param resources: Client 已注册的 Resource 集合 + :param resource_bindings: Client 与 Resource 的显式绑定集合 + :param requested_scopes: 已确认属于 Client 的 Scope 定义 + :param redirect_uri: 已精确验证的 Redirect URI + :param state: 原样绑定的客户端 state + :return: 选中的 Resource,或无 Resource 时返回 None + """ + + resource_scopes = [scope for scope in requested_scopes if scope.scope_type == 'resource'] + if requested_resource: + resource = next((item for item in resources if item.audience == requested_resource), None) + if resource is None: + raise AuthorizationService._redirect_error( + 'invalid_target', 'Requested resource is not allowed for this client', redirect_uri, state + ) + return resource + if not resource_scopes: + return None + default_pks = {binding.resource_pk for binding in resource_bindings if bool(binding.is_default)} + defaults = [resource for resource in resources if resource.resource_pk in default_pks] + if len(defaults) != 1: + raise AuthorizationService._redirect_error( + 'invalid_target', 'A resource audience is required for resource scope', redirect_uri, state + ) + return defaults[0] + + @staticmethod + def _redirect_error(error: str, description: str, redirect_uri: str, state: str | None) -> OAuthProtocolException: + """ + 创建已完成 Redirect 校验后的安全协议异常 + + :param error: OAuth 标准错误码 + :param description: 不泄漏内部数据的错误描述 + :param redirect_uri: 已注册的完整 Redirect URI + :param state: 客户端原样 state + :return: 标记为可安全重定向的协议异常 + """ + + return OAuthProtocolException( + error, + description, + 400, + redirect_uri=redirect_uri, + state=state, + redirect_uri_verified=True, + issuer=OidcConfig.oidc_issuer, + ) + + @staticmethod + def consent_is_satisfied(context: AuthorizationContext, grant: SysOAuthGrant | None) -> bool: + """ + 判断有效 Grant 是否覆盖当前授权请求 + + :param context: 已验证授权上下文 + :param grant: 当前用户和 Client 的 Grant + :return: Grant 有效且覆盖全部非预授权 Scope/Resource 时为 True + """ + + if not context.requires_consent: + return True + prompt_tokens = frozenset((context.prompt or '').split()) + if ( + 'consent' in prompt_tokens + or grant is None + or grant.status != 'active' + or grant.client_policy_version != context.client.policy_version + ): + return False + expires_at = TimezoneUtil.to_utc(grant.expires_at) if grant.expires_at is not None else None + if expires_at is not None and expires_at <= TimezoneUtil.utc_now(): + return False + non_pre_authorized = set(context.scopes) - set(context.pre_authorized_scopes) - set(context.required_scopes) + expected_resources = set(context.resources) + + return non_pre_authorized.issubset(set(grant.remembered_scopes or [])) and expected_resources.issubset( + set(grant.remembered_resources or []) + ) + + +class AuthorizationCodeReuseError(OAuthProtocolException): + """ + 表示一个已成功消费的 Authorization Code 被再次提交 + """ + + def __init__(self) -> None: + """ + 初始化对象状态 + + :return: None + """ + + super().__init__('invalid_grant', 'Authorization code is invalid or expired', 400) + self.must_commit = True + + +class AuthorizationCodeService: + """ + 授权码模块服务层 + + Redis 只保存 codeHash 和服务端白名单字段,不保存 Authorization Code 明文 + """ + + _CONSUME_SCRIPT = """ +local value = redis.call('GET', KEYS[1]) +if not value then + if redis.call('EXISTS', KEYS[2]) == 1 then + local consumed = redis.call('GET', KEYS[3]) + if consumed then + local ok_consumed, consumed_payload = pcall(cjson.decode, consumed) + if ok_consumed and type(consumed_payload) == 'table' and consumed_payload.codeHash + and consumed_payload.codeHash ~= ARGV[1] then + return -3 + end + end + return -4 + end + return nil +end +local ok, payload = pcall(cjson.decode, value) +if not ok or type(payload) ~= 'table' then return -2 end +if tostring(payload.version) ~= ARGV[2] or payload.codeHash ~= ARGV[1] then return -3 end +redis.call('DEL', KEYS[1]) +redis.call('SET', KEYS[2], ARGV[3], 'EX', ARGV[4]) +redis.call('SET', KEYS[3], value, 'EX', ARGV[4]) +return value +""" + _REQUIRED_FIELDS = frozenset( + { + 'clientPk', + 'redirectUri', + 'userId', + 'subjectId', + 'authVersion', + 'sid', + 'grantId', + 'scopes', + 'resources', + 'nonce', + 'codeChallenge', + 'codeChallengeMethod', + 'authTime', + } + ) + _MAX_SCOPES = 100 + _MAX_RESOURCES = 1 + _REUSE_TOMBSTONE_TTL_SECONDS = 300 + _CONSUMED_TOMBSTONE_RESULT = -4 + + @classmethod + async def issue( + cls, + redis: Redis, + payload: Mapping[str, Any], + *, + ttl_seconds: int | None = None, + pepper: str | None = None, + ) -> str: + """ + 签发短期 Authorization Code 并以 NX 写入 Redis + + :param redis: 异步 Redis 客户端 + :param payload: 已由服务端构建的白名单标量载荷 + :param ttl_seconds: 可选 Code TTL,省略时使用配置 + :param pepper: 独立 OIDC Token Hash Pepper + :return: 仅返回一次的 Authorization Code 明文 + :raises ValueError: 载荷或 Pepper 不符合安全约束 + :raises OAuthProtocolException: Redis 写入失败时抛出 + """ + + record = cls._validate_payload(payload) + code = generate_authorization_code() + parsed = parse_opaque_token(code, 'ac1') + digest = token_digest(code, pepper or OidcConfig.oidc_token_hash_pepper) + record['codeHash'] = digest + record['version'] = 1 + serialized = OidcUtil.serialize_json(record, error_message='授权码载荷必须支持 JSON 序列化') + ttl = OidcConfig.oidc_authorization_code_ttl_seconds if ttl_seconds is None else ttl_seconds + if not isinstance(ttl, int) or isinstance(ttl, bool) or ttl <= 0: + raise ValueError('授权码有效期必须为正整数') + created = await redis.set(OidcRedisKey.authorization_code(parsed.token_id), serialized, ex=ttl, nx=True) + if not created: + raise OAuthProtocolException('server_error', 'Authorization code could not be created', 500) + return code + + @classmethod + async def consume( + cls, + redis: Redis, + code: str, + *, + pepper: str | None = None, + ) -> dict[str, Any]: + """ + 使用 Redis Lua 原子校验并消费 Authorization Code + + :param redis: 异步 Redis 客户端 + :param code: 客户端提交的 Authorization Code 明文 + :param pepper: 独立 OIDC Token Hash Pepper + :return: 去除内部 codeHash/version 后的服务端载荷 + :raises OAuthProtocolException: Code 格式、Secret、状态或 TTL 无效时抛出统一错误 + """ + + try: + parsed = parse_opaque_token(code, 'ac1') + digest = token_digest(code, pepper or OidcConfig.oidc_token_hash_pepper) + except (OpaqueTokenError, TypeError, ValueError): + raise OAuthProtocolException('invalid_grant', 'Authorization code is invalid') from None + value = await redis.eval( + cls._CONSUME_SCRIPT, + 3, + OidcRedisKey.authorization_code(parsed.token_id), + OidcRedisKey.authorization_code_consumed(parsed.token_id), + OidcRedisKey.authorization_code_consumed_payload(parsed.token_id), + digest, + '1', + 'consumed', + max(OidcConfig.oidc_authorization_code_ttl_seconds, cls._REUSE_TOMBSTONE_TTL_SECONDS), + ) + if value == cls._CONSUMED_TOMBSTONE_RESULT: + raise AuthorizationCodeReuseError + if (isinstance(value, int) and value in {-2, -3}) or not value: + raise OAuthProtocolException('invalid_grant', 'Authorization code is invalid or expired') + try: + if isinstance(value, bytes): + value = value.decode('utf-8') + payload = json.loads(value) + except (TypeError, UnicodeDecodeError, json.JSONDecodeError): + raise OAuthProtocolException('invalid_grant', 'Authorization code state is invalid') from None + if not isinstance(payload, dict): + raise OAuthProtocolException('invalid_grant', 'Authorization code state is invalid') + payload.pop('codeHash', None) + payload.pop('version', None) + try: + return cls._validate_payload(payload) + except ValueError: + raise OAuthProtocolException('invalid_grant', 'Authorization code state is invalid') from None + + @classmethod + async def consumed_payload( + cls, + redis: Redis, + code: str, + *, + pepper: str | None = None, + ) -> dict[str, Any] | None: + """ + 读取已消费授权码的短期绑定载荷 + + :param redis: Authorization Code Redis 客户端 + :param code: 客户端提交的 Authorization Code 明文 + :param pepper: 独立 OIDC Token Hash Pepper + :return: 与授权码摘要匹配的绑定载荷,不存在或摘要不匹配时返回 None + """ + + try: + parsed = parse_opaque_token(code, 'ac1') + digest = token_digest(code, pepper or OidcConfig.oidc_token_hash_pepper) + except (OpaqueTokenError, TypeError, ValueError): + return None + value = await redis.get(OidcRedisKey.authorization_code_consumed_payload(parsed.token_id)) + if not value: + return None + try: + if isinstance(value, bytes): + value = value.decode('utf-8') + payload = json.loads(value) + except (TypeError, UnicodeDecodeError, json.JSONDecodeError): + return None + if not isinstance(payload, dict) or payload.get('codeHash') != digest: + return None + payload.pop('codeHash', None) + payload.pop('version', None) + try: + return cls._validate_payload(payload) + except ValueError: + return None + + @classmethod + async def invalidate(cls, redis: Redis, code: str) -> None: + """ + 精确删除一个已签发但未安全完成交付的授权码 + + :param redis: Authorization Code Redis 客户端 + :param code: 仅用于解析 code_id,不会写入 Redis + :return: None + """ + + try: + parsed = parse_opaque_token(code, 'ac1') + except (OpaqueTokenError, TypeError, ValueError): + return + await redis.delete(OidcRedisKey.authorization_code(parsed.token_id)) + await redis.delete(OidcRedisKey.authorization_code_consumed_payload(parsed.token_id)) + + @classmethod + def _validate_payload(cls, payload: Mapping[str, Any]) -> dict[str, Any]: + """ + 校验 Authorization Code 的服务端白名单标量载荷 + + :param payload: 待校验映射 + :return: 可稳定 JSON 序列化的复制载荷 + :raises ValueError: 存在未知、缺失或类型不安全字段时抛出 + """ + + if not isinstance(payload, Mapping): + raise ValueError('授权码载荷必须为映射对象') + keys = set(payload) + if keys != cls._REQUIRED_FIELDS: + raise ValueError('授权码载荷包含不允许的字段') + OidcUtil.positive_int(payload['clientPk'], 'clientPk') + OidcUtil.positive_int(payload['userId'], 'userId') + OidcUtil.nonnegative_int(payload['authVersion'], 'authVersion') + OidcUtil.nonempty_string(payload['redirectUri'], 'redirectUri', 1000) + OidcUtil.nonempty_string(payload['subjectId'], 'subjectId', 36) + OidcUtil.nonempty_string(payload['sid'], 'sid', 36) + grant_id = payload['grantId'] + if grant_id is not None: + OidcUtil.nonempty_string(grant_id, 'grantId', 36) + OidcUtil.nonempty_string(payload['nonce'], 'nonce', 1024) + challenge = OidcUtil.nonempty_string(payload['codeChallenge'], 'codeChallenge', 128) + if not OidcUtil.is_s256_challenge(challenge): + raise ValueError('codeChallenge 必须为不带填充的 Base64URL 编码 SHA-256 摘要') + if payload['codeChallengeMethod'] != 'S256': + raise ValueError('codeChallengeMethod 必须为 S256') + auth_time = OidcUtil.nonempty_string(payload['authTime'], 'authTime', 64) + try: + parsed_auth_time = datetime.fromisoformat(auth_time.replace('Z', '+00:00')) + except ValueError as exc: + raise ValueError('认证时间 authTime 必须为 ISO-8601 格式') from exc + if parsed_auth_time.tzinfo is None or parsed_auth_time.utcoffset() != timedelta(0): + raise ValueError('认证时间 authTime 必须为带时区的 UTC 时间戳') + scopes = OidcUtil.string_list(payload['scopes'], 'scopes', cls._MAX_SCOPES) + resources = OidcUtil.string_list(payload['resources'], 'resources', cls._MAX_RESOURCES) + if 'openid' not in scopes or len(set(scopes)) != len(scopes): + raise ValueError('权限范围必须包含 openid,且不得重复') + if len(set(resources)) != len(resources): + raise ValueError('资源列表不得包含重复项') + if len(resources) > cls._MAX_RESOURCES: + raise ValueError('一次请求只支持一个业务资源') + return { + 'clientPk': payload['clientPk'], + 'redirectUri': payload['redirectUri'], + 'userId': payload['userId'], + 'subjectId': payload['subjectId'], + 'authVersion': payload['authVersion'], + 'sid': payload['sid'], + 'grantId': grant_id, + 'scopes': scopes, + 'resources': resources, + 'nonce': payload['nonce'], + 'codeChallenge': challenge, + 'codeChallengeMethod': 'S256', + 'authTime': payload['authTime'], + } + + +class InteractionCompletionService: + """ + 交互完成模块服务层 + """ + + @staticmethod + async def complete(redis: Redis, interaction_id: str, db: AsyncSession) -> 'InteractionCompletionResult': + """ + 完成服务端校验并跳转至已注册的 Client Redirect URI + + :param redis: 交互 Redis + :param interaction_id: Interaction 标识 + :param db: 异步数据库会话 + :return: 外部 Client 重定向响应 + """ + + record = await InteractionService.get_record(redis, interaction_id) + if record.get('status') not in {'completed', 'denied'}: + raise OidcInteractionException(interaction_id, '认证交互尚未完成', error='interaction_required') + marker = await InteractionFlowService.reserve_completion(redis, interaction_id) + if marker is None: + raise OidcInteractionException( + interaction_id, '认证交互已完成,请勿重复提交', error='invalid_request', status_code=409 + ) + redirect_uri = record.get('redirectUri') + try: + registered = await OAuthClientDao.find_exact_uri(db, record['clientPk'], 'redirect', redirect_uri) + except Exception: + await InteractionFlowService.best_effort_delete(redis, marker) + await InteractionFlowService.rollback(db) + raise OAuthProtocolException('server_error', 'Validated redirect URI is unavailable', 500) from None + if registered is None: + await InteractionFlowService.best_effort_delete(redis, marker) + raise OAuthProtocolException('server_error', 'Validated redirect URI is unavailable', 500) + if record['status'] == 'denied': + try: + await db.commit() + except Exception: + await InteractionFlowService.best_effort_delete(redis, marker) + await InteractionFlowService.rollback(db) + raise OAuthProtocolException('server_error', 'Authorization completion failed', 500) from None + return InteractionCompletionService._redirect(registered.uri, record, error='access_denied') + return await InteractionCompletionService._issue_code(db, redis, record, registered.uri, marker) + + @staticmethod + async def active_session(db: AsyncSession, record: dict[str, Any]) -> SysSsoSession | None: + """ + 读取活动会话 + + :param db: 异步数据库会话 + :param record: 交互记录 + :return: 活动 SSO Session 或 None + """ + + sid = record.get('authenticatedSid') + if not isinstance(sid, str): + return None + session = await SsoSessionDao.get_active(db, sid, now=TimezoneUtil.utc_now()) + if session is None: + return None + if ( + session.user_id != record.get('userId') + or session.subject_id != record.get('subjectId') + or session.auth_version != record.get('authVersion') + ): + return None + return session + + @staticmethod + async def _issue_code( + db: AsyncSession, redis: Any, record: dict[str, Any], redirect_uri: str, marker: str + ) -> 'InteractionCompletionResult': + """ + 签发授权码 + + :param db: 异步数据库会话 + :param redis: Redis 客户端 + :param record: 交互记录 + :param redirect_uri: 已验证的重定向 URI + :param marker: 完成标记 Key + :return: 授权码 + """ + + try: + active = await InteractionCompletionService.active_session(db, record) + except Exception: + await InteractionFlowService.best_effort_delete(redis, marker) + await InteractionFlowService.rollback(db) + raise OAuthProtocolException('server_error', 'Current login could not be verified', 500) from None + if active is None: + await InteractionFlowService.best_effort_delete(redis, marker) + raise OAuthProtocolException('login_required', 'A current login is required', 400) + payload = { + 'clientPk': record['clientPk'], + 'redirectUri': redirect_uri, + 'userId': active.user_id, + 'subjectId': active.subject_id, + 'authVersion': active.auth_version, + 'sid': active.sid, + 'grantId': record.get('grantId'), + 'scopes': record['scopes'], + 'resources': record['resources'], + 'nonce': record['nonce'], + 'codeChallenge': record['codeChallenge'], + 'codeChallengeMethod': record['codeChallengeMethod'], + 'authTime': (TimezoneUtil.to_optional_utc(active.auth_time) or TimezoneUtil.utc_now()).isoformat(), + } + code: str | None = None + try: + grant = await AuthorizationService._completion_grant(db, record, active) + payload['grantId'] = grant.grant_id + code = await AuthorizationCodeService.issue(redis, payload, pepper=OidcConfig.oidc_token_hash_pepper) + await AuditService.record( + db, + OidcAuditEvent.AUTHORIZE_SUCCEEDED, + 'success', + client_id=record.get('clientId'), + user_id=active.user_id, + subject_id=str(active.subject_id), + sid=active.sid, + grant_id=grant.grant_id, + ) + await db.commit() + except Exception as exc: + if code is not None: + try: + await AuthorizationCodeService.invalidate(redis, code) + except Exception: + pass + await InteractionFlowService.best_effort_delete(redis, marker) + await InteractionFlowService.rollback(db) + if isinstance(exc, OAuthProtocolException): + raise + raise OAuthProtocolException('server_error', 'Authorization code could not be completed', 500) from None + return InteractionCompletionService._redirect(redirect_uri, record, code=code) + + @staticmethod + def _redirect( + redirect_uri: str, + record: dict[str, Any], + *, + code: str | None = None, + error: str | None = None, + ) -> 'InteractionCompletionResult': + """ + 构建重定向响应 + + :param redirect_uri: 已验证的重定向 URI + :param record: 交互记录 + :param code: 授权码或验证码 + :param error: 协议错误码 + :return: 重定向结果 + """ + + query = [] + if code is not None: + query.append(('code', code)) + if error is not None: + query.append(('error', error)) + if record.get('state') is not None: + query.append(('state', record['state'])) + query.append(('iss', OidcConfig.oidc_issuer)) + location = OidcUtil.replace_query_parameters( + redirect_uri, query, {'code', 'error', 'state', 'iss'}, fragment='' + ) + return InteractionCompletionResult(location=location) + + +@dataclass(frozen=True) +class InteractionCompletionResult: + """ + 完成交互后的重定向目标;Controller 负责构造 HTTP 跳转 + """ + + location: str diff --git a/ruoyi-fastapi-backend/module_identity/service/consent_service.py b/ruoyi-fastapi-backend/module_identity/service/consent_service.py new file mode 100644 index 000000000..f2cd565b7 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/consent_service.py @@ -0,0 +1,437 @@ +from collections.abc import Iterable +from dataclasses import dataclass + +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException, OidcInteractionException +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao, OAuthGrantSnapshot +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant +from module_identity.entity.vo.interaction_vo import InteractionConsentModel, InteractionResultModel +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import ( + AuthorizationContext, + AuthorizationService, + ClientSnapshot, + ResourceSnapshot, + ScopeSnapshot, +) +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.interaction_service import InteractionFlowService + + +@dataclass(frozen=True, slots=True) +class ConsentResult: + """ + 授权确认结果 + + :ivar approved: 是否批准授权 + :ivar scopes: 服务端从原请求范围内收敛后的 Scope + :ivar grant: 本次同意关联的可撤销 Grant + """ + + approved: bool + scopes: tuple[str, ...] + grant: SysOAuthGrant | None = None + previous_grant: OAuthGrantSnapshot | None = None + persisted_grant: OAuthGrantSnapshot | None = None + + +class ConsentService: + """ + 授权同意模块服务层 + + 用户提交只能取消原始请求中的可选 Scope,不能增加 Scope,也不能取消服务端必需 Scope + 每次同意均保存可撤销 Grant,记住同意仅控制后续是否需要再次确认 + """ + + @staticmethod + def validate_submission(context: AuthorizationContext, approved: bool, scopes: Iterable[str]) -> tuple[str, ...]: + """ + 校验并收敛用户提交的授权 Scope + + :param context: 已完成协议校验的授权上下文 + :param approved: 用户是否批准授权 + :param scopes: 用户提交的 Scope 集合 + :return: 按原请求顺序返回的有效 Scope + :raises OAuthProtocolException: 拒绝授权、扩大 Scope 或取消必需 Scope 时抛出 + """ + + requested = tuple(context.scopes) + submitted = tuple(dict.fromkeys(scopes)) + requested_set = set(requested) + submitted_set = set(submitted) + if not approved: + raise ConsentService._redirect_error(context, 'access_denied', 'User denied the authorization request') + if not submitted_set.issubset(requested_set): + raise ConsentService._redirect_error(context, 'invalid_scope', 'Submitted scope is not allowed') + if not set(context.required_scopes).issubset(submitted_set): + raise ConsentService._redirect_error(context, 'invalid_scope', 'Required scope cannot be removed') + return tuple(scope for scope in requested if scope in submitted_set) + + @classmethod + async def remember_consent( + cls, + db: AsyncSession, + context: AuthorizationContext, + user_id: int, + subject_id: str, + scopes: Iterable[str], + ) -> SysOAuthGrant: + """ + 持久化用户同意的 Scope + + :param db: 异步数据库会话 + :param context: 已完成协议校验的授权上下文 + :param user_id: 本地用户 ID + :param subject_id: 稳定 OIDC Subject + :param scopes: 已通过校验的 Scope 集合 + :return: 新建或合并后的 Grant + + 调用方负责提交事务;本方法不执行 commit。Grant 撤销与 Token 联动不在本接口内完成 + """ + + selected_scopes = cls.validate_submission(context, True, scopes) + + return await OAuthGrantDao.merge_active_grant( + db=db, + user_id=user_id, + subject_id=subject_id, + client_pk=context.client.client_pk, + granted_scopes=list(selected_scopes), + granted_resources=list(context.resources), + client_policy_version=context.client.policy_version, + ) + + @classmethod + async def submit_consent( + cls, + db: AsyncSession, + context: AuthorizationContext, + approved: bool, + scopes: Iterable[str], + remember_consent: bool = False, + user_id: int | None = None, + subject_id: str | None = None, + ) -> ConsentResult: + """ + 处理授权确认并写入可撤销 Grant + + :param db: 异步数据库会话 + :param context: 已完成协议校验的授权上下文 + :param approved: 用户是否批准授权 + :param scopes: 用户提交的 Scope 集合 + :param remember_consent: 是否在后续请求复用本次同意,不影响授权记录的保存 + :param user_id: 持久化 Grant 所需的本地用户 ID + :param subject_id: 持久化 Grant 所需的稳定 Subject + :return: 授权处理结果 + + 调用方负责 commit;拒绝授权不会写入 Grant + """ + + selected_scopes = cls.validate_submission(context, approved, scopes) + if user_id is None or subject_id is None: + raise ValueError('创建授权记录必须提供用户编号和主体标识') + await OAuthAccessPolicyDao.lock_client(db, context.client.client_pk) + if await OAuthAccessPolicyDao.is_blocked(db, user_id, context.client.client_pk, for_update=True): + raise OAuthProtocolException('access_denied', 'Access to this application is blocked') + current = await OAuthGrantDao.get_active_for_user_client(db, user_id, context.client.client_pk, for_update=True) + previous_grant = OAuthGrantDao.snapshot(current) if current is not None else None + grant = await OAuthGrantDao.merge_active_grant( + db=db, + user_id=user_id, + subject_id=subject_id, + client_pk=context.client.client_pk, + granted_scopes=list(selected_scopes), + granted_resources=list(context.resources), + client_policy_version=context.client.policy_version, + remember_consent=remember_consent, + ) + return ConsentResult( + approved=True, + scopes=selected_scopes, + grant=grant, + previous_grant=previous_grant, + persisted_grant=OAuthGrantDao.snapshot(grant) if grant is not None else None, + ) + + @staticmethod + async def compensate_persisted_grant(db: AsyncSession, result: ConsentResult) -> None: + """ + 补偿已提交但 Interaction CAS 失败的持久授权 + + :param db: 用于独立补偿事务的数据库会话 + :param result: 已提交的授权结果及更新前快照 + :return: None + """ + + if result.grant is None or result.persisted_grant is None: + return + try: + if result.previous_grant is None: + restored = await OAuthGrantDao.revoke_snapshot( + db, + result.persisted_grant, + reason='interaction_transition_failed', + ) + else: + restored = await OAuthGrantDao.restore_snapshot( + db, + result.previous_grant, + result.persisted_grant, + ) + if restored: + await db.commit() + return + await db.rollback() + await AuditService.record_independent( + db, + OidcAuditEvent.CONSENT_GRANTED, + 'failure', + risk_level='high', + user_id=result.persisted_grant.user_id, + subject_id=result.persisted_grant.subject_id, + grant_id=result.persisted_grant.grant_id, + failure_code='consent_compensation_conflict', + detail={'reason': 'grant_changed_after_interaction_commit'}, + ) + except Exception: + await db.rollback() + raise + + @staticmethod + def consent_is_satisfied(context: AuthorizationContext, grant: SysOAuthGrant | None) -> bool: + """ + 判断已有 Grant 是否可以跳过本次同意页 + + :param context: 已完成协议校验的授权上下文 + :param grant: 当前用户和 Client 的候选 Grant + :return: 仅当策略版本、范围和 Resource 均匹配时返回 True + """ + + return AuthorizationService.consent_is_satisfied(context, grant) + + @staticmethod + def _redirect_error(context: AuthorizationContext, error: str, description: str) -> OAuthProtocolException: + """ + 创建带已验证 Redirect 的授权错误 + + :param context: 已完成 Redirect 校验的上下文 + :param error: OAuth 标准错误码 + :param description: 安全错误描述 + :return: 可安全重定向的协议异常 + """ + + return OAuthProtocolException( + error, + description, + 400, + redirect_uri=context.redirect_uri, + state=context.state, + redirect_uri_verified=True, + issuer=OidcConfig.oidc_issuer, + ) + + +class InteractionConsentService: + """ + 认证中心授权同意流程模块服务层 + """ + + @staticmethod + async def context_from_record(db: AsyncSession, record: dict[str, object]) -> AuthorizationContext: + """ + 从 Interaction 白名单记录和当前数据库策略重建同意上下文 + + :param db: 异步数据库会话 + :param record: Interaction 内部记录 + :return: 当前策略下的授权上下文 + """ + + client = await OAuthClientDao.get_by_pk(db, record['clientPk'], active_only=True) + if client is None: + raise OidcInteractionException( + record['interactionId'], '客户端已停用或不可用', error='invalid_request', status_code=409 + ) + bindings = await OAuthClientDao.list_scope_bindings(db, client.client_pk) + scopes = await AuthorizationService.load_scope_models(db, bindings) + by_code = {scope.scope_code: scope for scope in scopes if scope.status == '0'} + requested = tuple(record['scopes']) + if any(scope not in by_code for scope in requested): + raise OidcInteractionException( + record['interactionId'], '权限范围策略已变更,请重新授权', error='invalid_scope' + ) + requested_resources = tuple(record['resources']) + resource_rows = await OAuthClientDao.list_resources(db, client.client_pk) + resource = next((item for item in resource_rows if item.audience in requested_resources), None) + if requested_resources and resource is None: + raise OidcInteractionException( + record['interactionId'], '资源访问策略已变更,请重新授权', error='invalid_scope', status_code=409 + ) + required = frozenset({'openid'} | {code for code in requested if not bool(by_code[code].consent_required)}) + pre_authorized = frozenset( + code + for code in requested + for binding in bindings + if binding.scope_pk == by_code[code].scope_pk and bool(binding.pre_authorized) + ) + + return AuthorizationContext( + client=ClientSnapshot( + client_pk=client.client_pk, + client_id=client.client_id, + policy_version=client.policy_version, + grant_types=tuple(client.grant_types or ()), + response_types=tuple(client.response_types or ()), + require_pkce=bool(client.require_pkce), + require_consent=bool(client.require_consent), + trusted_client=bool(client.trusted_client), + ), + redirect_uri=record['redirectUri'], + scopes=requested, + scope_models=tuple( + ScopeSnapshot( + scope_pk=by_code[code].scope_pk, + scope_code=code, + scope_type=by_code[code].scope_type, + resource_pk=by_code[code].resource_pk, + consent_required=bool(by_code[code].consent_required), + ) + for code in requested + ), + pre_authorized_scopes=pre_authorized, + resource=ResourceSnapshot(resource.resource_pk, resource.audience) if resource is not None else None, + state=record.get('state'), + nonce=record['nonce'], + code_challenge=record['codeChallenge'], + code_challenge_method=record['codeChallengeMethod'], + prompt=' '.join(record.get('prompt') or ()), + max_age=record.get('maxAge'), + required_scopes=required, + ) + + @staticmethod + async def consent( + redis: Redis, + interaction_id: str, + body: InteractionConsentModel, + db: AsyncSession, + csrf_token: str | None, + ) -> InteractionResultModel: + """ + 校验并提交授权同意,只允许原始 Scope 子集 + + :param redis: 交互 Redis + :param interaction_id: Interaction 标识 + :param body: 同意状态和 Scope 参数 + :param db: 异步数据库会话 + :param csrf_token: Interaction CSRF 原文 + :return: 同源完成跳转动作 + """ + + record = await InteractionFlowService.csrf_record(redis, interaction_id, csrf_token) + InteractionFlowService.require_status(record, 'awaiting_consent') + context = await InteractionConsentService.context_from_record(db, record) + try: + result = await ConsentService.submit_consent( + db, + context, + body.approved, + body.scopes, + body.remember_consent, + user_id=record.get('userId'), + subject_id=record.get('subjectId'), + ) + except OAuthProtocolException as exc: + if exc.error != 'access_denied': + await AuditService.record_interaction_failure( + db, + OidcAuditEvent.AUTHORIZE_DENIED, + client_id=record.get('clientId'), + user_id=record.get('userId'), + failure_code=exc.error, + ) + raise + await AuditService.record( + db, + OidcAuditEvent.AUTHORIZE_DENIED, + 'success', + client_id=record.get('clientId'), + user_id=record.get('userId'), + subject_id=str(record.get('subjectId')) if record.get('subjectId') else None, + failure_code='access_denied', + ) + await InteractionFlowService.commit_transition( + db, AfterCommitCoordinator(), redis, interaction_id, 'denied' + ) + return InteractionFlowService.interaction_result(interaction_id, 'redirect') + + grant_id = result.grant.grant_id if result.grant is not None else record.get('grantId') + await AuditService.record( + db, + OidcAuditEvent.CONSENT_GRANTED, + 'success', + client_id=record.get('clientId'), + user_id=record.get('userId'), + subject_id=str(record.get('subjectId')) if record.get('subjectId') else None, + grant_id=grant_id, + ) + + async def compensate() -> None: + """ + 执行事务补偿 + + :return: None + """ + + await ConsentService.compensate_persisted_grant(db, result) + + await InteractionFlowService.commit_transition( + db, + AfterCommitCoordinator(), + redis, + interaction_id, + 'completed', + {'scopes': list(result.scopes), 'grantId': grant_id}, + compensate=compensate if result.grant is not None else None, + ) + + return InteractionFlowService.interaction_result(interaction_id, 'redirect') + + @staticmethod + async def cancel( + redis: Redis, + interaction_id: str, + db: AsyncSession, + csrf_token: str | None, + ) -> InteractionResultModel: + """ + 以 CSRF 保护的原子状态迁移拒绝当前授权 + + :param redis: 交互 Redis + :param interaction_id: Interaction 标识 + :param db: 异步数据库会话 + :param csrf_token: Interaction CSRF 原文 + :return: 同源完成跳转动作 + """ + + record = await InteractionFlowService.csrf_record(redis, interaction_id, csrf_token) + if record.get('status') not in {'awaiting_login', 'awaiting_consent', 'password_change_required'}: + raise OidcInteractionException( + interaction_id, '认证交互已失效,请重新发起认证', error='invalid_request', status_code=409 + ) + await AuditService.record( + db, + OidcAuditEvent.AUTHORIZE_DENIED, + 'success', + client_id=record.get('clientId'), + user_id=record.get('userId'), + failure_code='access_denied', + ) + await InteractionFlowService.commit_transition(db, AfterCommitCoordinator(), redis, interaction_id, 'denied') + + return InteractionFlowService.interaction_result(interaction_id, 'redirect') diff --git a/ruoyi-fastapi-backend/module_identity/service/discovery_service.py b/ruoyi-fastapi-backend/module_identity/service/discovery_service.py new file mode 100644 index 000000000..42df9d998 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/discovery_service.py @@ -0,0 +1,143 @@ +from sqlalchemy.ext.asyncio import AsyncSession + +from config.env import OidcConfig +from module_identity.dao.oauth_client_dao import OAuthClientDao + + +class DiscoveryService: + """ + OIDC Discovery 模块服务层 + """ + + @staticmethod + async def supported_scopes(db: AsyncSession) -> list[str]: + """ + 读取启用 Scope 并稳定合并到协议元数据 + + :param db: 异步数据库会话 + :return: 启用 Scope 名称列表 + """ + + scopes = await OAuthClientDao.list_scope_definitions(db, active_only=True) + dynamic = tuple( + value if isinstance(value, str) else value.scope_code + for value in scopes + if isinstance(value, str) or isinstance(value.scope_code, str) + ) + + return list( + dict.fromkeys((*('openid', 'profile', 'email', 'phone', 'dept', 'roles', 'offline_access'), *dynamic)) + ) + + @staticmethod + def _issuer() -> str: + """ + 读取签发者地址 + + :return: OIDC 签发者地址 + """ + + return OidcConfig.oidc_issuer.rstrip('/') + + @classmethod + def openid_metadata(cls) -> dict[str, object]: + """ + 构建 OpenID 元数据 + + :return: 协议元数据映射 + """ + + issuer = cls._issuer() + + return { + 'issuer': issuer, + 'authorization_endpoint': f'{issuer}/oauth2/authorize', + 'token_endpoint': f'{issuer}/oauth2/token', + 'userinfo_endpoint': f'{issuer}/oauth2/userinfo', + 'jwks_uri': f'{issuer}/oauth2/jwks', + 'revocation_endpoint': f'{issuer}/oauth2/revoke', + 'introspection_endpoint': f'{issuer}/oauth2/introspect', + 'end_session_endpoint': f'{issuer}/oauth2/logout', + 'scopes_supported': ['openid', 'profile', 'email', 'phone', 'dept', 'roles', 'offline_access'], + 'response_types_supported': ['code'], + 'response_modes_supported': ['query'], + 'grant_types_supported': ['authorization_code', 'refresh_token', 'client_credentials'], + 'subject_types_supported': ['public'], + 'id_token_signing_alg_values_supported': ['RS256'], + 'token_endpoint_auth_methods_supported': ['none', 'client_secret_basic'], + 'code_challenge_methods_supported': ['S256'], + 'claims_supported': [ + 'sub', + 'name', + 'preferred_username', + 'picture', + 'email', + 'email_verified', + 'phone_number', + 'phone_number_verified', + 'dept_id', + 'dept_name', + 'roles', + 'auth_time', + 'acr', + 'amr', + 'sid', + ], + 'authorization_response_iss_parameter_supported': True, + 'backchannel_logout_supported': True, + 'backchannel_logout_session_supported': True, + } + + @classmethod + def oauth_metadata(cls) -> dict[str, object]: + """ + 构建 OAuth 元数据 + + :return: 协议元数据映射 + """ + + issuer = cls._issuer() + + return { + 'issuer': issuer, + 'authorization_endpoint': f'{issuer}/oauth2/authorize', + 'token_endpoint': f'{issuer}/oauth2/token', + 'jwks_uri': f'{issuer}/oauth2/jwks', + 'revocation_endpoint': f'{issuer}/oauth2/revoke', + 'introspection_endpoint': f'{issuer}/oauth2/introspect', + 'scopes_supported': ['openid', 'profile', 'email', 'phone', 'dept', 'roles', 'offline_access'], + 'response_types_supported': ['code'], + 'response_modes_supported': ['query'], + 'grant_types_supported': ['authorization_code', 'refresh_token', 'client_credentials'], + 'token_endpoint_auth_methods_supported': ['none', 'client_secret_basic'], + 'code_challenge_methods_supported': ['S256'], + 'authorization_response_iss_parameter_supported': True, + } + + @classmethod + async def openid_metadata_with_scopes(cls, db: AsyncSession) -> dict[str, object]: + """ + 构建带动态 Scope 的 OpenID 元数据 + + :param db: 异步数据库会话 + :return: 协议元数据映射 + """ + + payload = cls.openid_metadata() + payload['scopes_supported'] = await cls.supported_scopes(db) + + return payload + + @classmethod + async def oauth_metadata_with_scopes(cls, db: AsyncSession) -> dict[str, object]: + """ + 构建带动态 Scope 的 OAuth 元数据 + + :param db: 异步数据库会话 + :return: 协议元数据映射 + """ + + payload = cls.oauth_metadata() + payload['scopes_supported'] = await cls.supported_scopes(db) + + return payload diff --git a/ruoyi-fastapi-backend/module_identity/service/identity_service.py b/ruoyi-fastapi-backend/module_identity/service/identity_service.py new file mode 100644 index 000000000..8bb1df8c4 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/identity_service.py @@ -0,0 +1,982 @@ +from collections.abc import Awaitable, Callable, Iterable, Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Any, Literal + +from redis.asyncio import Redis +from sqlalchemy import Row +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import CommonConstant, OidcAuditEvent +from common.enums import RedisInitKeyConfig +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_admin.dao.login_dao import login_by_account +from module_admin.entity.do.dept_do import SysDept +from module_admin.entity.do.user_do import SysUser +from module_admin.entity.vo.login_vo import UserLogin +from module_identity.dao.identity_subject_dao import IdentitySubjectDao +from module_identity.dao.identity_user_dao import IdentityUserDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_token_dao import OAuthTokenDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.audit_service import AuditService +from utils.oidc_util import OidcUtil +from utils.pwd_util import PwdUtil +from utils.time_util import TimezoneUtil + +_DUMMY_PASSWORD_HASH = '$2b$12$ySHJfAWxzh49cIc7M5L21e5GlPyA7QhE2GkLn9XuUTqmKIRqhWIja' +_LOGIN_FAILURE_TTL = timedelta(minutes=10) +_CAPTCHA_CONSUME_SCRIPT = """ +local value = redis.call('get', KEYS[1]) +if value then + redis.call('del', KEYS[1]) +end +return value +""" +_OIDC_FAILURE_SCRIPT = """ +if redis.call('exists', KEYS[2]) == 1 then + return -1 +end +local count = redis.call('incr', KEYS[1]) +if count == 1 then + redis.call('expire', KEYS[1], ARGV[2]) +end +if count > tonumber(ARGV[1]) then + redis.call('del', KEYS[1]) + redis.call('set', KEYS[2], '1', 'EX', ARGV[2]) + return -1 +end +return count +""" + + +CredentialReason = Literal[ + 'invalid_credentials', + 'ip_blocked', + 'captcha_missing', + 'captcha_invalid', + 'account_locked', + 'user_disabled', +] + + +@dataclass(frozen=True) +class CredentialAuthenticationResult: + """ + 凭据校验成功后的稳定结果,不包含任何 Token + + :param user: 已通过密码检查的本地用户对象 + :param dept: 用户所属的部门对象或 None + :param acr: 密码认证上下文引用 + :param amr: 实际使用的认证方法集合 + :param remember_me: 后续 SSO 服务处理的记住登录选择 + :param password_change_required: 是否必须修改初始或过期密码 + :param password_change_reason: 修改密码原因 + """ + + user: SysUser + dept: SysDept | None + acr: str + amr: tuple[str, ...] + remember_me: bool + password_change_required: bool + password_change_reason: str | None + + +class CredentialAuthenticationError(Exception): + """ + 凭据校验失败的内部分类异常 + + :param reason: 脱敏后的失败原因分类 + :param legacy_message: Legacy 登录必须保持的中文错误消息 + """ + + def __init__(self, reason: CredentialReason, legacy_message: str) -> None: + """ + 初始化对象状态 + + :param reason: 凭据校验失败原因 + :param legacy_message: Legacy 登录错误消息 + :return: None + """ + + self.reason = reason + self.legacy_message = legacy_message + super().__init__(legacy_message) + + +class CredentialAuthenticationService: + """ + 凭据认证模块服务层 + """ + + ACR_PASSWORD = 'urn:ruoyi:acr:pwd' + + @classmethod + async def authenticate_legacy( + cls, + redis: Redis, + query_db: AsyncSession, + login_user: UserLogin, + *, + client_ip: str | None, + skip_captcha: bool = False, + ) -> Row[tuple[SysUser, SysDept]]: + """ + 执行 Legacy 登录的原有凭据检查,不签发 JWT + + :param redis: Legacy 登录使用的 Redis 客户端 + :param query_db: 异步数据库会话 + :param login_user: Legacy 登录参数 + :param client_ip: Controller 请求解析得到的客户端 IP + :param skip_captcha: 开发环境 API 文档登录是否跳过验证码 + :return: 通过检查的用户和部门行 + :raises CredentialAuthenticationError: 使用 Legacy 中文语义分类失败 + """ + + await cls._check_ip(client_ip, redis) + user_name = login_user.user_name + lock_key = f'{RedisInitKeyConfig.ACCOUNT_LOCK.key}:{user_name}' + account_lock = await redis.get(lock_key) + if user_name == account_lock: + raise CredentialAuthenticationError('account_locked', '账号已锁定,请稍后再试') + + if login_user.captcha_enabled and not skip_captcha: + await cls._check_legacy_captcha(redis, login_user) + + user = await login_by_account(query_db, user_name) + if not user: + raise CredentialAuthenticationError('invalid_credentials', '用户不存在') + if not cls._verify(login_user.password, user[0].password): + await cls._record_legacy_password_error(redis, user_name) + if user[0].status == '1': + raise CredentialAuthenticationError('user_disabled', '用户已停用') + await redis.delete(f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:{user_name}') + + return user + + @classmethod + async def authenticate_oidc( + cls, + redis: Redis, + query_db: AsyncSession, + *, + client_ip: str | None = None, + user_name: str, + password: str, + code: str | None = None, + uuid: str | None = None, + captcha_enabled: bool = False, + remember_me: bool = False, + ) -> CredentialAuthenticationResult: + """ + 执行 OIDC 交互的本地凭据校验,不触发 Legacy Token 流程 + + 用户不存在和密码错误共享 ``invalid_credentials`` 分类;用户名只用于 + HMAC 计算,OIDC 错误计数与锁定键不会包含用户名原文 + + :param redis: OIDC 交互使用的 Redis 客户端 + :param client_ip: Controller 请求解析得到的客户端 IP + :param query_db: 异步数据库会话 + :param user_name: 用户名,仅在内存中参与查询和摘要 + :param password: 用户密码,不写入日志或 Redis + :param code: 可选一次性验证码 + :param uuid: 验证码 Redis 标识 + :param captcha_enabled: 是否要求验证码 + :param remember_me: 交由后续 SSO 服务处理的记住登录选项 + :return: 凭据成功结果,不含 Token + :raises CredentialAuthenticationError: 脱敏分类的凭据失败 + """ + + await cls._check_ip(client_ip, redis) + username_digest = OidcUtil.hash_sensitive_identifier(user_name, OidcConfig.oidc_token_hash_pepper) + failure_key = OidcRedisKey.login_user_rate_limit(username_digest) + lock_key = f'{failure_key}:lock' + if await redis.get(lock_key): + raise CredentialAuthenticationError('account_locked', '账号已锁定,请稍后再试') + if captcha_enabled: + await cls._check_oidc_captcha(redis, code, uuid) + + user_row = await login_by_account(query_db, user_name) + stored_hash = user_row[0].password if user_row else _DUMMY_PASSWORD_HASH + password_valid = cls._verify(password, stored_hash) + if user_row is None or not password_valid: + await cls._record_oidc_password_error(redis, failure_key, lock_key) + raise CredentialAuthenticationError('invalid_credentials', '账号或密码错误') + if user_row[0].status == '1': + raise CredentialAuthenticationError('user_disabled', '用户已停用') + + await redis.delete(failure_key) + await redis.delete(lock_key) + password_change_required, change_reason = await cls._password_policy(redis, user_row[0]) + methods = ('pwd', 'captcha') if captcha_enabled else ('pwd',) + + return CredentialAuthenticationResult( + user=user_row[0], + dept=user_row[1], + acr=cls.ACR_PASSWORD, + amr=methods, + remember_me=remember_me, + password_change_required=password_change_required, + password_change_reason=change_reason, + ) + + @staticmethod + def _verify(password: str, stored_hash: str | None) -> bool: + """ + 使用项目 PwdUtil 验证密码,异常哈希按失败处理 + + :param password: 用户密码 + :param stored_hash: 已存储的密码摘要 + :return: 密码是否匹配 + """ + + try: + return bool(PwdUtil.verify_password(password, stored_hash or _DUMMY_PASSWORD_HASH)) + except (TypeError, ValueError): + return False + + @classmethod + async def _check_ip(cls, client_ip: str | None, redis: Redis) -> None: + """ + 拒绝系统配置黑名单中的客户端 IP + + :param client_ip: 客户端 IP 地址 + :param redis: Redis 客户端 + :return: None + """ + + value = await redis.get(f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.login.blackIPList') + if client_ip in (value.split(',') if value else []): + raise CredentialAuthenticationError('ip_blocked', '当前IP禁止登录') + + @classmethod + async def _check_legacy_captcha(cls, redis: Redis, login_user: UserLogin) -> None: + """ + 按原有非消费语义校验 Legacy 验证码 + + :param redis: Redis 客户端 + :param login_user: Legacy 登录请求 + :return: None + """ + + value = await redis.get(f'{RedisInitKeyConfig.CAPTCHA_CODES.key}:{login_user.uuid}') + if not value: + raise CredentialAuthenticationError('captcha_missing', '验证码已失效') + if login_user.code != str(value): + raise CredentialAuthenticationError('captcha_invalid', '验证码错误') + + @classmethod + async def _check_oidc_captcha(cls, redis: Redis, code: str | None, uuid: str | None) -> None: + """ + 原子消费 OIDC 验证码,阻止验证码重放 + + :param redis: Redis 客户端 + :param code: OIDC 验证码 + :param uuid: 验证码标识 + :return: None + """ + + if not uuid: + raise CredentialAuthenticationError('captcha_missing', '验证码已失效或不存在') + key = f'{RedisInitKeyConfig.CAPTCHA_CODES.key}:{uuid}' + value = await redis.eval(_CAPTCHA_CONSUME_SCRIPT, 1, key) + if not value: + raise CredentialAuthenticationError('captcha_missing', '验证码已失效或不存在') + expected = value.decode('utf-8') if isinstance(value, bytes) else str(value) + if code != expected: + raise CredentialAuthenticationError('captcha_invalid', '验证码错误') + + @classmethod + async def _record_legacy_password_error(cls, redis: Redis, user_name: str) -> None: + """ + 保持 Legacy 错误计数、阈值和十分钟 TTL 完全不变 + + :param redis: Redis 客户端 + :param user_name: 用户名 + :return: None + """ + + key = f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:{user_name}' + cached = await redis.get(key) + count = int(cached or 0) + 1 + await redis.set(key, count, ex=_LOGIN_FAILURE_TTL) + if count > CommonConstant.PASSWORD_ERROR_COUNT: + await redis.delete(key) + await redis.set(f'{RedisInitKeyConfig.ACCOUNT_LOCK.key}:{user_name}', user_name, ex=_LOGIN_FAILURE_TTL) + raise CredentialAuthenticationError( + 'account_locked', '10分钟内密码已输错超过5次,账号已锁定,请10分钟后再试' + ) + raise CredentialAuthenticationError('invalid_credentials', '密码错误') + + @classmethod + async def _record_oidc_password_error(cls, redis: Redis, failure_key: str, lock_key: str) -> None: + """ + 记录 OIDC 摘要命名空间错误状态,不写入用户名 + + :param redis: Redis 客户端 + :param failure_key: OIDC 失败计数 Key + :param lock_key: OIDC 锁定 Key + :return: None + """ + + result = await redis.eval( + _OIDC_FAILURE_SCRIPT, + 2, + failure_key, + lock_key, + str(CommonConstant.PASSWORD_ERROR_COUNT), + str(int(_LOGIN_FAILURE_TTL.total_seconds())), + ) + if int(result or 0) < 0: + raise CredentialAuthenticationError('account_locked', '账号已锁定,请稍后再试') + + @classmethod + async def _password_policy(cls, redis: Redis, user: SysUser) -> tuple[bool, str | None]: + """ + 复用系统初始密码和密码有效期配置,处理 naive/aware 时间 + + :param redis: Redis 客户端 + :param user: 系统用户对象 + :return: 密码修改要求及原因 + """ + + init_modify = await redis.get(f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.account.initPasswordModify') + update_time = getattr(user, 'pwd_update_date', None) + if init_modify == '1' and update_time is None: + return True, 'initial_password' + days_value = await redis.get(f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.account.passwordValidateDays') + try: + days = int(days_value or 0) + except (TypeError, ValueError): + days = 0 + if days > 0: + if update_time is None: + return True, 'password_expired' + now = TimezoneUtil.utc_now() + update_time = TimezoneUtil.to_utc(update_time) + if now > update_time + timedelta(days=days): + return True, 'password_expired' + return False, None + + +class ClaimService: + """ + OIDC Claim 模块服务层 + """ + + SCOPE_CLAIMS = { + 'openid': frozenset({'sub'}), + 'profile': frozenset({'name', 'preferred_username', 'picture', 'updated_at'}), + 'email': frozenset({'email', 'email_verified'}), + 'phone': frozenset({'phone_number', 'phone_number_verified'}), + 'roles': frozenset({'roles'}), + 'role': frozenset({'roles'}), + 'dept': frozenset({'dept_id', 'dept_name'}), + 'department': frozenset({'dept_id', 'dept_name'}), + } + FORBIDDEN_CLAIMS = frozenset( + { + 'user_id', + 'userid', + 'password', + 'passwd', + 'secret', + 'token', + 'access_token', + 'refresh_token', + 'client_secret', + } + ) + + @classmethod + def _claim_set(cls, value: Any) -> set[str]: + """ + 规范化策略或白名单中的 Claim 名称 + + :param value: 待规范化的 Claim 名称集合 + :return: 规范化 Claim 集合 + """ + + if value is True: + return set(cls.SCOPE_CLAIMS['openid'] | cls.SCOPE_CLAIMS['profile']) + if isinstance(value, str): + return OidcUtil.scope_set(value) + if isinstance(value, Mapping): + value = value.get('claims', value.get('allowed_claims', [])) + if isinstance(value, Iterable) and not isinstance(value, (bytes, str, Mapping)): + return {item for item in value if isinstance(item, str)} + return set() + + @classmethod + def _client_claims_for_scope(cls, scope: str, policy: Any) -> set[str]: + """ + 取得 Client 对某一 Scope 的 Claim 策略 + + :param scope: 待读取策略的 Scope 名称 + :param policy: Scope Claim 策略 + :return: Scope 对应 Claim 集合 + """ + + if isinstance(policy, Mapping): + if scope in policy: + value = policy[scope] + return set(cls.SCOPE_CLAIMS.get(scope, ())) if value is True else cls._claim_set(value) + scopes = policy.get('scopes') + if scopes is not None and scope not in OidcUtil.scope_set(scopes): + return set() + if scopes is not None: + return set(cls.SCOPE_CLAIMS.get(scope, ())) + return set() + allowed_scopes = OidcUtil.scope_set(policy) + + return set(cls.SCOPE_CLAIMS.get(scope, ())) if scope in allowed_scopes else set() + + @classmethod + def effective_claims( + cls, + requested_scopes: str | Iterable[str] | None, + client_scope_policy: Any, + resource_allowed_claims: Any, + ) -> set[str]: + """ + 计算请求 Scope、Client 策略和 Resource 白名单的三重交集 + + :param requested_scopes: 请求的空格分隔 Scope 或 Scope 集合 + :param client_scope_policy: Client 的 Scope 到 Claim 策略映射 + :param resource_allowed_claims: Resource 允许的 Claim 列表 + :return: 可安全生成的 Claim 名称集合 + """ + + requested = OidcUtil.scope_set(requested_scopes) + client_claims = set().union(*(cls._client_claims_for_scope(scope, client_scope_policy) for scope in requested)) + resource_claims = cls._claim_set(resource_allowed_claims) + + return (client_claims & resource_claims) - cls.FORBIDDEN_CLAIMS + + @classmethod + def _safe_values(cls, user: Any, subject_id: str | None, roles: Any, department: Any) -> dict[str, Any]: + """ + 抽取允许的用户属性,明确排除本地 ID 和密码 + + :param user: SysUser ORM 记录或用户 Claim 字段映射 + :param subject_id: 稳定身份主体 ID + :param roles: 角色集合 + :param department: 部门对象或部门映射 + :return: 安全 Claim 字段映射 + """ + + values = { + 'sub': subject_id or OidcUtil.read_field(user, 'subject_id'), + 'name': OidcUtil.read_field(user, 'name', OidcUtil.read_field(user, 'nick_name')), + 'preferred_username': OidcUtil.read_field( + user, 'preferred_username', OidcUtil.read_field(user, 'user_name') + ), + 'picture': OidcUtil.read_field(user, 'picture', OidcUtil.read_field(user, 'avatar')), + 'updated_at': OidcUtil.claim_numeric_date( + OidcUtil.read_field(user, 'updated_at', OidcUtil.read_field(user, 'update_time')) + ), + 'email': OidcUtil.read_field(user, 'email'), + 'email_verified': OidcUtil.read_field(user, 'email_verified'), + 'phone_number': OidcUtil.read_field(user, 'phone_number', OidcUtil.read_field(user, 'phonenumber')), + 'phone_number_verified': OidcUtil.read_field(user, 'phone_number_verified'), + 'roles': roles if roles is not None else OidcUtil.read_field(user, 'roles'), + } + dept = department if department is not None else OidcUtil.read_field(user, 'dept') + if dept is None: + dept = { + 'dept_id': OidcUtil.read_field(user, 'dept_id'), + 'dept_name': OidcUtil.read_field(user, 'dept_name'), + } + if isinstance(dept, Mapping): + values['dept_id'] = dept.get('dept_id') + values['dept_name'] = dept.get('dept_name') + else: + values['dept_id'] = getattr(dept, 'dept_id', None) + values['dept_name'] = getattr(dept, 'dept_name', None) + return values + + @classmethod + def build_claims( + cls, + user: Any, + requested_scopes: str | Iterable[str] | None, + client_scope_policy: Any, + resource_allowed_claims: Any, + *, + subject_id: str | None = None, + roles: Any = None, + department: Any = None, + ) -> dict[str, Any]: + """ + 为用户构建经过三重授权和敏感字段过滤的 OIDC Claims + + :param user: 已加载的用户 ORM 对象或安全字段映射 + :param requested_scopes: 请求 Scope + :param client_scope_policy: Client Scope 到 Claim 的授权策略 + :param resource_allowed_claims: Resource Claim 白名单 + :param subject_id: 稳定 OIDC Subject,不能使用本地 user_id 替代 + :param roles: 已按项目角色 DAO 查询的角色 key 集合 + :param department: 已按项目部门 DAO 查询的部门对象或映射 + :return: 最小 OIDC Claim 映射 + """ + + requested = OidcUtil.scope_set(requested_scopes) + allowed = cls.effective_claims(requested, client_scope_policy, resource_allowed_claims) + stable_subject = subject_id or OidcUtil.read_field(user, 'subject_id') + if 'openid' in requested: + if not isinstance(stable_subject, str) or not stable_subject.strip(): + raise OAuthProtocolException('server_error', 'User identity mapping is unavailable', 500) + if 'sub' not in allowed: + raise OAuthProtocolException('invalid_scope', 'OpenID scope requires the sub claim', 400) + values = cls._safe_values(user, subject_id, roles, department) + role_keys: set[str] = set() + if isinstance(client_scope_policy, Mapping): + for scope in requested: + rule = client_scope_policy.get(scope) + if isinstance(rule, Mapping) and isinstance(rule.get('allowed_role_keys'), list): + role_keys.update( + key for key in rule['allowed_role_keys'] if isinstance(key, str) and '*' not in key + ) + if 'roles' in values: + values['roles'] = [role for role in values['roles'] if role in role_keys] + return { + key: value + for key, value in values.items() + if key in allowed and key not in cls.FORBIDDEN_CLAIMS and value is not None + } + + @classmethod + async def load_roles_and_department(cls, db: AsyncSession, user_id: int) -> tuple[list[str], SysDept | None]: + """ + 按现有 ORM DAO 使用的启用状态规则加载角色和部门 + + :param db: 异步数据库会话 + :param user_id: 本地用户 ID,仅用于内部查询,不会进入 Claim + :return: 角色名称列表和部门对象 + """ + + return await IdentityUserDao.get_claim_attributes(db, user_id) + + @classmethod + async def resolve_scope_policy( + cls, + db: AsyncSession, + client_pk: int, + scopes: Iterable[str], + *, + resource_allowed_claims: Iterable[str] | None = None, + ) -> tuple[dict[str, Any], set[str]]: + """ + 解析 Client Scope 策略及最终 Claim 白名单 + + :param db: 异步数据库会话 + :param client_pk: Client 内部主键 + :param scopes: 当前已授权 Scope + :param resource_allowed_claims: 可选 Resource Claim 白名单 + :return: Scope 策略映射和可用 Claim 集合 + """ + + requested = set(scopes) + bindings = await OAuthClientDao.list_scope_bindings(db, client_pk) + definitions = await OAuthClientDao.list_scope_definitions(db) + by_pk = {item.scope_pk: item for item in definitions} + policy: dict[str, Any] = {} + allowed: set[str] = set() + for binding in bindings: + definition = by_pk.get(binding.scope_pk) + if definition is None or definition.status != '0' or definition.scope_code not in requested: + continue + policy[definition.scope_code] = binding.claim_filter if binding.claim_filter is not None else True + allowed.update(cls._claim_set(definition.claims)) + if resource_allowed_claims is not None: + allowed.intersection_update(cls._claim_set(resource_allowed_claims)) + return policy, allowed + + +class IdentitySubjectService: + """ + 身份主体模块服务层 + """ + + @classmethod + async def require_by_user_id( + cls, + db: AsyncSession, + user_id: int, + *, + audit_db: AsyncSession | None = None, + audit_writer: Callable[[int], Awaitable[object]] | None = None, + ) -> SysIdentitySubject: + """ + 获取用户的稳定 Subject,缺失时拒绝继续签发 + + :param db: 与当前业务事务相同的异步数据库会话 + :param user_id: 本地用户 ID + :param audit_db: 可选独立审计会话,避免主事务回滚时丢失完整性告警 + :param audit_writer: 可选外部审计写入器,提交边界由写入器负责 + 生产签发链路必须提供其中之一;两者均省略时审计仅随当前事务 flush, + 主事务回滚会一并回滚,不能作为持久化告警 + :return: 已存在的主体关联 + :raises OAuthProtocolException: 主体关联缺失或用户标识无效时抛出 + """ + + if not isinstance(user_id, int) or isinstance(user_id, bool) or user_id <= 0: + await cls._write_missing_audit(db, user_id, audit_db=audit_db, audit_writer=audit_writer) + raise OAuthProtocolException('server_error', 'User identity mapping is unavailable', 500) + subject = await IdentitySubjectDao.get_by_user_id(db, user_id) + if subject is None: + await cls._write_missing_audit(db, user_id, audit_db=audit_db, audit_writer=audit_writer) + raise OAuthProtocolException('server_error', 'User identity mapping is unavailable', 500) + return subject + + @classmethod + async def _write_missing_audit( + cls, + db: AsyncSession, + user_id: int, + *, + audit_db: AsyncSession | None = None, + audit_writer: Callable[[int], Awaitable[object]] | None = None, + ) -> None: + """ + 写入主体完整性审计;审计失败不能把缺失主体变成可签发状态 + + 未传入独立会话或写入器时只在调用方事务中 flush,调用方必须自行承担 + 回滚丢失告警的后果;持久化安全告警应使用 ``audit_writer`` + + :param db: 异步数据库会话 + :param user_id: 本地用户 ID + :param audit_db: 独立审计数据库会话 + :param audit_writer: 外部审计写入器 + :return: None + """ + + try: + if audit_writer is not None: + await audit_writer(user_id) + return + await AuditService.record( + audit_db or db, + event_type=OidcAuditEvent.IDENTITY_SUBJECT_MISSING, + result='failure', + risk_level='high', + user_id=user_id if isinstance(user_id, int) and not isinstance(user_id, bool) else None, + failure_code='identity_integrity', + detail={'reason': 'subject_missing'}, + ) + except Exception: + # 主体完整性拒绝优先于审计写入成功,调用方仍会得到 fail-closed 结果 + return + + @classmethod + async def create_for_new_user( + cls, + db: AsyncSession, + *, + user_id: int, + create_by: str | None = None, + subject_id: str | None = None, + ) -> SysIdentitySubject: + """ + 在用户创建事务中建立稳定 Subject,不提交调用方事务 + + :param db: 用户创建事务使用的异步数据库会话 + :param user_id: 已创建的本地用户 ID + :param create_by: 主体记录创建者 + :param subject_id: 测试迁移或导入场景使用的预生成 Subject + :return: 已存在或新建的主体关联 + :raises ValueError: 用户 ID 无效时抛出 + """ + + if not isinstance(user_id, int) or isinstance(user_id, bool) or user_id <= 0: + raise ValueError('身份用户编号 identity_user_id 必须为正整数') + return await IdentitySubjectDao.create_for_user(db, user_id, create_by, subject_id) + + @classmethod + async def repair_missing_subject( + cls, db: AsyncSession, user_id: int, create_by: str = 'identity-repair' + ) -> SysIdentitySubject: + """ + 幂等修复单个用户的缺失 Subject + + :param db: 异步数据库会话 + :param user_id: 待修复的本地用户 ID + :param create_by: 回填记录的创建者标识 + :return: 已存在或本次创建的主体关联 + :raises ValueError: 用户 ID 无效时抛出 + """ + + if not isinstance(user_id, int) or isinstance(user_id, bool) or user_id <= 0: + raise ValueError('身份用户编号 identity_user_id 必须为正整数') + try: + return await IdentitySubjectDao.create_for_user(db, user_id, create_by=create_by) + except IntegrityError: + existing = await IdentitySubjectDao.get_by_user_id(db, user_id) + if existing is not None: + return existing + raise + + @classmethod + async def repair_missing_subjects( + cls, + db: AsyncSession, + user_ids: Iterable[int] | None = None, + create_by: str = 'identity-repair', + ) -> Sequence[SysIdentitySubject]: + """ + 幂等批量修复缺失 Subject,不提交调用方事务 + + :param db: 异步数据库会话 + :param user_ids: 可选用户 ID 集合;省略时检查所有未删除用户 + :param create_by: 回填记录的创建者标识 + :return: 本次调用新建的主体记录集合 + """ + + if user_ids is not None: + values = list(user_ids) + if any(not isinstance(item, int) or isinstance(item, bool) or item <= 0 for item in values): + raise ValueError('身份用户编号 identity_user_id 必须为正整数') + return await IdentitySubjectDao.backfill_for_users( + db, await IdentitySubjectDao.list_missing_user_ids(db, user_ids), create_by + ) + + @classmethod + async def increment_auth_version( + cls, db: AsyncSession, user_id: int, expected_version: int | None = None + ) -> SysIdentitySubject: + """ + 原子递增用户认证版本并返回最新主体 + + :param db: 异步数据库会话 + :param user_id: 本地用户 ID + :param expected_version: 可选的乐观锁版本 + :return: 更新后的主体关联 + :raises OAuthProtocolException: 主体不存在或版本竞争失败时抛出 + """ + + if not await IdentitySubjectDao.increment_auth_version(db, user_id, expected_version): + raise OAuthProtocolException('server_error', 'User identity version could not be updated', 500) + subject = await IdentitySubjectDao.get_by_user_id(db, user_id) + if subject is None: + raise OAuthProtocolException('server_error', 'User identity mapping is unavailable', 500) + return subject + + @classmethod + async def get_by_subject_id(cls, db: AsyncSession, subject_id: str) -> SysIdentitySubject | None: + """ + 按稳定 Subject 查询主体关联 + + :param db: 异步数据库会话 + :param subject_id: 稳定 OIDC Subject + :return: 主体关联,不存在时返回 None + """ + + return await IdentitySubjectDao.get_by_subject_id(db, subject_id) + + +IdentitySecurityEvent = Literal[ + 'user_disabled', + 'user_deleted', + 'password_changed', + 'role_assignment_changed', + 'department_changed', + 'role_claim_changed', + 'role_disabled', + 'role_deleted', +] + + +class IdentitySecurityEventError(ValueError): + """ + 身份安全事件缺少主体映射或输入不满足事务约束 + """ + + +@dataclass(frozen=True, slots=True) +class IdentitySecurityEventResult: + """ + 一次身份安全失效操作的不可变结果 + """ + + affected_users: tuple[int, ...] + revoked_sessions: int + revoked_refresh_tokens: int + + +class IdentitySecurityEventService: + """ + 身份安全事件模块服务层 + """ + + _EVENTS = frozenset( + { + 'user_disabled', + 'user_deleted', + 'password_changed', + 'role_assignment_changed', + 'department_changed', + 'role_claim_changed', + 'role_disabled', + 'role_deleted', + } + ) + _ALWAYS_REVOKE_SSO = frozenset({'user_disabled', 'user_deleted', 'password_changed'}) + _CLAIM_EVENTS = frozenset( + {'role_assignment_changed', 'department_changed', 'role_claim_changed', 'role_disabled', 'role_deleted'} + ) + _TERMINAL_REFRESH_STATUSES = frozenset({'revoked', 'expired', 'reuse_detected'}) + _MAX_BATCH_USERS = 10_000 + _MAX_SID_LENGTH = 36 + + @classmethod + async def handle_user_event( + cls, + db: AsyncSession, + user_id: int, + event: IdentitySecurityEvent, + *, + actor: str | None = None, + exclude_sid: str | None = None, + revoke_sso_on_claim_change: bool = False, + now: datetime | None = None, + ) -> IdentitySecurityEventResult: + """ + 处理单个用户事件,但不提交调用方事务 + + :param db: 与用户或角色变更共用的数据库会话 + :param user_id: 本地用户 ID + :param event: 设计文档定义的身份安全事件 + :param actor: 可选审计操作者 + :param exclude_sid: 修改密码时允许保留的当前交互 Session + :param revoke_sso_on_claim_change: Claim 变化时是否立即撤销 SSO + :param now: 可注入的 项目当前时间 + :return: 版本递增及凭据撤销数量 + """ + + return await cls.handle_users_event( + db, + (user_id,), + event, + actor=actor, + exclude_sid=exclude_sid, + revoke_sso_on_claim_change=revoke_sso_on_claim_change, + now=now, + ) + + @classmethod + async def handle_users_event( + cls, + db: AsyncSession, + user_ids: Iterable[int], + event: IdentitySecurityEvent, + *, + actor: str | None = None, + exclude_sid: str | None = None, + revoke_sso_on_claim_change: bool = False, + now: datetime | None = None, + ) -> IdentitySecurityEventResult: + """ + 按固定锁顺序批量处理身份安全事件,不提交事务 + + :param db: 异步数据库会话 + :param user_ids: 本地用户 ID 集合 + :param event: 身份安全事件 + :param actor: 审计操作者标识 + :param exclude_sid: 需要排除的 Session 标识 + :param revoke_sso_on_claim_change: Claim 变化时是否撤销 SSO + :param now: 当前时间 + :return: 安全事件处理结果 + :raises IdentitySecurityEventError: 事件、用户或主体映射不完整 + """ + + try: + ids = OidcUtil.normalize_user_ids(user_ids, max_size=cls._MAX_BATCH_USERS) + except ValueError as exc: + raise IdentitySecurityEventError(str(exc)) from exc + if event not in cls._EVENTS: + raise IdentitySecurityEventError('身份安全事件无效') + if exclude_sid is not None and ( + event != 'password_changed' + or not OidcUtil.is_trimmed_identifier(exclude_sid, max_length=cls._MAX_SID_LENGTH) + ): + raise IdentitySecurityEventError('当前安全事件不允许指定保留会话 exclude_sid') + if not isinstance(revoke_sso_on_claim_change, bool): + raise IdentitySecurityEventError('身份声明变更时撤销会话的配置必须为布尔值') + if now is not None and not isinstance(now, datetime): + raise IdentitySecurityEventError('当前时间必须为 datetime 对象') + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + actor_value = actor.strip()[:64] if isinstance(actor, str) and actor.strip() else 'identity-security-event' + + subjects = list(await IdentitySubjectDao.list_for_users_for_update(db, ids)) + if {row.user_id for row in subjects} != set(ids): + raise IdentitySecurityEventError('缺少用户主体映射') + await IdentitySubjectDao.increment_auth_versions(db, ids, actor_value, current) + revoked_refresh_tokens = await OAuthTokenDao.revoke_for_users(db, ids, reason=event, now=current) + + revoke_sso = event in cls._ALWAYS_REVOKE_SSO or (event in cls._CLAIM_EVENTS and revoke_sso_on_claim_change) + revoked_sessions = 0 + if revoke_sso: + revoked_sessions = await SsoSessionDao.revoke_for_users( + db, ids, reason=event, now=current, exclude_sid=exclude_sid + ) + + await AuditService.record( + db, + event_type=OidcAuditEvent.SECURITY_VERSION_CHANGED, + result='success', + risk_level='high', + user_id=ids[0] if len(ids) == 1 else None, + detail={ + 'event': event, + 'actor': actor_value, + 'affected_users': len(ids), + 'revoked_sessions': revoked_sessions, + 'revoked_refresh_tokens': revoked_refresh_tokens, + }, + ) + + return IdentitySecurityEventResult(ids, revoked_sessions, revoked_refresh_tokens) + + @classmethod + async def handle_role_event( + cls, + db: AsyncSession, + role_id: int, + event: Literal['role_claim_changed', 'role_disabled', 'role_deleted'], + *, + actor: str | None = None, + revoke_sso_on_claim_change: bool = False, + now: datetime | None = None, + ) -> IdentitySecurityEventResult: + """ + 查询角色当前成员并批量触发身份安全失效 + + :param db: 异步数据库会话 + :param role_id: 角色 ID + :param event: 身份安全事件 + :param actor: 审计操作者标识 + :param revoke_sso_on_claim_change: Claim 变化时是否撤销 SSO + :param now: 当前时间 + :return: 安全事件处理结果 + """ + + if not isinstance(role_id, int) or isinstance(role_id, bool) or role_id <= 0: + raise IdentitySecurityEventError('角色编号 role_id 必须为正整数') + if event not in {'role_claim_changed', 'role_disabled', 'role_deleted'}: + raise IdentitySecurityEventError('角色安全事件无效') + user_ids = tuple(sorted(set(await IdentityUserDao.list_user_ids_by_role_id(db, role_id)))) + if not user_ids: + return IdentitySecurityEventResult((), 0, 0) + return await cls.handle_users_event( + db, + user_ids, + event, + actor=actor, + revoke_sso_on_claim_change=revoke_sso_on_claim_change, + now=now, + ) diff --git a/ruoyi-fastapi-backend/module_identity/service/infrastructure_service.py b/ruoyi-fastapi-backend/module_identity/service/infrastructure_service.py new file mode 100644 index 000000000..02b34d95d --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/infrastructure_service.py @@ -0,0 +1,150 @@ +from collections.abc import Awaitable, Callable + +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + + +class RateLimitExceeded(ValueError): + """ + 请求超过 Redis 原子窗口限额 + """ + + def __init__(self, retry_after: int) -> None: + """ + 创建限流异常 + + :param retry_after: 建议重试等待秒数 + :return: None + """ + + super().__init__('请求过于频繁,请稍后重试') + self.retry_after = max(1, retry_after) + + +class RateLimitUnavailable(RuntimeError): + """ + Redis 限流依赖不可用,调用方必须 fail closed + """ + + +class OidcRateLimiter: + """ + 认证限流模块服务层 + """ + + _SCRIPT = """ +local value = redis.call('INCR', KEYS[1]) +if value == 1 then redis.call('EXPIRE', KEYS[1], ARGV[1]) end +local ttl = redis.call('TTL', KEYS[1]) +return {value, ttl} +""" + + @classmethod + async def enforce(cls, redis: Redis, key: str, *, limit: int, window_seconds: int) -> None: + """ + 原子增加计数,超限或 Redis 故障均不静默放行 + + :param redis: Redis 客户端 + :param key: 不含 Token、Secret 或用户名原文的限流 Key + :param limit: 窗口最大请求数 + :param window_seconds: 窗口秒数 + :return: None + :raises RateLimitExceeded: 请求超过限额 + :raises RateLimitUnavailable: Redis 执行失败 + """ + + try: + result = await redis.eval(cls._SCRIPT, 1, key, window_seconds) + count, ttl = int(result[0]), max(1, int(result[1])) + except Exception as exc: + raise RateLimitUnavailable('认证限流服务暂不可用') from exc + if count > limit: + raise RateLimitExceeded(ttl) + + +AfterCommitCallback = Callable[[], Awaitable[None]] + + +class AfterCommitCoordinator: + """ + 提交后副作用协调模块服务层 + + ``register`` 只登记闭包,不会执行;生产调用方必须使用 ``commit`` 或 + ``rollback`` 完成事务,避免在数据库事实落盘前修改缓存 + """ + + def __init__(self) -> None: + """ + 创建空的提交后回调队列 + + :return: None + """ + + self._callbacks: list[AfterCommitCallback] = [] + self._callback_errors: list[Exception] = [] + + @property + def pending_count(self) -> int: + """ + 返回尚未处理的回调数量 + + :return: 尚未执行的回调数量 + """ + + return len(self._callbacks) + + @property + def callback_errors(self) -> tuple[Exception, ...]: + """ + 返回已记录的提交后回调异常 + + :return: 提交后回调异常元组 + """ + + return tuple(self._callback_errors) + + async def register(self, callback: AfterCommitCallback) -> None: + """ + 登记一个只应在数据库提交成功后运行的异步闭包 + + :param callback: 不读取可变 ORM 状态的异步副作用闭包 + :return: None + :raises TypeError: callback 不是可调用对象 + """ + + if not callable(callback): + raise TypeError('事务提交后回调必须为可调用对象') + self._callbacks.append(callback) + + async def commit(self, db: AsyncSession) -> None: + """ + 先提交数据库,再按登记顺序执行副作用 + + :param db: 要提交的异步数据库会话 + :return: None + :raises Exception: 数据库提交失败 + """ + + callbacks, self._callbacks = self._callbacks, [] + try: + await db.commit() + except Exception: + callbacks.clear() + raise + for callback in callbacks: + try: + await callback() + except Exception as exc: # noqa: PERF203 + self._callback_errors.append(exc) + callbacks.clear() + + async def rollback(self, db: AsyncSession) -> None: + """ + 回滚数据库并丢弃全部提交后副作用 + + :param db: 要回滚的异步数据库会话 + :return: None + """ + + self._callbacks.clear() + await db.rollback() diff --git a/ruoyi-fastapi-backend/module_identity/service/interaction_service.py b/ruoyi-fastapi-backend/module_identity/service/interaction_service.py new file mode 100644 index 000000000..78128e0e2 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/interaction_service.py @@ -0,0 +1,1187 @@ +import hmac +import json +import re +import secrets +from collections.abc import Awaitable, Callable, Collection, Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import TYPE_CHECKING, Any +from uuid import uuid4 + +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from common.enums import RedisInitKeyConfig +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException, OidcInteractionException, ServiceException +from module_admin.service.captcha_service import CaptchaService +from module_admin.service.user_service import UserService +from module_identity.dao.identity_user_dao import IdentityUserDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.entity.vo.interaction_vo import ( + CaptchaResponseModel, + ChangePasswordModel, + InteractionLoginModel, + InteractionResultModel, +) +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.audit_service import AuditService +from module_identity.service.identity_service import ( + CredentialAuthenticationError, + CredentialAuthenticationService, + IdentitySecurityEventError, + IdentitySecurityEventService, + IdentitySubjectService, +) +from module_identity.service.infrastructure_service import ( + AfterCommitCoordinator, + OidcRateLimiter, + RateLimitExceeded, + RateLimitUnavailable, +) +from module_identity.service.session_service import SsoSessionError, SsoSessionService +from utils.oidc_util import OidcUtil +from utils.pwd_util import PwdUtil +from utils.time_util import TimezoneUtil + +if TYPE_CHECKING: + from module_identity.entity.do.oauth_client_do import SysOAuthClient + from module_identity.entity.do.oauth_grant_do import SysSsoSession + from module_identity.entity.do.oauth_resource_do import SysOAuthScope + + +_PUBLIC_SCOPE_DESCRIPTIONS = { + 'openid': '用于确认你的身份并登录应用。', + 'profile': '允许应用读取昵称、头像等基本资料。', + 'email': '允许应用读取你的电子邮箱。', + 'phone': '允许应用读取你的手机号码。', + 'roles': '允许应用读取已向它开放的角色信息。', + 'dept': '允许应用读取你的部门信息。', + 'offline_access': '浏览器关闭或登录到期后,仍允许应用在授权有效期内继续访问;你可以撤销授权。', +} + + +@dataclass(frozen=True, slots=True) +class InteractionCreated: + """ + Interaction 创建结果的最小一次性返回值 + + :ivar interaction_id: 新建 Interaction 标识 + :ivar csrf_token: 仅首次返回的原始 CSRF Token + :ivar initial_status: 服务端决定的初始状态 + """ + + interaction_id: str + csrf_token: str + initial_status: str + + +class InteractionService: + """ + 交互流程模块服务层 + + Redis 只保存 CSRF 摘要;原始 CSRF 仅由创建接口返回一次。页面读取接口只返回 + 白名单投影,不建立新 CSRF,也不返回授权协议内部绑定字段 + """ + + _ALLOWED_STATUSES = frozenset( + {'awaiting_login', 'awaiting_consent', 'password_change_required', 'completed', 'denied', 'expired'} + ) + _TERMINAL_STATUSES = frozenset({'completed', 'denied', 'expired'}) + _TRANSITIONS = { + 'awaiting_login': frozenset({'awaiting_consent', 'password_change_required', 'completed', 'denied', 'expired'}), + 'awaiting_consent': frozenset({'completed', 'denied', 'expired'}), + 'password_change_required': frozenset({'awaiting_consent', 'completed', 'denied', 'expired'}), + } + _PROTECTED_FIELDS = frozenset({'requestedAt', 'csrfHash', 'version', 'interactionId'}) + _UPDATABLE_FIELDS = frozenset( + { + 'authenticatedSid', + 'userId', + 'subjectId', + 'authVersion', + 'consentRequired', + 'scopes', + 'grantId', + 'credentialProofHash', + 'rememberMe', + } + ) + _PROMPT_VALUES = frozenset({'none', 'login', 'consent'}) + _MAX_STATE_LENGTH = 1024 + _MAX_SCOPE_LENGTH = 500 + _MAX_SCOPES = 100 + _MAX_INTERACTION_ID_LENGTH = 128 + _MAX_GRANT_ID_LENGTH = 36 + _NO_TTL_RESULT = -5 + _TRANSITION_SCRIPT = """ +local current_json = redis.call('GET', KEYS[1]) +if not current_json then return -1 end +local ok, current = pcall(cjson.decode, current_json) +if not ok or type(current) ~= 'table' then return -4 end +if redis.call('PTTL', KEYS[1]) <= 0 then + redis.call('DEL', KEYS[1]) + return -5 +end +if tonumber(current.version) ~= tonumber(ARGV[1]) then return -2 end +local expected = cjson.decode(ARGV[2]) +local matched = false +for _, status in ipairs(expected) do + if current.status == status then matched = true break end +end +if not matched then return -3 end +redis.call('SET', KEYS[1], ARGV[3], 'KEEPTTL') +return 1 +""" + + @classmethod + async def create( + cls, + redis: Redis, + payload: Mapping[str, Any], + *, + ttl_seconds: int | None = None, + pepper: str | None = None, + ) -> InteractionCreated: + """ + 创建短期 Interaction,并原子占用随机 ID + + :param redis: 异步 Redis 客户端 + :param payload: AuthorizationContext 生成的服务端白名单载荷 + :param ttl_seconds: 可选 Interaction TTL + :param pepper: CSRF 摘要 Pepper,至少 32 bytes + :return: 只含标识、原始 CSRF 和初始状态的冻结结果 + :raises OidcInteractionException: prompt=none 无法静默完成时抛出 + :raises ValueError: payload、Prompt 或 TTL 不合法时抛出 + """ + + record = cls._validate_payload(payload) + prompt = record['prompt'] + authenticated_sid = record.get('authenticatedSid') + has_sso = bool(authenticated_sid) + if 'none' in prompt and not has_sso: + raise OidcInteractionException(message='需要登录后继续', error='login_required', status_code=400) + if 'none' in prompt and bool(record['consentRequired']): + raise OidcInteractionException(message='需要用户确认授权', error='consent_required', status_code=400) + + if 'login' in prompt or not has_sso: + initial_status = 'awaiting_login' + elif bool(record['consentRequired']): + initial_status = 'awaiting_consent' + else: + initial_status = 'completed' + + interaction_id = str(uuid4()) + csrf_token = secrets.token_urlsafe(32) + csrf_hash = OidcUtil.csrf_digest(csrf_token, pepper or OidcConfig.oidc_token_hash_pepper) + record.update( + { + 'interactionId': interaction_id, + 'requestedAt': datetime.now(timezone.utc).isoformat(), + 'csrfHash': csrf_hash, + 'status': initial_status, + 'version': 1, + } + ) + ttl = OidcConfig.oidc_interaction_ttl_seconds if ttl_seconds is None else ttl_seconds + if not isinstance(ttl, int) or isinstance(ttl, bool) or ttl <= 0: + raise ValueError('认证交互有效期必须为正整数') + created = await redis.set( + OidcRedisKey.interaction(interaction_id), + OidcUtil.serialize_json(record, error_message='认证交互载荷必须支持 JSON 序列化'), + ex=ttl, + nx=True, + ) + if not created: + raise OidcInteractionException(message='认证交互创建失败', error='server_error', status_code=500) + return InteractionCreated(interaction_id, csrf_token, initial_status) + + @classmethod + async def get(cls, redis: Redis, interaction_id: str, db: AsyncSession) -> dict[str, Any]: + """ + 从当前启用的应用及权限定义构建页面,不公开管理备注或协议内部字段 + + :param redis: 异步 Redis 客户端 + :param interaction_id: Interaction 标识 + :param db: 默认平台数据库会话 + :return: 页面安全载荷 + :raises OidcInteractionException: Interaction 缺失、过期或状态损坏时抛出 + """ + + record = await cls._get_record(redis, interaction_id) + client = await OAuthClientDao.get_by_pk(db, record['clientPk']) + if client is None or client.client_id != record['clientId']: + raise OidcInteractionException(message='应用已不可用,请返回应用重新登录', status_code=400) + scopes = {scope.scope_code: scope for scope in await OAuthClientDao.list_scopes(db, client.client_pk)} + if not set(record['scopes']).issubset(scopes): + raise OidcInteractionException(message='应用权限已变更,请返回应用重新登录', status_code=400) + ttl = await redis.ttl(OidcRedisKey.interaction(interaction_id)) + + return cls._page_projection(record, ttl, client, scopes) + + @classmethod + async def get_record(cls, redis: Redis, interaction_id: str) -> dict[str, Any]: + """ + 获取仅供后端状态处理使用的完整 Interaction 记录 + + :param redis: 异步 Redis 客户端 + :param interaction_id: Interaction 标识 + :return: 完整内部记录;调用方不得直接返回给页面 + :raises OidcInteractionException: Interaction 缺失或记录损坏时抛出 + """ + + return await cls._get_record(redis, interaction_id) + + @classmethod + async def transition( + cls, + redis: Redis, + interaction_id: str, + expected_statuses: Collection[str], + target_status: str, + updates: Mapping[str, Any] | None = None, + expected_version: int | None = None, + ) -> dict[str, Any]: + """ + 使用 Redis Lua CAS 原子推进 Interaction 状态并保留 TTL + + :param redis: 异步 Redis 客户端 + :param interaction_id: Interaction 标识 + :param expected_statuses: 允许作为当前状态的集合 + :param target_status: 目标状态 + :param updates: 可更新的后端认证字段 + :param expected_version: 调用方读取状态时的版本号 + :return: 更新后的下一步动作与剩余 TTL;页面元数据由 get 读取 + :raises OidcInteractionException: Interaction 缺失或并发状态已变化时抛出 + :raises ValueError: 状态、字段或流转方向不合法时抛出 + """ + + if target_status not in cls._ALLOWED_STATUSES: + raise ValueError('认证交互的目标状态不受支持') + expected = set(expected_statuses) + if not expected or not expected.issubset(cls._ALLOWED_STATUSES): + raise ValueError('认证交互的来源状态无效') + if any(target_status not in cls._TRANSITIONS.get(status, frozenset()) for status in expected): + raise ValueError('不允许执行此认证交互状态变更') + changes = dict(updates or {}) + unsafe = cls._PROTECTED_FIELDS.intersection(changes) + if unsafe: + raise ValueError(f'认证交互的受保护字段不可修改:{sorted(unsafe)}') + if not set(changes).issubset(cls._UPDATABLE_FIELDS): + raise ValueError('认证交互包含不允许更新的字段') + current = await cls._get_record(redis, interaction_id) + compare_version = current.get('version') if expected_version is None else expected_version + if not isinstance(compare_version, int) or isinstance(compare_version, bool) or compare_version <= 0: + raise ValueError('认证交互的预期版本无效') + next_record = dict(current) + next_record.update(changes) + next_record['status'] = target_status + next_record['version'] = int(current.get('version', 0)) + 1 + cls._validate_record(next_record) + serialized = OidcUtil.serialize_json(next_record, error_message='认证交互载荷必须支持 JSON 序列化') + result = await redis.eval( + cls._TRANSITION_SCRIPT, + 1, + OidcRedisKey.interaction(interaction_id), + compare_version, + json.dumps(sorted(expected), separators=(',', ':')), + serialized, + ) + if result == 1: + ttl = await redis.ttl(OidcRedisKey.interaction(interaction_id)) + return { + 'interactionId': interaction_id, + 'nextAction': cls._next_action(target_status), + 'expiresIn': max(0, int(ttl)), + } + if result == -1: + raise OidcInteractionException(message='认证交互不存在或已过期', error='invalid_request', status_code=404) + if result in {-2, -3}: + raise OidcInteractionException( + message='认证交互状态已变更,请刷新后重试', error='invalid_request', status_code=409 + ) + if result == cls._NO_TTL_RESULT: + raise OidcInteractionException(message='认证交互有效期无效', error='invalid_request', status_code=404) + raise OidcInteractionException(message='认证交互状态更新失败', error='server_error', status_code=500) + + @classmethod + def verify_csrf(cls, record: Mapping[str, Any], csrf_token: str, *, pepper: str | None = None) -> bool: + """ + 使用恒定时间比较验证 Interaction CSRF + + :param record: 后端读取的完整 Interaction 记录 + :param csrf_token: 请求携带的原始 CSRF Token + :param pepper: CSRF 摘要 Pepper + :return: 摘要匹配时为 True + """ + + stored = record.get('csrfHash') if isinstance(record, Mapping) else None + if not isinstance(stored, str) or not isinstance(csrf_token, str) or not csrf_token: + return False + try: + actual = OidcUtil.csrf_digest(csrf_token, pepper or OidcConfig.oidc_token_hash_pepper) + except (TypeError, ValueError): + return False + return hmac.compare_digest(actual, stored) + + @classmethod + async def _get_record(cls, redis: Redis, interaction_id: str) -> dict[str, Any]: + """ + 读取并校验后端完整记录 + + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :return: 通过校验的完整 Interaction 记录 + """ + + if ( + not isinstance(interaction_id, str) + or not interaction_id + or len(interaction_id) > cls._MAX_INTERACTION_ID_LENGTH + ): + raise OidcInteractionException(message='认证交互不存在或已过期', error='invalid_request', status_code=404) + key = OidcRedisKey.interaction(interaction_id) + value = await redis.get(key) + if value is None: + raise OidcInteractionException(message='认证交互不存在或已过期', error='invalid_request', status_code=404) + try: + if isinstance(value, bytes): + value = value.decode('utf-8') + record = json.loads(value) + except (TypeError, UnicodeDecodeError, json.JSONDecodeError): + raise OidcInteractionException(message='认证交互状态无效', error='server_error', status_code=500) from None + ttl = await redis.ttl(key) + if ttl <= 0: + await redis.delete(key) + raise OidcInteractionException(message='认证交互不存在或已过期', error='invalid_request', status_code=404) + try: + cls._validate_record(record) + except ValueError as exc: + raise OidcInteractionException(message='认证交互状态无效', error='server_error', status_code=500) from exc + return record + + @classmethod + def _validate_record(cls, record: Mapping[str, Any]) -> None: + """ + 校验 Redis 中 Interaction 的完整结构和身份字段组合 + + :param record: 交互记录 + :return: None + """ + + if not isinstance(record, Mapping): + raise ValueError('认证交互记录必须为映射对象') + required = {'interactionId', 'requestedAt', 'csrfHash', 'status', 'version'} + if not required.issubset(record): + raise ValueError('认证交互记录不完整') + if not isinstance(record['status'], str) or record['status'] not in cls._ALLOWED_STATUSES: + raise ValueError('认证交互状态标识无效') + if not isinstance(record['version'], int) or isinstance(record['version'], bool) or record['version'] <= 0: + raise ValueError('认证交互版本无效') + if not isinstance(record['csrfHash'], str) or not re.fullmatch(r'[0-9a-f]{64}', record['csrfHash']): + raise ValueError('认证交互的 CSRF 摘要无效') + requested_at = record['requestedAt'] + if not isinstance(requested_at, str): + raise ValueError('认证交互的请求时间 requestedAt 无效') + try: + parsed = datetime.fromisoformat(requested_at.replace('Z', '+00:00')) + except ValueError as exc: + raise ValueError('认证交互的请求时间 requestedAt 无效') from exc + if parsed.tzinfo is None or parsed.utcoffset() != timezone.utc.utcoffset(parsed): + raise ValueError('认证交互的请求时间 requestedAt 必须使用 UTC') + cls._validate_payload(record) + identity_values = [record.get('userId'), record.get('subjectId'), record.get('authVersion')] + if any(value is not None for value in identity_values) and not all( + value is not None for value in identity_values + ): + raise ValueError('认证交互的用户身份字段必须完整') + + @classmethod + def _validate_payload(cls, payload: Mapping[str, Any]) -> dict[str, Any]: # noqa: PLR0912, PLR0915 + """ + 校验 AuthorizationContext 生成的 Interaction 载荷 + + :param payload: Interaction 载荷 + :return: 通过校验的交互载荷 + """ + + if not isinstance(payload, Mapping): + raise ValueError('认证交互载荷必须为映射对象') + allowed = { + 'interactionId', + 'requestedAt', + 'csrfHash', + 'status', + 'version', + 'clientPk', + 'clientId', + 'redirectUri', + 'responseType', + 'scopes', + 'resources', + 'state', + 'nonce', + 'codeChallenge', + 'codeChallengeMethod', + 'prompt', + 'maxAge', + 'consentRequired', + 'grantId', + 'authenticatedSid', + 'userId', + 'subjectId', + 'authVersion', + 'credentialProofHash', + 'rememberMe', + } + if not set(payload).issubset(allowed): + raise ValueError('认证交互载荷包含未知字段') + if 'interactionId' in payload and ( + not isinstance(payload['interactionId'], str) + or not payload['interactionId'] + or len(payload['interactionId']) > cls._MAX_INTERACTION_ID_LENGTH + ): + raise ValueError('认证交互标识 interactionId 无效') + required = { + 'clientPk', + 'clientId', + 'redirectUri', + 'responseType', + 'scopes', + 'resources', + 'state', + 'nonce', + 'codeChallenge', + 'codeChallengeMethod', + 'prompt', + 'maxAge', + 'consentRequired', + } + if not required.issubset(payload): + raise ValueError('认证交互载荷缺少必填字段') + if ( + not isinstance(payload['clientPk'], int) + or isinstance(payload['clientPk'], bool) + or payload['clientPk'] <= 0 + ): + raise ValueError('客户端主键 clientPk 必须为正整数') + for field, limit in (('clientId', 64), ('redirectUri', 1000), ('nonce', 1024), ('codeChallenge', 128)): + if not isinstance(payload[field], str) or not payload[field] or len(payload[field]) > limit: + raise ValueError(f'{field} 无效') + if payload['responseType'] != 'code' or payload['codeChallengeMethod'] != 'S256': + raise ValueError('认证交互包含不支持的协议字段') + if not OidcUtil.is_s256_challenge(payload['codeChallenge']): + raise ValueError('PKCE 挑战值 codeChallenge 无效') + if ( + not isinstance(payload['scopes'], (list, tuple)) + or len(payload['scopes']) > cls._MAX_SCOPES + or not all(isinstance(item, str) for item in payload['scopes']) + ): + raise ValueError('权限范围必须为字符串列表') + if 'openid' not in payload['scopes'] or len(set(payload['scopes'])) != len(payload['scopes']): + raise ValueError('权限范围必须包含 openid,且不得重复') + if ( + not isinstance(payload['resources'], (list, tuple)) + or len(payload['resources']) > 1 + or not all(isinstance(item, str) for item in payload['resources']) + ): + raise ValueError('资源列表最多只能包含一个字符串') + if ( + 'grantId' in payload + and payload['grantId'] is not None + and ( + not isinstance(payload['grantId'], str) + or not payload['grantId'] + or len(payload['grantId']) > cls._MAX_GRANT_ID_LENGTH + ) + ): + raise ValueError('授权记录标识 grantId 无效') + if any(not item or len(item) > cls._MAX_SCOPE_LENGTH for item in (*payload['scopes'], *payload['resources'])): + raise ValueError('权限范围或资源值无效') + if payload['state'] is not None and ( + not isinstance(payload['state'], str) or len(payload['state']) > cls._MAX_STATE_LENGTH + ): + raise ValueError('授权请求状态 state 无效') + if payload['maxAge'] is not None and ( + not isinstance(payload['maxAge'], int) or isinstance(payload['maxAge'], bool) or payload['maxAge'] < 0 + ): + raise ValueError('认证新鲜度参数 maxAge 无效') + if 'rememberMe' in payload and not isinstance(payload['rememberMe'], bool): + raise ValueError('保持登录标识 rememberMe 无效') + if not isinstance(payload['consentRequired'], bool): + raise ValueError('授权确认标识 consentRequired 无效') + for field, limit in (('authenticatedSid', 36), ('subjectId', 36)): + if ( + field in payload + and payload[field] is not None + and (not isinstance(payload[field], str) or not payload[field] or len(payload[field]) > limit) + ): + raise ValueError(f'{field} 无效') + if ( + 'userId' in payload + and payload['userId'] is not None + and ( + not isinstance(payload['userId'], int) or isinstance(payload['userId'], bool) or payload['userId'] <= 0 + ) + ): + raise ValueError('用户编号 userId 无效') + if ( + 'authVersion' in payload + and payload['authVersion'] is not None + and ( + not isinstance(payload['authVersion'], int) + or isinstance(payload['authVersion'], bool) + or payload['authVersion'] < 0 + ) + ): + raise ValueError('身份安全版本 authVersion 无效') + if ( + 'credentialProofHash' in payload + and payload['credentialProofHash'] is not None + and ( + not isinstance(payload['credentialProofHash'], str) + or not re.fullmatch(r'[0-9a-f]{64}', payload['credentialProofHash']) + ) + ): + raise ValueError('凭据证明摘要 credentialProofHash 无效') + prompt = payload['prompt'] + if prompt is None: + prompts: list[str] = [] + elif isinstance(prompt, str): + prompts = prompt.split() + elif isinstance(prompt, (list, tuple)) and all(isinstance(item, str) for item in prompt): + prompts = list(prompt) + else: + raise ValueError('prompt 参数无效') + if ( + len(set(prompts)) != len(prompts) + or any(item not in cls._PROMPT_VALUES for item in prompts) + or ('none' in prompts and len(prompts) > 1) + ): + raise ValueError('prompt 参数组合无效') + result = dict(payload) + result['scopes'] = list(payload['scopes']) + result['resources'] = list(payload['resources']) + result['prompt'] = prompts + if 'authenticatedSid' not in result: + result['authenticatedSid'] = None + return result + + @staticmethod + def _next_action(status: str) -> str: + """ + 将内部交互状态转换为页面的下一步动作 + + :param status: 当前交互状态 + :return: 页面下一步动作 + """ + + return { + 'awaiting_login': 'login', + 'awaiting_consent': 'consent', + 'password_change_required': 'changePassword', + 'completed': 'redirect', + 'denied': 'redirect', + 'expired': 'redirect', + }[status] + + @staticmethod + def _page_projection( + record: Mapping[str, Any], + ttl: int, + client: 'SysOAuthClient', + scopes: Mapping[str, 'SysOAuthScope'], + ) -> dict[str, Any]: + """ + 构建不含协议内部绑定和 CSRF 摘要的页面投影 + + :param record: 交互记录 + :param ttl: 剩余有效期秒数 + :param client: 当前启用的应用 + :param scopes: 当前应用允许的启用权限定义 + :return: 交互页面数据 + """ + + next_action = InteractionService._next_action(record['status']) + + return { + 'interactionId': record['interactionId'], + 'client': { + 'clientId': client.client_id, + 'clientName': client.client_name, + 'logoUri': client.logo_uri, + 'policyUri': client.policy_uri, + 'tosUri': client.tos_uri, + }, + 'requestedScopes': [ + { + 'scope': scope, + 'name': scopes[scope].scope_name, + 'description': _PUBLIC_SCOPE_DESCRIPTIONS.get(scope), + 'sensitive': bool(scopes[scope].sensitive), + 'required': scope == 'openid' or not bool(scopes[scope].consent_required), + } + for scope in record['scopes'] + ], + 'nextAction': next_action, + 'captchaEnabled': False, + 'expiresIn': max(0, int(ttl)), + } + + +@dataclass(frozen=True) +class CaptchaOutcome: + """ + 验证码业务结果;HTTP 状态和响应头由 Controller 决定 + """ + + result: CaptchaResponseModel | None = None + rate_limited: bool = False + retry_after: int | None = None + unavailable: bool = False + + +class InteractionFlowService: + """ + 交互状态模块服务层 + """ + + @staticmethod + async def commit_transition( + db: AsyncSession, + coordinator: AfterCommitCoordinator, + redis: Redis, + interaction_id: str, + target: str, + updates: dict[str, Any] | None = None, + compensate: Callable[[], Awaitable[None]] | None = None, + ) -> None: + """ + 提交数据库事务并推进交互状态 + + :param db: 异步数据库会话 + :param coordinator: 提交后副作用协调器 + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :param target: 交互状态目标值 + :param updates: 状态更新字段 + :param compensate: 失败时执行的补偿回调 + :return: None + """ + + initial = await InteractionService.get_record(redis, interaction_id) + expected_status = initial['status'] + expected_version = initial['version'] + transition_errors: list[Exception] = [] + + async def transition() -> None: + """ + 推进交互状态 + + :return: 状态转换结果 + """ + + try: + await InteractionService.transition( + redis, interaction_id, {expected_status}, target, updates, expected_version=expected_version + ) + except Exception as exc: + transition_errors.append(exc) + raise + + await coordinator.register(transition) + try: + await coordinator.commit(db) + except Exception: + await db.rollback() + raise + if transition_errors: + if compensate is not None: + try: + await compensate() + except Exception: + pass + raise OidcInteractionException( + interaction_id, '认证交互状态更新失败', error='server_error', status_code=500 + ) + if coordinator.callback_errors: + raise OidcInteractionException( + interaction_id, '认证交互缓存更新回调执行失败', error='server_error', status_code=500 + ) + + @staticmethod + async def csrf_record(redis: Redis, interaction_id: str, csrf_token: str | None) -> dict[str, Any]: + """ + 校验 CSRF 并返回交互记录 + + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :param csrf_token: CSRF Token + :return: 已校验的交互记录 + """ + + record = await InteractionService.get_record(redis, interaction_id) + if not InteractionService.verify_csrf(record, csrf_token or '', pepper=OidcConfig.oidc_token_hash_pepper): + raise OidcInteractionException( + interaction_id, 'CSRF 校验失败,请重新发起认证', error='invalid_request', status_code=403 + ) + return record + + @staticmethod + def require_status(record: dict[str, Any], expected: str) -> None: + """ + 校验业务前置条件 + + :param record: 交互记录 + :param expected: 期望的状态 + :return: None + """ + + if record.get('status') != expected: + raise OidcInteractionException( + record.get('interactionId'), + '认证交互状态已变更,请刷新后重试', + error='invalid_request', + status_code=409, + ) + + @staticmethod + async def captcha(redis: Redis, interaction_id: str, client_ip: str | None) -> CaptchaOutcome: + """ + 生成交互验证码响应 + + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :param client_ip: 客户端 IP 地址 + :return: 验证码响应结果 + """ + + try: + await OidcRateLimiter.enforce( + redis, + OidcRedisKey.interaction_captcha_rate_limit( + OidcUtil.hash_sensitive_identifier( + f'{interaction_id}:{client_ip or ""}', + OidcConfig.oidc_token_hash_pepper, + ) + ), + limit=5, + window_seconds=60, + ) + except RateLimitExceeded as exc: + return CaptchaOutcome(rate_limited=True, retry_after=exc.retry_after) + except RateLimitUnavailable: + return CaptchaOutcome(unavailable=True) + await InteractionService.get_record(redis, interaction_id) + enabled = await InteractionFlowService.captcha_enabled(redis) + if enabled: + image, answer = await CaptchaService.create_captcha_image_service() + captcha_id = str(uuid4()) + await redis.set(f'{RedisInitKeyConfig.CAPTCHA_CODES.key}:{captcha_id}', answer, ex=timedelta(minutes=2)) + result = CaptchaResponseModel(captcha_enabled=True, uuid=captcha_id, img=image) + else: + result = CaptchaResponseModel(captcha_enabled=False) + return CaptchaOutcome(result=result) + + @staticmethod + async def captcha_enabled(redis: Redis) -> bool: + """ + 读取交互验证码开关 + + :param redis: Redis 客户端 + :return: 是否启用验证码 + """ + + value = await redis.get(f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.account.captchaEnabled') + + return value in {'true', b'true', True} + + @staticmethod + async def reserve_completion(redis: Redis, interaction_id: str) -> str | None: + """ + 保留交互完成标记 Key + + :param redis: Redis 客户端 + :param interaction_id: 交互流程标识 + :return: 完成标记 Key 或 None + """ + + ttl = await redis.ttl(OidcRedisKey.interaction(interaction_id)) + if not isinstance(ttl, int) or ttl <= 0: + return None + marker = OidcRedisKey.interaction(f'{interaction_id}-completion') + reserved = await redis.set(marker, 'reserved', ex=ttl, nx=True) + + return marker if reserved else None + + @staticmethod + def interaction_result(interaction_id: str, next_action: str) -> InteractionResultModel: + """ + 构建交互结果响应 + + :param interaction_id: 交互流程标识 + :param next_action: 下一步交互动作 + :return: 交互响应 + """ + + return InteractionResultModel( + next_action=next_action, + interaction_id=interaction_id, + redirect_url=f'/auth/interaction/{interaction_id}/complete' if next_action == 'redirect' else None, + ) + + @staticmethod + async def best_effort_delete(redis: Redis, key: str) -> None: + """ + 尽力删除 Redis Key 并吞掉删除异常 + + :param redis: Redis 客户端 + :param key: Redis Key + :return: None + """ + + try: + await redis.delete(key) + except Exception: + pass + + @staticmethod + async def rollback(db: AsyncSession) -> None: + """ + 回滚数据库事务 + + :param db: 异步数据库会话 + :return: None + """ + + try: + await db.rollback() + except Exception: + pass + + +@dataclass(frozen=True) +class InteractionLoginOutcome: + """ + 登录/改密业务结果;Cookie 和 JSON 响应由 Controller 写出 + """ + + result: InteractionResultModel | None = None + cookie: str | None = None + failure_message: str | None = None + cookie_max_age: int | None = None + + +class InteractionLoginService: + """ + 交互登录模块服务层 + """ + + @staticmethod + async def login( + redis: Redis, + interaction_id: str, + body: InteractionLoginModel, + db: AsyncSession, + csrf_token: str | None, + client_ip: str | None = None, + user_agent: str | None = None, + ) -> InteractionLoginOutcome: + """ + 校验本地凭据并推进认证交互 + + :param redis: 交互 Redis + :param interaction_id: Interaction 标识 + :param body: 登录凭据和验证码参数 + :param db: 异步数据库会话 + :param csrf_token: Interaction CSRF 原文 + :param client_ip: 客户端 IP 地址 + :param user_agent: 客户端 User-Agent + :return: 下一步交互动作 + """ + + record = await InteractionFlowService.csrf_record(redis, interaction_id, csrf_token) + InteractionFlowService.require_status(record, 'awaiting_login') + try: + result = await CredentialAuthenticationService.authenticate_oidc( + redis, + db, + client_ip=client_ip, + user_name=body.user_name, + password=body.password, + code=body.code, + uuid=body.uuid, + captcha_enabled=await InteractionFlowService.captcha_enabled(redis), + remember_me=body.remember_me, + ) + except CredentialAuthenticationError: + await AuditService.record_interaction_failure( + db, + OidcAuditEvent.LOGIN_FAILED, + client_id=record.get('clientId'), + failure_code='invalid_credentials', + ) + return InteractionLoginOutcome(failure_message='登录失败') + + subject = await IdentitySubjectService.require_by_user_id( + db, result.user.user_id, audit_writer=AuditService.interaction_subject_writer(db) + ) + coordinator = AfterCommitCoordinator() + cookie: str | None = None + session: SysSsoSession | None = None + target_status = ( + 'password_change_required' + if result.password_change_required + else 'awaiting_consent' + if bool(record.get('consentRequired')) + else 'completed' + ) + updates: dict[str, Any] = { + 'authenticatedSid': None, + 'userId': result.user.user_id, + 'subjectId': subject.subject_id, + 'authVersion': subject.auth_version, + 'rememberMe': result.remember_me, + } + if result.password_change_required: + updates['credentialProofHash'] = OidcUtil.credential_proof( + record['interactionId'], + result.user.user_id, + subject.subject_id, + subject.auth_version, + pepper=OidcConfig.oidc_token_hash_pepper, + ) + else: + cookie, session = await SsoSessionService.create( + db, + redis, + result.user.user_id, + subject.subject_id, + subject.auth_version, + result.acr, + result.amr, + ip_address=client_ip, + user_agent=user_agent, + remember_me=result.remember_me, + pepper=OidcConfig.oidc_token_hash_pepper, + coordinator=coordinator, + ) + updates['authenticatedSid'] = session.sid + await AuditService.record( + db, + OidcAuditEvent.LOGIN_SUCCEEDED, + 'success', + client_id=record.get('clientId'), + user_id=result.user.user_id, + subject_id=str(subject.subject_id), + sid=session.sid if session is not None else None, + ) + + async def compensate() -> None: + """ + 执行事务补偿 + + :return: None + """ + + if session is None: + return + await InteractionFlowService.best_effort_delete(redis, OidcRedisKey.interaction(interaction_id)) + try: + cleanup = AfterCommitCoordinator() + await SsoSessionService.revoke( + db, redis, session.sid, reason='interaction_transition_failed', coordinator=cleanup + ) + await cleanup.commit(db) + except Exception: + pass + try: + await AuditService.record_independent( + db, + OidcAuditEvent.LOGIN_FAILED, + 'failure', + risk_level='high', + client_id=record.get('clientId'), + user_id=result.user.user_id, + failure_code='interaction_transition_failed', + ) + except Exception: + pass + + await InteractionFlowService.commit_transition( + db, coordinator, redis, interaction_id, target_status, updates, compensate=compensate + ) + next_action = ( + 'changePassword' + if target_status == 'password_change_required' + else ('consent' if target_status == 'awaiting_consent' else 'redirect') + ) + model = InteractionResultModel( + next_action=next_action, + interaction_id=interaction_id, + redirect_url=f'/auth/interaction/{interaction_id}/complete' if target_status == 'completed' else None, + reason=result.password_change_reason if target_status == 'password_change_required' else None, + ) + + return InteractionLoginOutcome( + result=model, + cookie=cookie, + cookie_max_age=SsoSessionService.cookie_max_age(session) if session is not None else None, + ) + + @staticmethod + async def change_password( + redis: Redis, + interaction_id: str, + body: ChangePasswordModel, + db: AsyncSession, + csrf_token: str | None, + ) -> InteractionLoginOutcome: + """ + 修改初始或过期密码并重新建立当前 SSO Session + + :param redis: 交互 Redis + :param interaction_id: Interaction 标识 + :param body: 原密码、新密码和确认密码 + :param db: 异步数据库会话 + :param csrf_token: Interaction CSRF 原文 + :return: 下一步交互动作 + """ + + record = await InteractionFlowService.csrf_record(redis, interaction_id, csrf_token) + InteractionFlowService.require_status(record, 'password_change_required') + user = await IdentityUserDao.get_active_user(db, int(record.get('userId') or 0)) + proof = OidcUtil.credential_proof( + record['interactionId'], + int(record.get('userId') or 0), + str(record.get('subjectId') or ''), + int(record.get('authVersion') or 0), + pepper=OidcConfig.oidc_token_hash_pepper, + ) + if ( + user is None + or user.status != '0' + or user.del_flag != '0' + or not PwdUtil.verify_password(body.old_password, user.password) + or not hmac.compare_digest(str(record.get('credentialProofHash') or ''), proof) + ): + await AuditService.record_interaction_failure( + db, + OidcAuditEvent.LOGIN_FAILED, + client_id=record.get('clientId'), + user_id=record.get('userId'), + failure_code='password_change_failed', + ) + return InteractionLoginOutcome(failure_message='改密失败') + if PwdUtil.verify_password(body.new_password, user.password): + await AuditService.record_interaction_failure( + db, + OidcAuditEvent.LOGIN_FAILED, + client_id=record.get('clientId'), + user_id=record.get('userId'), + failure_code='password_reuse', + ) + return InteractionLoginOutcome(failure_message='新密码不能与旧密码相同') + + try: + await UserService.validate_password_services(redis, body.new_password) + user.password = PwdUtil.get_password_hash(body.new_password) + user.pwd_update_date = TimezoneUtil.utc_now() + await IdentitySubjectService.require_by_user_id( + db, user.user_id, audit_writer=AuditService.interaction_subject_writer(db) + ) + coordinator = AfterCommitCoordinator() + await SsoSessionService.revoke_user( + db, redis, user.user_id, reason='password_changed', coordinator=coordinator + ) + await IdentitySecurityEventService.handle_user_event( + db, user.user_id, 'password_changed', now=TimezoneUtil.utc_now() + ) + subject = await IdentitySubjectService.require_by_user_id( + db, user.user_id, audit_writer=AuditService.interaction_subject_writer(db) + ) + cookie, session = await SsoSessionService.create( + db, + redis, + user.user_id, + subject.subject_id, + subject.auth_version, + 'urn:ruoyi:acr:pwd', + ('pwd',), + remember_me=record.get('rememberMe', False), + pepper=OidcConfig.oidc_token_hash_pepper, + coordinator=coordinator, + ) + target_status = 'awaiting_consent' if bool(record.get('consentRequired')) else 'completed' + await AuditService.record( + db, + OidcAuditEvent.LOGIN_SUCCEEDED, + 'success', + client_id=record.get('clientId'), + user_id=user.user_id, + subject_id=str(subject.subject_id), + sid=session.sid, + ) + + async def compensate() -> None: + """ + 执行事务补偿 + + :return: None + """ + + await InteractionFlowService.best_effort_delete(redis, OidcRedisKey.interaction(interaction_id)) + try: + cleanup = AfterCommitCoordinator() + await SsoSessionService.revoke( + db, redis, session.sid, reason='interaction_transition_failed', coordinator=cleanup + ) + await cleanup.commit(db) + except Exception: + pass + try: + await AuditService.record_independent( + db, + OidcAuditEvent.LOGIN_FAILED, + 'failure', + risk_level='high', + client_id=record.get('clientId'), + user_id=user.user_id, + failure_code='password_change_saga_failed', + ) + except Exception: + pass + + await InteractionFlowService.commit_transition( + db, + coordinator, + redis, + interaction_id, + target_status, + { + 'authenticatedSid': session.sid, + 'userId': user.user_id, + 'subjectId': subject.subject_id, + 'authVersion': subject.auth_version, + }, + compensate=compensate, + ) + except (OAuthProtocolException, ServiceException, SsoSessionError, IdentitySecurityEventError): + await AuditService.record_interaction_failure( + db, + OidcAuditEvent.LOGIN_FAILED, + client_id=record.get('clientId'), + user_id=record.get('userId'), + failure_code='password_change_failed', + ) + return InteractionLoginOutcome(failure_message='改密失败') + model = InteractionResultModel( + next_action='consent' if target_status == 'awaiting_consent' else 'redirect', + interaction_id=interaction_id, + redirect_url=f'/auth/interaction/{interaction_id}/complete' if target_status == 'completed' else None, + ) + + return InteractionLoginOutcome( + result=model, cookie=cookie, cookie_max_age=SsoSessionService.cookie_max_age(session) + ) diff --git a/ruoyi-fastapi-backend/module_identity/service/key_service.py b/ruoyi-fastapi-backend/module_identity/service/key_service.py new file mode 100644 index 000000000..d043882f2 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/key_service.py @@ -0,0 +1,862 @@ +import asyncio +import inspect +import secrets +from collections.abc import AsyncIterator, Awaitable, Callable +from contextlib import asynccontextmanager +from datetime import datetime, timedelta +from typing import Any, Literal, TypeVar + +from anyio import Path as AsyncPath +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import ServiceException +from module_identity.dao.oidc_key_dao import OidcKeyDao +from module_identity.entity.do.oidc_key_do import SysOidcSigningKey +from module_identity.entity.vo.oidc_key_vo import OidcKeyRotateModel, OidcKeyViewModel +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.audit_service import AuditService +from utils.log_util import logger +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + +RS256 = 'RS256' +_PUBLISHED_STATUSES = frozenset({'pending', 'active', 'retiring'}) +_MIN_RSA_BITS = 2048 +_MIN_ENCRYPTION_KEY_BYTES = 32 +_ENCRYPTION_SALT_BYTES = 16 +_ROTATION_LOCK_TTL_SECONDS = 30 +_ROTATION_LOCK_RELEASE_SCRIPT = """ +if redis.call('get', KEYS[1]) == ARGV[1] then + return redis.call('del', KEYS[1]) +end +return 0 +""" + +T = TypeVar('T') + + +class KeyServiceError(ValueError): + """ + 表示 OIDC 签名密钥管理错误 + """ + + +class KeyService: + """ + OIDC 签名密钥模块服务层 + """ + + CACHE_MAX_AGE = 300 + + @classmethod + def _require_enabled(cls, *, management: bool = False) -> None: + """ + 校验 OIDC 签名密钥功能是否可用 + + :param management: 是否允许管理模式 + :return: None + :raises KeyServiceError: OIDC 未启用 + """ + + if not OidcConfig.oidc_enabled and not management: + raise KeyServiceError('统一认证中心未启用') + + @staticmethod + @asynccontextmanager + async def _rotation_lock(redis: Any) -> AsyncIterator[None]: + """ + 获取密钥轮换锁 + + :param redis: 应用共享 Redis 客户端;离线数据库单元测试可传 ``None`` + :return: 异步上下文管理器 + :raises KeyServiceError: 锁已被其他实例持有 + """ + + if redis is None: + yield + return + token = secrets.token_urlsafe(24) + acquired = await redis.set( + OidcRedisKey.signing_key_rotation_lock(), token, nx=True, ex=_ROTATION_LOCK_TTL_SECONDS + ) + if not acquired: + raise KeyServiceError('签名密钥正在轮换,请稍后重试') + try: + yield + finally: + await redis.eval( + _ROTATION_LOCK_RELEASE_SCRIPT, + 1, + OidcRedisKey.signing_key_rotation_lock(), + token, + ) + + @staticmethod + async def _private_material_async( # noqa: PLR0912 + record: Any, + decrypt_private_key: Callable[[str], bytes | str | Awaitable[bytes | str]] | None, + ) -> bytes: + """ + 异步读取私钥材料 + + :param record: 数据库密钥记录 + :param decrypt_private_key: 可同步或异步的解密回调 + :return: PEM 编码私钥字节 + :raises KeyServiceError: 私钥来源不合法或无法解密 + """ + + reference = getattr(record, 'private_key_ref', None) + ciphertext = getattr(record, 'private_key_ciphertext', None) + if bool(reference) == bool(ciphertext): + raise KeyServiceError('必须且只能配置一个签名私钥来源') + if ciphertext: + if decrypt_private_key is None: + try: + value = OidcUtil.decrypt_signing_private_key( + str(ciphertext), OidcConfig.oidc_signing_key_encryption_key.encode() + ) + except Exception as exc: + raise KeyServiceError('解密签名私钥需要有效的加密密钥') from exc + else: + value = decrypt_private_key(str(ciphertext)) + if inspect.isawaitable(value): + value = await value + if isinstance(value, str): + value = value.encode() + if not isinstance(value, bytes) or not value: + raise KeyServiceError('签名私钥解密器返回的数据无效') + return value + if OidcConfig.oidc_signing_key_source != 'file': + if decrypt_private_key is None: + raise KeyServiceError('外部签名私钥引用必须配置加载器') + value = decrypt_private_key(str(reference)) + if inspect.isawaitable(value): + value = await value + if isinstance(value, str): + value = value.encode() + if not isinstance(value, bytes) or not value: + raise KeyServiceError('签名私钥加载器返回的数据无效') + return value + path_value = str(reference or OidcConfig.oidc_signing_private_key_path).strip() + if not path_value: + raise KeyServiceError('签名私钥文件路径不能为空') + try: + return await AsyncPath(path_value).read_bytes() + except OSError as exc: + raise KeyServiceError('无法读取签名私钥文件') from exc + + @classmethod + def _validate_record_window(cls, record: Any, now: datetime, *, require_active: bool = True) -> None: + """ + 校验签名密钥状态、算法和时间窗口 + + :param record: 数据库密钥记录 + :param now: 当前项目时间 + :param require_active: 是否要求密钥处于 active 状态 + :return: None + :raises KeyServiceError: 记录不能用于签名 + """ + + expected_status = 'active' if require_active else 'pending' + if getattr(record, 'status', None) != expected_status: + raise KeyServiceError(f'签名密钥状态不符合要求,预期状态为 {expected_status}') + if getattr(record, 'alg', None) != RS256 or OidcConfig.oidc_signing_algorithm != RS256: + raise KeyServiceError('仅支持 RS256 签名密钥') + if not OidcUtil.is_valid_kid(getattr(record, 'kid', None)): + raise KeyServiceError('签名密钥标识 kid 包含不允许的字符') + start = TimezoneUtil.to_optional_utc(getattr(record, 'signing_start_at', None)) + stop = TimezoneUtil.to_optional_utc(getattr(record, 'signing_stop_at', None)) + if require_active and (start is None or start > now or (stop is not None and stop <= now)): + raise KeyServiceError('签名密钥不在有效签发时间范围内') + + @classmethod + async def load_private_key_async( + cls, + record: Any, + *, + now: datetime | None = None, + decrypt_private_key: Callable[[str], bytes | str | Awaitable[bytes | str]] | None = None, + require_active: bool = True, + management: bool = False, + ) -> RSAPrivateKey: + """ + 异步加载并校验签名密钥 RSA 私钥 + + :param record: 数据库中的签名密钥记录 + :param now: 可注入的当前时间 + :param decrypt_private_key: 解密 ciphertext 的可注入回调 + :param require_active: 是否要求密钥处于 active 状态 + :param management: 是否允许在 OIDC 关闭时执行管理校验 + :return: 已验证且不会被序列化的 RSA 私钥 + :raises KeyServiceError: 状态、来源、格式或公私钥不匹配 + """ + + cls._require_enabled(management=management) + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + cls._validate_record_window(record, current, require_active=require_active) + material = await cls._private_material_async(record, decrypt_private_key) + try: + private_key = serialization.load_pem_private_key(material, password=None) + except (TypeError, ValueError) as exc: + raise KeyServiceError('RSA 签名私钥无效') from exc + if not isinstance(private_key, RSAPrivateKey): + raise KeyServiceError('签名密钥必须为 RSA 类型') + if private_key.key_size < _MIN_RSA_BITS: + raise KeyServiceError('RSA 签名密钥长度不得小于 2048 位') + derived = OidcUtil.rsa_public_jwk(private_key.public_key(), getattr(record, 'kid', '')) + try: + expected = OidcUtil.normalize_public_jwk(record, min_rsa_bits=_MIN_RSA_BITS) + except ValueError as exc: + raise KeyServiceError(str(exc)) from exc + if derived != expected: + raise KeyServiceError('签名私钥与公开 JWK 不匹配') + return private_key + + @classmethod + def load_private_key( + cls, + record: Any, + *, + now: datetime | None = None, + require_active: bool = True, + ) -> RSAPrivateKey: + """ + 同步加载签名密钥 RSA 私钥 + + :param record: 数据库中的签名密钥记录 + :param now: 可注入的当前时间 + :param require_active: 是否要求密钥处于 active 状态 + :return: 已验证的 RSA 私钥 + :raises KeyServiceError: 私钥不合法或无法同步加载 + """ + + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(cls.load_private_key_async(record, now=now, require_active=require_active)) + raise KeyServiceError('事件循环中不能同步加载签名私钥') + + @classmethod + async def get_signing_key( + cls, + db: AsyncSession, + *, + now: datetime | None = None, + decrypt_private_key: Callable[[str], bytes | str | Awaitable[bytes | str]] | None = None, + ) -> RSAPrivateKey: + """ + 获取当前 active 签名密钥的 RSA 私钥 + + :param db: 异步数据库会话 + :param now: 可注入当前时间 + :param decrypt_private_key: 可注入的私钥解密回调 + :return: 已验证 RSA 私钥 + :raises KeyServiceError: 没有匹配的 active 密钥 + """ + + cls._require_enabled() + record = await OidcKeyDao.get_active(db, alg=RS256) + if record is None: + raise KeyServiceError('尚无可用的活动签名密钥') + return await cls.load_private_key_async(record, now=now, decrypt_private_key=decrypt_private_key) + + @classmethod + async def build_jwks(cls, db: AsyncSession, *, now: datetime | None = None) -> dict[str, list[dict[str, str]]]: + """ + 构造发布中的签名公钥 JWKS + + :param db: 异步数据库会话 + :param now: 可注入当前时间 + :return: 标准 ``{'keys': [...]}`` JSON 结构 + :raises KeyServiceError: OIDC 关闭或公开密钥记录不合法 + """ + + cls._require_enabled() + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + records = await OidcKeyDao.list_published(db, now=current) + keys: list[dict[str, str]] = [] + for record in records: + publish_at = TimezoneUtil.to_optional_utc(getattr(record, 'publish_at', None)) + remove_at = TimezoneUtil.to_optional_utc(getattr(record, 'remove_from_jwks_at', None)) + if ( + getattr(record, 'status', None) not in _PUBLISHED_STATUSES + or getattr(record, 'alg', None) != RS256 + or publish_at is None + or publish_at > current + or (remove_at is not None and remove_at <= current) + ): + continue + try: + keys.append(OidcUtil.normalize_public_jwk(record, min_rsa_bits=_MIN_RSA_BITS)) + except ValueError as exc: + raise KeyServiceError(str(exc)) from exc + keys.sort(key=lambda item: item['kid']) + + return {'keys': keys} + + @classmethod + async def activate_key( + cls, + db: AsyncSession, + kid: str, + *, + now: datetime | None = None, + decrypt_private_key: Callable[[str], bytes | str | Awaitable[bytes | str]] | None = None, + actor: str = 'system:lifecycle', + redis: Any = None, + ) -> bool: + """ + 激活 pending 签名密钥并安排旧密钥退役 + + :param db: 异步数据库会话,调用方负责提交事务 + :param kid: 待激活的 pending 密钥 kid + :param now: 可注入当前时间 + :param decrypt_private_key: 可注入的私钥引用解密/加载回调 + :param actor: 操作人标识 + :param redis: 应用共享 Redis,用于跨实例轮换互斥;离线单元测试可省略 + :return: 成功完成切换时返回 True + :raises KeyServiceError: 目标密钥未到发布时间或算法不符 + """ + + cls._require_enabled(management=True) + if not OidcUtil.is_valid_kid(kid): + raise KeyServiceError('签名密钥标识 kid 包含不允许的字符') + async with cls._rotation_lock(redis): + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + await OidcKeyDao.lock_algorithm_for_update(db, alg=RS256) + target = await OidcKeyDao.get_by_kid_for_update(db, kid) + if target is None or target.alg != RS256 or target.status != 'pending': + raise KeyServiceError('目标签名密钥不处于待激活状态') + publish_at = TimezoneUtil.to_optional_utc(target.publish_at) + if publish_at is None or publish_at > current: + raise KeyServiceError('目标签名密钥尚未发布') + old = await OidcKeyDao.get_active(db, alg=RS256, for_update=True) + await cls.load_private_key_async( + target, + now=current, + decrypt_private_key=decrypt_private_key, + require_active=False, + management=True, + ) + changed = await OidcKeyDao.activate(db, kid, alg=RS256, now=current) + if changed and old is not None and old.kid != kid: + retention_at = current + timedelta( + seconds=max( + OidcConfig.oidc_key_rotation_overlap_seconds, + max( + OidcConfig.oidc_access_token_ttl_seconds, + OidcConfig.oidc_max_access_token_ttl_seconds, + OidcConfig.oidc_id_token_ttl_seconds, + ) + + OidcConfig.oidc_allowed_clock_skew_seconds, + ) + ) + previous_remove_at = TimezoneUtil.to_optional_utc(old.remove_from_jwks_at) + if previous_remove_at is None or previous_remove_at < retention_at: + await OidcKeyDao.set_retiring(db, old.kid, retention_at, current) + if changed: + await AuditService.record( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'success', + detail={'action': 'activated', 'kid': kid, 'actor': actor[:64]}, + ) + return changed + + @classmethod + async def create_pending_key( + cls, + db: AsyncSession, + *, + kid: str, + publish_at: datetime, + activate_at: datetime | None = None, + actor: str, + remark: str | None = None, + now: datetime | None = None, + ) -> SysOidcSigningKey: + """ + 创建加密保存的 pending 签名密钥记录 + + :param db: 异步数据库会话,提交边界由调用方控制 + :param kid: 新密钥标识 + :param publish_at: JWKS 发布时间 + :param activate_at: 计划签名开始时间 + :param actor: 管理员安全标识 + :param remark: 管理备注 + :param now: 可注入当前项目时间 + :return: 不含私钥明文的数据库实体 + :raises KeyServiceError: 配置、标识或密钥加密条件不满足 + """ + + cls._require_enabled(management=True) + if not isinstance(actor, str) or not actor.strip(): + raise KeyServiceError('签名密钥标识 kid 和操作者不能为空') + if not OidcUtil.is_valid_kid(kid): + raise KeyServiceError('签名密钥标识 kid 包含不允许的字符') + encryption_material = str(OidcConfig.oidc_signing_key_encryption_key or '').encode() + if len(encryption_material) < _MIN_ENCRYPTION_KEY_BYTES: + raise KeyServiceError('签名私钥加密密钥不能为空') + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + publish_at = TimezoneUtil.to_optional_utc(publish_at) + activate_at = TimezoneUtil.to_optional_utc(activate_at) + if activate_at is not None and activate_at < publish_at: + raise KeyServiceError('密钥生效时间不得早于发布时间') + if publish_at < current: + raise KeyServiceError('签名密钥发布时间不得早于当前时间') + await OidcKeyDao.lock_algorithm_for_update(db, alg=RS256) + existing = await OidcKeyDao.get_by_kid_for_update(db, kid) + if existing is not None: + raise KeyServiceError('签名密钥标识 kid 已存在') + private_key = rsa.generate_private_key(public_exponent=65537, key_size=_MIN_RSA_BITS) + public_jwk = OidcUtil.rsa_public_jwk(private_key.public_key(), kid) + pem = private_key.private_bytes( + serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption() + ) + salt = secrets.token_bytes(_ENCRYPTION_SALT_BYTES) + nonce = secrets.token_bytes(12) + ciphertext = OidcUtil.encrypt_signing_private_key(pem, encryption_material, salt=salt, nonce=nonce) + record = SysOidcSigningKey( + kid=kid, + key_use='sig', + alg=RS256, + public_jwk=public_jwk, + private_key_ciphertext=ciphertext, + status='pending', + publish_at=publish_at, + signing_start_at=activate_at or publish_at, + create_by=actor[:64], + create_time=current, + remark=remark[:500] if isinstance(remark, str) else None, + ) + record = await OidcKeyDao.create(db, record) + await AuditService.record( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'success', + detail={'action': 'created_pending', 'kid': kid, 'actor': actor[:64]}, + ) + + return record + + @classmethod + async def bootstrap_signing_key( + cls, + db: AsyncSession, + *, + kid: str, + actor: str, + redis: object, + now: datetime | None = None, + ) -> tuple[SysOidcSigningKey, bool]: + """ + 幂等创建并激活首把 OIDC 签名密钥 + + 已存在可用 active 密钥时只校验并返回;否则在 Redis 轮换锁与 + 数据库行锁保护下创建和激活密钥。事务提交由调用方统一处理。 + + :param db: 异步数据库会话 + :param kid: 首把签名密钥标识 + :param actor: 部署操作人标识 + :param redis: 异步 Redis 客户端 + :param now: 可注入的当前项目时间 + :return: 签名密钥记录与本次是否创建新密钥 + :raises KeyServiceError: 密钥配置、状态或材料不可用 + """ + + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + async with cls._rotation_lock(redis): + records = await OidcKeyDao.lock_algorithm_for_update(db, alg=RS256) + active = next( + ( + row + for row in reversed(records) + if row.status == 'active' + and TimezoneUtil.to_optional_utc(row.signing_start_at) is not None + and TimezoneUtil.to_optional_utc(row.signing_start_at) <= current + ), + None, + ) + if active is not None: + await cls.load_private_key_async(active, now=current, management=True) + return active, False + + existing = await OidcKeyDao.get_by_kid_for_update(db, kid) + created = existing is None + if existing is None: + existing = await cls.create_pending_key( + db, + kid=kid, + publish_at=current, + activate_at=current, + actor=actor, + remark='OIDC 部署初始化', + now=current, + ) + elif existing.status != 'pending': + raise KeyServiceError('初始化签名密钥不处于待激活状态') + + await cls.load_private_key_async( + existing, + now=current, + require_active=False, + management=True, + ) + if not await OidcKeyDao.activate(db, kid, alg=RS256, now=current): + raise KeyServiceError('初始化签名密钥激活失败') + await AuditService.record( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'success', + detail={'action': 'activated', 'kid': kid, 'actor': actor[:64]}, + ) + return existing, created + + @classmethod + async def retire_key( + cls, + db: AsyncSession, + kid: str, + *, + now: datetime | None = None, + actor: str = 'system:lifecycle', + ) -> bool: + """ + 将 active 签名密钥转为 retiring 状态 + + :param db: 异步数据库会话 + :param kid: 密钥标识 + :param now: 当前时间 + :param actor: 操作人标识 + :return: 本次是否将签名密钥转为退役中状态 + :raises KeyServiceError: 签名密钥状态、材料或配置不符合要求 + """ + + cls._require_enabled(management=True) + if not OidcUtil.is_valid_kid(kid): + raise KeyServiceError('签名密钥标识 kid 包含不允许的字符') + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + record = await OidcKeyDao.get_by_kid_for_update(db, kid) + if record is None or record.status not in {'active', 'retiring'}: + raise KeyServiceError('签名密钥未处于活动状态') + if record.status == 'retiring': + return False + if record.status == 'active': + active_count = sum(1 for item in await OidcKeyDao.lock_algorithm_for_update(db) if item.status == 'active') + if active_count <= 1: + raise KeyServiceError('必须保留至少一把有效的活动签名密钥') + record.status = 'retiring' + record.signing_stop_at = current + record.remove_from_jwks_at = current + timedelta( + seconds=max( + OidcConfig.oidc_key_rotation_overlap_seconds, + OidcConfig.oidc_max_access_token_ttl_seconds + OidcConfig.oidc_allowed_clock_skew_seconds, + ) + ) + await AuditService.record( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'success', + detail={'action': 'retired', 'kid': kid, 'actor': actor[:64]}, + ) + + return True + + @classmethod + async def retire_due(cls, db: AsyncSession) -> int: + """ + 退役已结束 JWKS 保留期的签名密钥 + + :param db: 异步数据库会话 + :return: 更新的记录数量 + """ + + cls._require_enabled() + + return await OidcKeyDao.retire_due(db) + + @classmethod + async def activate_due( + cls, + db: AsyncSession, + *, + redis: Any = None, + now: datetime | None = None, + audit_writer: Callable[..., Awaitable[Any]] | None = None, + ) -> int: + """ + 批量激活已到签名开始时间的 pending 密钥 + + :param db: 异步数据库会话;每次状态推进由调用方提交 + :param redis: 应用共享 Redis 轮换锁客户端 + :param now: 可注入当前项目时间 + :param audit_writer: 可注入的独立审计写入器,生产默认使用独立提交事务 + :return: 本次成功激活的密钥数量 + """ + + cls._require_enabled() + current = TimezoneUtil.to_optional_utc(now or TimezoneUtil.utc_now()) + pending = await OidcKeyDao.list_due_pending(db, now=current) + activated = 0 + for record in pending: + try: + if await cls.activate_key( + db, + record.kid, + now=current, + redis=redis, + actor='system:lifecycle', + ): + await db.commit() + activated += 1 + except Exception: # noqa: PERF203 + await db.rollback() + safe_kid = ( + record.kid if isinstance(record.kid, str) and OidcUtil.is_valid_kid(record.kid) else 'invalid' + ) + try: + writer = audit_writer or AuditService.record_independent + await writer( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'failure', + failure_code='key_activation_failed', + detail={'action': 'activation_failed', 'kid': safe_kid}, + ) + except Exception: + logger.warning('OIDC 签名密钥激活审计记录写入失败,密钥标识={}', safe_kid) + logger.warning('OIDC 签名密钥激活失败,事件=signing_key_rotated,密钥标识={}', safe_kid) + return activated + + @classmethod + async def delete_key( + cls, + db: AsyncSession, + kid: str, + *, + actor: str = 'system:lifecycle', + ) -> bool: + """ + 删除已退役且超过 JWKS 保留期的签名密钥 + + :param db: 异步数据库会话 + :param kid: 待删除的 kid + :param actor: 操作人标识 + :return: 删除成功时返回 True + :raises KeyServiceError: 密钥仍可能用于验签或不存在 + """ + + cls._require_enabled(management=True) + if not OidcUtil.is_valid_kid(kid): + raise KeyServiceError('签名密钥标识 kid 包含不允许的字符') + record = await OidcKeyDao.get_by_kid_for_update(db, kid) + if record is None or record.status != 'retired': + raise KeyServiceError('签名密钥尚未完成安全退役') + remove_at = TimezoneUtil.to_optional_utc(record.remove_from_jwks_at) + if remove_at is None or remove_at > TimezoneUtil.utc_now(): + raise KeyServiceError('签名密钥仍在 JWKS 公钥保留期内') + if not await OidcKeyDao.delete_retired(db, kid): + raise KeyServiceError('签名密钥删除失败') + await AuditService.record( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'success', + detail={'action': 'deleted', 'kid': kid, 'actor': actor[:64]}, + ) + + return True + + +class OidcKeyManagementService: + """ + OIDC 签名密钥管理模块服务层 + """ + + _ERROR_MESSAGES = { + '目标签名密钥不处于待激活状态': '签名密钥状态已变化,请刷新列表后重试', + '目标签名密钥尚未发布': '签名公钥尚未到公开时间,暂时不能开始使用', + '签名密钥正在轮换,请稍后重试': '其他实例正在处理签名密钥,请稍后重试', + } + + @staticmethod + def view(row: SysOidcSigningKey) -> dict[str, object]: + """ + 将签名密钥 ORM 记录转换为管理视图 + + :param row: 签名密钥 ORM 记录 + :return: 签名密钥管理视图字段映射 + """ + + return OidcKeyViewModel( + kid=row.kid, + key_use=row.key_use, + alg=row.alg, + public_jwk=row.public_jwk, + status=row.status, + publish_at=TimezoneUtil.to_optional_utc(row.publish_at), + signing_start_at=TimezoneUtil.to_optional_utc(row.signing_start_at), + signing_stop_at=TimezoneUtil.to_optional_utc(row.signing_stop_at), + remove_from_jwks_at=TimezoneUtil.to_optional_utc(row.remove_from_jwks_at), + create_time=TimezoneUtil.to_optional_utc(row.create_time), + ).model_dump(by_alias=True) + + @classmethod + async def bootstrap( + cls, + db: AsyncSession, + *, + kid: str, + actor: str, + redis: object, + now: datetime | None = None, + ) -> tuple[dict[str, object], bool]: + """ + 幂等创建并激活首把 OIDC 签名密钥 + + 已存在可用 active 密钥时只校验并返回;否则在 Redis 轮换锁与 + 数据库行锁保护下创建、激活并提交,适合部署流水线并发调用。 + + :param db: 异步数据库会话 + :param kid: 首把签名密钥标识 + :param actor: 部署操作人标识 + :param redis: 异步 Redis 客户端 + :param now: 可注入的当前项目时间 + :return: 管理视图与本次是否创建新密钥 + :raises ServiceException: 签名密钥初始化失败 + """ + + row, created = await cls._execute( + db, + lambda: KeyService.bootstrap_signing_key( + db, + kid=kid, + actor=actor, + redis=redis, + now=now, + ), + ) + await db.refresh(row) + + return cls.view(row), created + + @classmethod + async def list_page( + cls, + db: AsyncSession, + status: Literal['pending', 'active', 'retiring', 'retired', 'compromised'] | None, + page_num: int, + page_size: int, + ) -> tuple[list[dict[str, object]], int]: + """ + 分页查询签名密钥管理视图 + + :param db: 异步数据库会话 + :param status: 密钥状态筛选值 + :param page_num: 页码 + :param page_size: 页大小 + :return: 管理视图列表和总数 + """ + + rows = await OidcKeyDao.list_admin(db, status=status, offset=(page_num - 1) * page_size, limit=page_size) + total = await OidcKeyDao.count_admin(db, status=status) + + return [cls.view(row) for row in rows], total + + @staticmethod + async def rotate(db: AsyncSession, payload: OidcKeyRotateModel, actor: str) -> dict[str, object]: + """ + 创建 pending 签名密钥并返回管理视图 + + :param db: 异步数据库会话 + :param payload: 签名密钥轮换参数 + :param actor: 操作人标识 + :return: 新建签名密钥的管理视图字段映射 + """ + + result = await OidcKeyManagementService._execute( + db, + lambda: KeyService.create_pending_key( + db, + kid=payload.kid, + publish_at=payload.publish_at, + activate_at=payload.activate_at, + actor=actor, + remark=payload.remark, + ), + ) + + return OidcKeyManagementService.view(result) + + @staticmethod + async def activate(db: AsyncSession, kid: str, actor: str, redis: object) -> bool: + """ + 激活指定签名密钥并提交事务 + + :param db: 异步数据库会话 + :param kid: 密钥标识 + :param actor: 操作人标识 + :param redis: 异步 Redis 客户端 + :return: 本次是否激活签名密钥 + :raises ServiceException: 签名密钥无法开始使用 + """ + + return await OidcKeyManagementService._execute( + db, lambda: KeyService.activate_key(db, kid, actor=actor, redis=redis) + ) + + @staticmethod + async def retire(db: AsyncSession, kid: str, actor: str) -> bool: + """ + 退役指定签名密钥并提交事务 + + :param db: 异步数据库会话 + :param kid: 密钥标识 + :param actor: 操作人标识 + :return: 本次是否将签名密钥转为退役中状态 + """ + + return await OidcKeyManagementService._execute(db, lambda: KeyService.retire_key(db, kid, actor=actor)) + + @staticmethod + async def delete(db: AsyncSession, kid: str, actor: str) -> bool: + """ + 删除指定已退役签名密钥并提交事务 + + :param db: 异步数据库会话 + :param kid: 密钥标识 + :param actor: 操作人标识 + :return: 本次是否删除签名密钥 + """ + + return await OidcKeyManagementService._execute(db, lambda: KeyService.delete_key(db, kid, actor=actor)) + + @staticmethod + async def _execute(db: AsyncSession, operation: Callable[[], Awaitable[T]]) -> T: + """ + 执行签名密钥管理事务并统一处理回滚 + + :param db: 异步数据库会话 + :param operation: 事务操作回调 + :return: operation 回调的返回值 + :raises ServiceException: 签名密钥管理事务失败时抛出 + """ + + try: + result = await operation() + await db.commit() + return result + except ServiceException: + await db.rollback() + raise + except KeyServiceError as exc: + await db.rollback() + message = OidcKeyManagementService._ERROR_MESSAGES.get(str(exc), 'OIDC 签名密钥操作失败') + raise ServiceException(message=message) from exc + except Exception as exc: + await db.rollback() + raise ServiceException(message='OIDC 签名密钥操作失败') from exc diff --git a/ruoyi-fastapi-backend/module_identity/service/logout_confirmation_service.py b/ruoyi-fastapi-backend/module_identity/service/logout_confirmation_service.py new file mode 100644 index 000000000..09fbd5895 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/logout_confirmation_service.py @@ -0,0 +1,127 @@ +import hmac +import json +import re +import secrets +from urllib.parse import urlsplit + +import jwt +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from config.env import OidcConfig +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.service.session_service import LogoutService +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + + +class LogoutConfirmationService: + """ + 认证中心退出确认服务层 + """ + + COOKIE_NAME = '__Host-oidc_logout_confirmation' + TTL_SECONDS = 300 + _TOKEN_LENGTH = 43 + _CONSUME = """ +local value = redis.call('GET', KEYS[1]) +if value then redis.call('DEL', KEYS[1]) end +return value +""" + + @classmethod + async def issue(cls, redis: Redis, parameters: dict[str, str], sso_cookie: str | None) -> tuple[str, str]: + """ + 签发短期有效且仅可使用一次的退出确认凭据 + + :param redis: Redis连接对象 + :param parameters: 已校验的退出请求参数 + :param sso_cookie: 首次退出请求携带的SSO Cookie + :return: 退出确认凭据和浏览器随机数 + :raises RuntimeError: 确认记录保存失败 + """ + + token = secrets.token_urlsafe(32) + nonce = secrets.token_urlsafe(32) + record = json.dumps( + { + 'parameters': parameters, + 'binding': OidcUtil.logout_confirmation_digest( + sso_cookie, nonce, pepper=OidcConfig.oidc_token_hash_pepper + ), + 'sso_bound': sso_cookie is not None, + }, + separators=(',', ':'), + ) + stored = await redis.set( + 'oidc:logout:confirmation:' + OidcUtil.sha256_digest(token), + record, + ex=cls.TTL_SECONDS, + nx=True, + ) + if not stored: + raise RuntimeError('退出确认服务暂不可用') + return token, nonce + + @classmethod + async def consume( + cls, redis: Redis, token: str, sso_cookie: str | None, browser_nonce: str | None + ) -> dict[str, str]: + """ + 消费一次性退出确认凭据并校验浏览器归属 + + :param redis: Redis连接对象 + :param token: 用户提交的退出确认凭据 + :param sso_cookie: 确认请求携带的SSO Cookie + :param browser_nonce: 确认请求携带的浏览器随机数 + :return: 原始退出请求参数 + :raises ValueError: 凭据无效、已使用、已过期或浏览器绑定不匹配 + """ + + if not isinstance(token, str) or len(token) != cls._TOKEN_LENGTH or not browser_nonce: + raise ValueError('退出确认信息无效') + raw = await redis.eval(cls._CONSUME, 1, 'oidc:logout:confirmation:' + OidcUtil.sha256_digest(token)) + if raw is None: + raise ValueError('退出确认已过期或已使用') + record = json.loads(raw) + # 跨站首次POST可能不携带Lax会话Cookie,随机数仍绑定当前浏览器 + # 同源确认时仅校验并撤销当前浏览器的会话 + bound_sso = sso_cookie if record.get('sso_bound', True) else None + if not hmac.compare_digest( + record['binding'], + OidcUtil.logout_confirmation_digest(bound_sso, browser_nonce, pepper=OidcConfig.oidc_token_hash_pepper), + ): + raise ValueError('退出确认的浏览器会话已变更,请重新发起退出') + return record['parameters'] + + @classmethod + async def form_redirect_origin(cls, db: AsyncSession, parameters: dict[str, str]) -> str | None: + """ + 获取退出确认页内容安全策略允许的客户端回调源 + + 校验ID Token提示,并精确匹配已启用的注册回调地址。 + + :param db: orm对象 + :param parameters: 已校验的退出请求参数 + :return: 允许的回调源地址,校验失败时返回None + """ + + hint = parameters.get('id_token_hint') + uri = parameters.get('post_logout_redirect_uri') + if not hint or not OidcUtil.is_safe_post_logout_uri(uri): + return None + parsed = urlsplit(uri) + hostname = parsed.hostname.encode('idna').decode('ascii') + if not re.fullmatch(r'[A-Za-z0-9.:-]+', hostname): + return None + try: + _claims, client = await LogoutService._validate_id_token_hint(db, hint, TimezoneUtil.utc_now()) + except (ValueError, jwt.PyJWTError): + return None + registered = await OAuthClientDao.find_exact_uri(db, client.client_pk, 'post_logout', uri) + if registered is None or registered.status != '0': + return None + host = f'[{hostname}]' if ':' in hostname else hostname + port = f':{parsed.port}' if parsed.port is not None else '' + + return f'{parsed.scheme}://{host}{port}' diff --git a/ruoyi-fastapi-backend/module_identity/service/oauth_management_service.py b/ruoyi-fastapi-backend/module_identity/service/oauth_management_service.py new file mode 100644 index 000000000..f54aaa502 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/oauth_management_service.py @@ -0,0 +1,2068 @@ +import ipaddress +from collections.abc import Awaitable, Callable, Iterable, Sequence +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import TypeVar +from urllib.parse import parse_qsl, urlsplit +from uuid import uuid4 + +from sqlalchemy.exc import IntegrityError +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import ServiceException +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_resource_dao import OAuthResourceDao +from module_identity.entity.do.oauth_client_do import SysOAuthClient, SysOAuthClientSecret, SysOAuthClientUri +from module_identity.entity.do.oauth_resource_do import SysOAuthResource, SysOAuthScope +from module_identity.entity.vo.oauth_client_vo import ( + ClientCreateModel, + ClientPageQueryModel, + ClientSecretResponseModel, + ClientStatusModel, + ClientUpdateModel, + ClientUriModel, + ClientViewModel, +) +from module_identity.entity.vo.oauth_resource_vo import ( + ResourceCreateModel, + ResourcePageQueryModel, + ResourceStatusModel, + ResourceUpdateModel, + ResourceViewModel, + ScopeModel, + ScopePageQueryModel, + ScopeStatusModel, +) +from module_identity.security.client_auth import hash_client_secret +from module_identity.security.uri_validator import is_safe_backchannel_uri +from module_identity.service.audit_service import AuditService +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + +T = TypeVar('T') + + +@dataclass(frozen=True, slots=True) +class ResourceInvalidationTargets: + """ + 保存 Resource 变更后需要撤销的 Grant 和 Refresh Token 标识 + """ + + resource_id: str + client_ids: tuple[str, ...] + grant_ids: tuple[str, ...] + refresh_token_ids: tuple[str, ...] + + +class OAuthClientManagementError(ServiceException): + """ + 表示 OAuth Client 管理错误 + """ + + def __init__(self, message: str) -> None: + """ + 初始化 OAuth Client 管理异常 + + :param message: 错误消息 + :return: None + """ + + super().__init__(message=message) + self.args = (message,) + + +class OAuthManagementBaseService: + """ + OAuth 管理模块公共服务层 + """ + + @classmethod + async def _transaction( + cls, + db: AsyncSession, + operation: Callable[[], Awaitable[T]], + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> T: + """ + 执行事务操作 + + :param db: 异步数据库会话 + :param operation: 事务内的管理变更操作 + :param after_commit: 事务提交成功后执行的异步回调 + :return: operation 回调的返回值 + :raises OAuthClientManagementError: 敏感凭据相关 ServiceException 被转换为安全管理异常 + """ + + try: + result = await operation() + await db.commit() + except Exception as error: + await db.rollback() + if isinstance(error, IntegrityError): + raise OAuthClientManagementError('OAuth 资源或绑定关系已存在') from error + if isinstance(error, ServiceException) and any( + field in str(error).lower() for field in ('secret_hash', 'secret_key', 'token_hash', 'client_secret') + ): + raise OAuthClientManagementError('OAuth 管理请求被拒绝') from error + raise + if after_commit is not None: + await after_commit() + return result + + _OAUTH_RESPONSE_QUERY_KEYS = frozenset({'code', 'error', 'error_description', 'error_uri', 'state', 'iss'}) + _MAX_AUDIENCE_LENGTH = 500 + _MAX_CLAIMS = 64 + _MAX_CLAIM_NAME_LENGTH = 64 + _MAX_CLIENT_ID_LENGTH = 64 + _URI_TYPES = ( + ('redirect', 'redirect_uris'), + ('post_logout', 'post_logout_redirect_uris'), + ('backchannel_logout', 'backchannel_logout_uris'), + ('cors_origin', 'cors_origins'), + ) + _ALLOWED_CLAIMS = frozenset( + { + 'sub', + 'name', + 'preferred_username', + 'picture', + 'updated_at', + 'email', + 'email_verified', + 'phone_number', + 'phone_number_verified', + 'dept_id', + 'dept_name', + 'roles', + 'client_id', + 'scope', + 'sid', + } + ) + + @staticmethod + async def _record_audit( + db: AsyncSession, + event_type: str, + actor: str, + *, + client_id: str | None = None, + resource_id: str | None = None, + sid: str | None = None, + grant_id: str | None = None, + detail: dict[str, str] | None = None, + ) -> None: + """ + 记录管理审计 + + :param db: 异步数据库会话 + :param event_type: 审计事件类型 + :param actor: 操作人标识 + :param client_id: 客户端标识 + :param resource_id: 资源标识 + :param sid: Session 标识 + :param grant_id: Grant 标识 + :param detail: 审计详情 + :return: None + """ + + safe_detail = {'actor': actor[:64]} + if detail: + safe_detail.update({key: value[:200] for key, value in detail.items() if key in {'reason', 'action'}}) + await AuditService.record( + db, + event_type, + 'success', + client_id=client_id, + resource_id=resource_id, + sid=sid, + grant_id=grant_id, + detail=safe_detail, + ) + + @staticmethod + def _now(value: datetime | None) -> datetime: + """ + 将输入时间规范化为项目时间 + + :param value: 调用方提供的当前时间值 + :return: 规范化后的带时区的 UTC 时间 + :raises OAuthClientManagementError: now 不是 datetime 时抛出 + """ + + if value is None: + return TimezoneUtil.utc_now() + if not isinstance(value, datetime): + raise OAuthClientManagementError('当前时间必须为 datetime 对象') + return TimezoneUtil.to_utc(value) + + @staticmethod + def _actor(actor: str) -> str: + """ + 校验操作人标识 + + :param actor: 操作人标识 + :return: 规范化后的操作者标识 + :raises OAuthClientManagementError: actor 为空或不是字符串时抛出 + """ + + try: + return OidcUtil.actor_name(actor) + except ValueError as exc: + raise OAuthClientManagementError(str(exc)) from exc + + @classmethod + def _validate_client_ttls(cls, payload: ClientCreateModel) -> None: + """ + 校验 Client 令牌时效策略 + + :param payload: Client 创建参数 + :return: None + :raises OAuthClientManagementError: Client TTL 非正数、超出平台上限或闲置 TTL 超过绝对 TTL 时抛出 + """ + + access = payload.access_token_ttl_seconds or OidcConfig.oidc_access_token_ttl_seconds + refresh_idle = payload.refresh_token_idle_seconds or OidcConfig.oidc_refresh_token_idle_seconds + refresh_absolute = payload.refresh_token_absolute_seconds or OidcConfig.oidc_refresh_token_absolute_seconds + values = (access, refresh_idle, refresh_absolute, OidcConfig.oidc_max_access_token_ttl_seconds) + if any(isinstance(value, bool) or not isinstance(value, int) or value <= 0 for value in values): + raise OAuthClientManagementError('客户端令牌有效期配置无效') + if access > OidcConfig.oidc_max_access_token_ttl_seconds: + raise OAuthClientManagementError('客户端访问令牌有效期超过平台上限') + if refresh_idle > OidcConfig.oidc_refresh_token_idle_seconds: + raise OAuthClientManagementError('客户端刷新令牌闲置有效期超过平台上限') + if refresh_absolute > OidcConfig.oidc_refresh_token_absolute_seconds: + raise OAuthClientManagementError('客户端刷新令牌绝对有效期超过平台上限') + if refresh_idle > refresh_absolute: + raise OAuthClientManagementError('客户端刷新令牌闲置有效期不能超过绝对有效期') + + @classmethod + def _validate_resource_ttls(cls, payload: ResourceCreateModel) -> None: + """ + 校验 Resource Access Token 时效策略 + + :param payload: Resource 创建或更新参数 + :return: None + :raises OAuthClientManagementError: Resource TTL 非正数或超出平台上限时抛出 + """ + + value = payload.access_token_ttl_seconds or OidcConfig.oidc_access_token_ttl_seconds + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise OAuthClientManagementError('资源访问令牌有效期配置无效') + if value > OidcConfig.oidc_max_access_token_ttl_seconds: + raise OAuthClientManagementError('资源访问令牌有效期超过平台上限') + + @classmethod + def _validate_uri(cls, uri_type: str, uri: str) -> str: + """ + 校验并规范化 Client 注册 URI + + :param uri_type: URI 类型 + :param uri: 回调 URI + :return: 校验后的完整 URI + :raises OAuthClientManagementError: URI 格式不合法、包含保留参数或 Back-Channel 地址不安全时抛出 + """ + + try: + value = ClientUriModel(uri_type=uri_type, uri=uri).uri + except (TypeError, ValueError) as exc: + raise OAuthClientManagementError(str(exc)) from exc + if uri_type == 'backchannel_logout': + parsed = urlsplit(value) + if parsed.scheme != 'https': + raise OAuthClientManagementError('后端退出通知地址必须使用 HTTPS') + try: + address = ipaddress.ip_address(parsed.hostname or '') + except ValueError: + address = None + if address is not None and not OidcUtil.is_public_ip(parsed.hostname or ''): + raise OAuthClientManagementError('后端退出通知地址不得使用回环、私网或保留 IP 地址') + if uri_type == 'redirect': + query_keys = {key for key, _ in parse_qsl(urlsplit(value).query, keep_blank_values=True)} + forbidden = query_keys & cls._OAUTH_RESPONSE_QUERY_KEYS + if forbidden: + raise OAuthClientManagementError('回调地址包含 OAuth 响应保留参数') + return value + + @classmethod + def _uri_values(cls, payload: ClientCreateModel) -> dict[str, list[str]]: + """ + 收集并校验 Client 注册 URI + + :param payload: Client 创建参数 + :return: 按 URI 类型分组的注册地址映射 + :raises OAuthClientManagementError: URI 重复或格式不合法时抛出 + """ + + values: dict[str, list[str]] = {} + for uri_type, field_name in cls._URI_TYPES: + uris = list(getattr(payload, field_name)) + if len(set(uris)) != len(uris): + raise OAuthClientManagementError(f'{field_name} 不得包含重复项') + values[uri_type] = [cls._validate_uri(uri_type, uri) for uri in uris] + return values + + @classmethod + async def _uri_values_async(cls, payload: ClientCreateModel) -> dict[str, list[str]]: + """ + 收集 URI 并执行 Back-Channel DNS 校验 + + :param payload: Client 创建参数 + :return: 完成 DNS 校验的注册地址映射 + :raises OAuthClientManagementError: Back-Channel URI DNS 目标不是公网地址时抛出 + """ + + values = cls._uri_values(payload) + backchannels = values.get('backchannel_logout', []) + for uri in backchannels: + if not await is_safe_backchannel_uri(uri): + raise OAuthClientManagementError('后端退出通知地址的 DNS 解析结果不是公网地址') + return values + + @classmethod + async def _load_bindings( + cls, db: AsyncSession, payload: ClientCreateModel + ) -> tuple[list[SysOAuthScope], list[SysOAuthResource]]: + """ + 查询并校验 Client 的 Scope 与 Resource 绑定 + + :param db: 异步数据库会话 + :param payload: Client 创建参数 + :return: 已校验的 Scope ORM 列表和 Resource ORM 列表 + :raises OAuthClientManagementError: Scope 或 Resource 不存在、未启用或绑定关系不合法时抛出 + """ + + try: + scope_codes = OidcUtil.unique_codes(payload.scope_codes, 'scope_codes') + resource_ids = OidcUtil.unique_codes(payload.resource_ids, 'resource_ids') + pre_authorized = set( + OidcUtil.unique_codes(payload.pre_authorized_scope_codes, 'pre_authorized_scope_codes') + ) + except ValueError as exc: + raise OAuthClientManagementError(str(exc)) from exc + if not pre_authorized.issubset(scope_codes): + raise OAuthClientManagementError('预授权权限必须包含在客户端允许申请的权限范围内') + + scopes: list[SysOAuthScope] = [] + if scope_codes: + scopes = list(await OAuthClientDao.get_active_scopes_by_codes(db, scope_codes)) + if {scope.scope_code for scope in scopes} != set(scope_codes): + raise OAuthClientManagementError('权限列表包含未启用或不存在的权限范围') + + resources: list[SysOAuthResource] = [] + if resource_ids: + resources = list(await OAuthClientDao.get_active_resources_by_ids(db, resource_ids)) + if {resource.resource_id for resource in resources} != set(resource_ids): + raise OAuthClientManagementError('资源列表包含未启用或不存在的资源') + + resource_pks = {resource.resource_pk for resource in resources} + if any(scope.scope_type == 'resource' and scope.resource_pk not in resource_pks for scope in scopes): + raise OAuthClientManagementError('资源权限必须绑定客户端已获准访问的资源') + return scopes, resources + + @classmethod + async def _replace_bindings( + cls, + db: AsyncSession, + client_pk: int, + payload: ClientCreateModel, + scopes: Sequence[SysOAuthScope], + resources: Sequence[SysOAuthResource], + now: datetime, + ) -> None: + """ + 写入 Client 的 Scope、Resource 与 URI 绑定 + + :param db: 异步数据库会话 + :param client_pk: 客户端主键 + :param payload: Client 创建参数 + :param scopes: Scope 集合 + :param resources: Resource 集合 + :param now: 当前时间 + :return: None + """ + + uri_values = await cls._uri_values_async(payload) + await OAuthClientDao.replace_bindings( + db, + client_pk, + scopes, + resources, + uri_values, + set(payload.pre_authorized_scope_codes), + now, + allowed_role_keys=payload.allowed_role_keys, + ) + + @staticmethod + def _client_values(payload: ClientCreateModel) -> dict[str, object]: + """ + 转换 Client DTO 为 ORM 字段 + + :param payload: Client 创建参数 + :return: 可写入 Client ORM 的字段映射 + """ + + return { + 'client_name': payload.client_name, + 'client_type': payload.client_type, + 'token_endpoint_auth_method': payload.token_endpoint_auth_method, + 'grant_types': list(payload.grant_types), + 'response_types': list(payload.response_types), + 'require_pkce': int(payload.require_pkce), + 'require_consent': int(payload.require_consent), + 'trusted_client': int(payload.trusted_client), + 'access_token_ttl_seconds': payload.access_token_ttl_seconds, + 'refresh_token_idle_seconds': payload.refresh_token_idle_seconds, + 'refresh_token_absolute_seconds': payload.refresh_token_absolute_seconds, + 'logo_uri': payload.logo_uri, + 'policy_uri': payload.policy_uri, + 'tos_uri': payload.tos_uri, + 'remark': payload.remark, + } + + @staticmethod + def _secret_response(client_id: str, secret: SysOAuthClientSecret, plaintext: str) -> ClientSecretResponseModel: + """ + 构造 Client Secret 返回模型 + + :param client_id: 客户端标识 + :param secret: 已保存的 Client Secret ORM 记录 + :param plaintext: 明文密钥 + :return: 包含一次性明文的 ClientSecretResponseModel + """ + + return ClientSecretResponseModel( + client_id=client_id, + secret_id=secret.secret_id, + client_secret=plaintext, + secret_hint=secret.secret_hint, + not_before=secret.not_before, + expires_at=secret.expires_at, + ) + + @classmethod + async def _new_secret( + cls, + db: AsyncSession, + client: SysOAuthClient, + actor: str, + created_at: datetime, + *, + not_before: datetime, + expires_at: datetime | None = None, + ) -> ClientSecretResponseModel: + """ + 创建并 flush Client Secret ORM 记录 + + :param db: 异步数据库会话 + :param client: Client ORM 记录 + :param actor: 操作人标识 + :param created_at: 创建时间 + :param not_before: 生效时间 + :param expires_at: 过期时间 + :return: 已 flush 的 ClientSecretResponseModel + """ + + plaintext = OidcUtil.generate_client_secret() + secret = SysOAuthClientSecret( + secret_id=str(uuid4()), + client_pk=client.client_pk, + secret_hash=hash_client_secret(plaintext), + secret_hint=f'...{plaintext[-6:]}', + status='active', + not_before=not_before, + expires_at=expires_at, + create_by=actor, + create_time=created_at, + ) + await OAuthClientDao.add_secret(db, secret) + + return cls._secret_response(client.client_id, secret, plaintext) + + @classmethod + def _validate_claims(cls, values: Iterable[str], field_name: str = 'claims') -> list[str]: + """ + 校验 Resource Claim 白名单 + + :param values: Resource Claim 名称列表 + :param field_name: 字段名称 + :return: 已校验的 Resource Claim 名称列表 + :raises OAuthClientManagementError: Claim 列表不是允许的字符串列表或包含重复项时抛出 + """ + + if not isinstance(values, list) or len(values) > cls._MAX_CLAIMS: + raise OAuthClientManagementError(f'{field_name} 必须为列表,且最多包含 64 项') + result: list[str] = [] + for value in values: + if not isinstance(value, str) or not value or len(value) > cls._MAX_CLAIM_NAME_LENGTH: + raise OAuthClientManagementError(f'{field_name} 包含无效的声明') + if value not in cls._ALLOWED_CLAIMS: + raise OAuthClientManagementError(f'{field_name} 包含不允许发布的声明') + if value in result: + raise OAuthClientManagementError(f'{field_name} 不得包含重复项') + result.append(value) + return result + + @staticmethod + def _resource_values(payload: ResourceCreateModel) -> dict[str, object]: + """ + 转换 Resource DTO 为 ORM 字段 + + :param payload: Resource 创建或更新参数 + :return: 可写入 Resource ORM 的字段映射 + """ + + return { + 'resource_name': payload.resource_name, + 'audience': payload.audience, + 'token_format': payload.token_format, + 'signing_alg': payload.signing_alg, + 'access_token_ttl_seconds': payload.access_token_ttl_seconds, + 'allowed_claims': list(payload.allowed_claims), + 'remark': payload.remark, + } + + @classmethod + async def _resolve_introspection_client(cls, db: AsyncSession, client_id: str | None) -> SysOAuthClient | None: + """ + 解析 Resource 使用的 introspection Client + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :return: 启用的 introspection Client ORM 记录或 None + :raises OAuthClientManagementError: introspection Client 不存在、未启用或不是机密 Client 时抛出 + """ + + if client_id is None: + return None + if not isinstance(client_id, str) or not client_id or len(client_id) > cls._MAX_CLIENT_ID_LENGTH: + raise OAuthClientManagementError('内省客户端标识 introspection_client_id 无效') + client = await OAuthClientDao.find_active_introspection_client(db, client_id) + if client is None: + raise OAuthClientManagementError('内省客户端必须是已启用的机密客户端') + return client + + @staticmethod + def _resource_view(resource: SysOAuthResource, introspection_client_id: str | None = None) -> ResourceViewModel: + """ + 转换 Resource ORM 记录为 ResourceViewModel + + :param resource: Resource ORM 记录 + :param introspection_client_id: 用于内省的客户端标识 + :return: ResourceViewModel 管理视图 + """ + + return ResourceViewModel.model_validate( + { + 'resource_id': resource.resource_id, + 'resource_name': resource.resource_name, + 'audience': resource.audience, + 'token_format': resource.token_format, + 'signing_alg': resource.signing_alg, + 'access_token_ttl_seconds': resource.access_token_ttl_seconds, + 'introspection_client_id': introspection_client_id, + 'allowed_claims': list(resource.allowed_claims or []), + 'status': resource.status, + 'remark': resource.remark, + } + ) + + @staticmethod + def _scope_view(scope: SysOAuthScope, resource_id: str | None) -> ScopeModel: + """ + 转换 Scope ORM 记录为 ScopeModel + + :param scope: Scope ORM 记录 + :param resource_id: 资源标识 + :return: ScopeModel 管理视图 + """ + + return ScopeModel.model_validate( + { + 'scope_code': scope.scope_code, + 'scope_name': scope.scope_name, + 'scope_type': scope.scope_type, + 'resource_id': resource_id, + 'claims': list(scope.claims or []), + 'consent_required': bool(scope.consent_required), + 'sensitive': bool(scope.sensitive), + 'status': scope.status, + 'remark': scope.remark, + } + ) + + @classmethod + async def _lock_resource_clients(cls, db: AsyncSession, resource_pk: int) -> list[SysOAuthClient]: + """ + 锁定关联 Resource 的 Client 记录 + + :param db: 异步数据库会话 + :param resource_pk: 资源主键 + :return: 已锁定的 Client ORM 记录列表 + """ + + return list(await OAuthClientDao.lock_clients_for_resource(db, resource_pk)) + + @staticmethod + def _bump_clients(clients: Iterable[SysOAuthClient], actor: str, now: datetime) -> None: + """ + 递增 Client 策略版本并记录操作者 + + :param clients: 需要递增策略版本的客户端集合 + :param actor: 操作人标识 + :param now: 当前时间 + :return: None + """ + + for client in clients: + client.policy_version = int(client.policy_version or 0) + 1 + client.update_by, client.update_time = actor, now + + @staticmethod + async def _revoke_client_credentials(db: AsyncSession, client_pk: int, now: datetime) -> None: + """ + 撤销 Client 的 Grant 与 Refresh Token + + :param db: 异步数据库会话 + :param client_pk: 客户端主键 + :param now: 当前时间 + :return: None + """ + + await OAuthClientDao.revoke_client_credentials(db, client_pk, now) + + +class OAuthClientManagementService(OAuthManagementBaseService): + """ + OAuth Client 管理模块服务层 + """ + + @classmethod + async def create_client( + cls, + db: AsyncSession, + payload: ClientCreateModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ClientViewModel: + """ + 创建 OAuth Client 及其绑定配置 + + :param db: 异步数据库会话 + :param payload: Client 创建参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 创建后的 Client 详情 + """ + + return await cls._transaction(db, lambda: cls._create_client(db, payload, actor=actor, now=now), after_commit) + + @classmethod + async def update_client( + cls, + db: AsyncSession, + payload: ClientUpdateModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ClientViewModel: + """ + 更新 OAuth Client 及其绑定配置 + + :param db: 异步数据库会话 + :param payload: Client 更新参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 更新后的 Client 详情 + """ + + return await cls._transaction(db, lambda: cls._update_client(db, payload, actor=actor, now=now), after_commit) + + @classmethod + async def disable_clients( + cls, + db: AsyncSession, + identifiers: list[str], + actor: str, + *, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> None: + """ + 禁用指定 OAuth Client 并撤销其凭据 + + :param db: 异步数据库会话 + :param identifiers: Client 公开标识列表 + :param actor: 操作者标识 + :param after_commit: 事务提交成功后执行的异步回调 + :return: None + """ + + async def operation() -> None: + """ + 在同一事务中批量禁用指定 Client + + :return: None + """ + + for identifier in identifiers: + await cls._soft_disable(db, identifier, actor=actor) + + await cls._transaction(db, operation, after_commit) + + @classmethod + async def change_client_status( + cls, + db: AsyncSession, + payload: ClientStatusModel, + actor: str, + *, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ClientViewModel: + """ + 变更 OAuth Client 启用状态 + + :param db: 异步数据库会话 + :param payload: Client 状态变更参数 + :param actor: 操作者标识 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 更新后的 Client 详情 + """ + + return await cls._transaction(db, lambda: cls._change_status(db, payload, actor=actor), after_commit) + + @classmethod + async def rotate_secret( + cls, + db: AsyncSession, + client_id: str, + actor: str, + *, + not_before: datetime | None = None, + expires_at: datetime | None = None, + retirement_seconds: int | None = None, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ClientSecretResponseModel: + """ + 轮换 OAuth Client Secret 并安排旧凭据退役 + + :param db: 异步数据库会话 + :param client_id: Client 公开标识 + :param actor: 操作者标识 + :param not_before: 生效时间 + :param expires_at: 过期时间 + :param retirement_seconds: 退役等待秒数 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 只包含一次性明文的 Secret 响应 + """ + + return await cls._transaction( + db, + lambda: cls._rotate_secret( + db, + client_id, + actor=actor, + not_before=not_before, + expires_at=expires_at, + retirement_seconds=retirement_seconds, + now=now, + ), + after_commit, + ) + + @classmethod + async def revoke_secret( + cls, + db: AsyncSession, + client_id: str, + secret_id: str, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> bool: + """ + 撤销指定 Client Secret + + :param db: 异步数据库会话 + :param client_id: Client 公开标识 + :param secret_id: Secret 标识 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 本次是否发生状态变更 + """ + + return await cls._transaction( + db, lambda: cls._revoke_secret(db, client_id, secret_id, actor=actor, now=now), after_commit + ) + + @classmethod + async def add_uri( + cls, + db: AsyncSession, + client_id: str, + payload: ClientUriModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> int: + """ + 添加注册 URI + + :param db: 异步数据库会话 + :param client_id: Client 公开标识 + :param payload: URI 注册参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 新增 URI 的内部标识 + """ + + return await cls._transaction( + db, lambda: cls._add_uri(db, client_id, payload, actor=actor, now=now), after_commit + ) + + @classmethod + async def remove_uri( + cls, + db: AsyncSession, + client_id: str, + uri_id: int, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> bool: + """ + 删除注册 URI + + :param db: 异步数据库会话 + :param client_id: Client 公开标识 + :param uri_id: URI 内部标识 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 本次是否发生状态变更 + """ + + return await cls._transaction( + db, lambda: cls._remove_uri(db, client_id, uri_id, actor=actor, now=now), after_commit + ) + + @classmethod + async def _create_client( + cls, + db: AsyncSession, + payload: ClientCreateModel, + *, + actor: str, + now: datetime | None = None, + ) -> ClientViewModel: + """ + 在事务中创建 Client ORM 记录及关联配置 + + :param db: 异步数据库会话 + :param payload: Client 创建参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ClientViewModel 详情 + :raises OAuthClientManagementError: Client DTO 校验失败或公开标识已存在时抛出 + """ + + if not isinstance(payload, ClientCreateModel) or isinstance(payload, ClientUpdateModel): + raise OAuthClientManagementError('客户端创建参数必须为 ClientCreateModel') + actor_value, current = cls._actor(actor), cls._now(now) + cls._validate_client_ttls(payload) + scopes, resources = await cls._load_bindings(db, payload) + client = SysOAuthClient( + client_id=OidcUtil.generate_client_id(), + subject_type='public', + policy_version=1, + status='0', + create_by=actor_value, + create_time=current, + update_by=actor_value, + update_time=current, + **cls._client_values(payload), + ) + await OAuthClientDao.add_client(db, client) + await cls._replace_bindings(db, client.client_pk, payload, scopes, resources, current) + await cls._record_audit(db, OidcAuditEvent.CLIENT_CREATED, actor_value, client_id=client.client_id) + + return await cls.detail(db, client.client_id) + + @classmethod + async def _update_client( + cls, + db: AsyncSession, + payload: ClientUpdateModel, + *, + actor: str, + now: datetime | None = None, + ) -> ClientViewModel: + """ + 在事务中更新 Client ORM 记录及关联配置 + + :param db: 异步数据库会话 + :param payload: Client 更新参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ClientViewModel 详情 + :raises OAuthClientManagementError: Client 不存在、DTO 校验失败或策略绑定不合法时抛出 + """ + + if not isinstance(payload, ClientUpdateModel): + raise OAuthClientManagementError('客户端更新参数必须为 ClientUpdateModel') + actor_value, current = cls._actor(actor), cls._now(now) + cls._validate_client_ttls(payload) + client = await OAuthClientDao.get_by_client_id(db, payload.client_id, active_only=False, for_update=True) + if client is None: + raise OAuthClientManagementError('客户端不存在') + if client.client_type != payload.client_type: + raise OAuthClientManagementError('客户端类型不可修改') + if client.client_type == 'public' and payload.token_endpoint_auth_method != 'none': + raise OAuthClientManagementError('公开客户端不能使用客户端密钥认证') + current_view = await cls.detail(db, client.client_id) + policy_changed = cls._client_policy_changed(current_view, payload) + scopes, resources = await cls._load_bindings(db, payload) + for key, value in cls._client_values(payload).items(): + setattr(client, key, value) + client.update_by, client.update_time = actor_value, current + if policy_changed: + client.policy_version = int(client.policy_version or 0) + 1 + await cls._replace_bindings(db, client.client_pk, payload, scopes, resources, current) + await OAuthClientDao.persist_client_policy_change(db, client) + if policy_changed: + await cls._record_audit( + db, + OidcAuditEvent.SECURITY_VERSION_CHANGED, + actor_value, + client_id=client.client_id, + detail={'action': 'client_updated'}, + ) + return await cls.detail(db, client.client_id) + + @staticmethod + def _client_policy_changed(current: ClientViewModel, payload: ClientUpdateModel) -> bool: + """ + 比较 Client 授权与令牌策略是否变化 + + :param current: 当前客户端 ORM 记录 + :param payload: Client 更新参数 + :return: Client 授权或令牌策略是否发生变化 + """ + + scalar_fields = ( + 'client_type', + 'token_endpoint_auth_method', + 'require_pkce', + 'require_consent', + 'trusted_client', + 'access_token_ttl_seconds', + 'refresh_token_idle_seconds', + 'refresh_token_absolute_seconds', + ) + if any(getattr(current, field) != getattr(payload, field) for field in scalar_fields): + return True + unordered_fields = ( + 'grant_types', + 'response_types', + 'scope_codes', + 'pre_authorized_scope_codes', + 'allowed_role_keys', + ) + if any(set(getattr(current, field)) != set(getattr(payload, field)) for field in unordered_fields): + return True + ordered_fields = ( + 'resource_ids', + 'redirect_uris', + 'post_logout_redirect_uris', + 'backchannel_logout_uris', + 'cors_origins', + ) + + return any(tuple(getattr(current, field)) != tuple(getattr(payload, field)) for field in ordered_fields) + + @classmethod + async def detail(cls, db: AsyncSession, client_id: str) -> ClientViewModel: + """ + 查询 OAuth Client 详情 + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :return: ClientViewModel 详情 + :raises OAuthClientManagementError: Client 不存在时抛出 + """ + + detail_rows = await OAuthClientDao.get_client_detail_rows(db, client_id) + if detail_rows is None: + raise OAuthClientManagementError('客户端不存在') + client, scope_rows, resource_ids, uri_rows = detail_rows + bindings = await OAuthClientDao.list_scope_bindings(db, client.client_pk) + allowed_role_keys = sorted( + { + role + for binding in bindings + if isinstance(binding.claim_filter, dict) + for role in binding.claim_filter.get('allowed_role_keys', []) + if isinstance(role, str) + } + ) + uri_map = {uri_type: [] for uri_type, _ in cls._URI_TYPES} + for row in uri_rows: + if row.status != '0': + continue + uri_map.setdefault(row.uri_type, []).append(row.uri) + payload = { + 'client_id': client.client_id, + 'client_name': client.client_name, + 'client_type': client.client_type, + 'token_endpoint_auth_method': client.token_endpoint_auth_method, + 'grant_types': list(client.grant_types or []), + 'response_types': list(client.response_types or []), + 'require_pkce': bool(client.require_pkce), + 'require_consent': bool(client.require_consent), + 'trusted_client': bool(client.trusted_client), + 'scope_codes': [scope.scope_code for scope, _ in scope_rows], + 'pre_authorized_scope_codes': [scope.scope_code for scope, pre in scope_rows if bool(pre)], + 'allowed_role_keys': allowed_role_keys, + 'resource_ids': list(resource_ids), + 'redirect_uris': uri_map['redirect'], + 'post_logout_redirect_uris': uri_map['post_logout'], + 'backchannel_logout_uris': uri_map['backchannel_logout'], + 'cors_origins': uri_map['cors_origin'], + 'access_token_ttl_seconds': client.access_token_ttl_seconds, + 'refresh_token_idle_seconds': client.refresh_token_idle_seconds, + 'refresh_token_absolute_seconds': client.refresh_token_absolute_seconds, + 'logo_uri': client.logo_uri, + 'policy_uri': client.policy_uri, + 'tos_uri': client.tos_uri, + 'remark': client.remark, + 'status': client.status, + 'policy_version': client.policy_version, + 'create_time': client.create_time, + 'update_time': client.update_time, + } + + return ClientViewModel.model_validate(payload) + + @classmethod + async def _add_uri( + cls, + db: AsyncSession, + client_id: str, + payload: ClientUriModel, + *, + actor: str, + now: datetime | None = None, + ) -> int: + """ + 添加 Client 注册 URI + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :param payload: URI 注册参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: 新增 URI 的内部主键 + :raises OAuthClientManagementError: Client 不存在或 URI 已注册时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + client = await OAuthClientDao.get_by_client_id(db, client_id, active_only=False, for_update=True) + if client is None: + raise OAuthClientManagementError('客户端不存在') + uri = cls._validate_uri(payload.uri_type, payload.uri) + if payload.uri_type == 'backchannel_logout' and not await is_safe_backchannel_uri(uri): + raise OAuthClientManagementError('后端退出通知地址的 DNS 解析结果不是公网地址') + uri_hash = OidcUtil.sha256_digest(uri) + if await OAuthClientDao.find_uri(db, client.client_pk, payload.uri_type, uri_hash, uri) is not None: + raise OAuthClientManagementError('该地址已注册') + row = SysOAuthClientUri( + client_pk=client.client_pk, + uri_type=payload.uri_type, + uri=uri, + uri_hash=uri_hash, + is_default=int(payload.is_default), + status=payload.status, + create_time=current, + ) + await OAuthClientDao.add_uri(db, row) + client.policy_version = int(client.policy_version or 0) + 1 + client.update_by, client.update_time = actor_value, current + await OAuthClientDao.persist_client_policy_change(db, client) + await cls._record_audit( + db, + OidcAuditEvent.SECURITY_VERSION_CHANGED, + actor_value, + client_id=client_id, + detail={'action': 'uri_added'}, + ) + + return row.uri_id + + @classmethod + async def _remove_uri( + cls, db: AsyncSession, client_id: str, uri_id: int, *, actor: str, now: datetime | None = None + ) -> bool: + """ + 删除 Client 注册 URI + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :param uri_id: 待删除的 URI 主键 + :param actor: 操作人标识 + :param now: 当前时间 + :return: 本次是否停用注册 URI + :raises OAuthClientManagementError: Client 不存在或注册 URI 不存在 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + client = await OAuthClientDao.get_by_client_id(db, client_id, active_only=False, for_update=True) + if client is None: + raise OAuthClientManagementError('客户端不存在') + row = await OAuthClientDao.get_uri_for_update(db, uri_id) + if row is None: + raise OAuthClientManagementError('注册地址不存在') + if row.client_pk != client.client_pk: + raise OAuthClientManagementError('该地址不属于当前客户端') + if row.status == '1': + return False + row.status = '1' + client.policy_version = int(client.policy_version or 0) + 1 + client.update_by, client.update_time = actor_value, current + await OAuthClientDao.persist_client_policy_change(db, client) + await cls._record_audit( + db, + OidcAuditEvent.SECURITY_VERSION_CHANGED, + actor_value, + client_id=client_id, + detail={'action': 'uri_removed'}, + ) + + return True + + @classmethod + async def list_clients(cls, db: AsyncSession, query: ClientPageQueryModel | None = None) -> list[ClientViewModel]: + """ + 分页查询 OAuth Client + + :param db: 异步数据库会话 + :param query: Client 分页查询参数 + :return: ClientViewModel 列表 + """ + + page = query or ClientPageQueryModel() + rows = await OAuthClientDao.list_clients_page(db, page) + + return [await cls.detail(db, item.client_id) for item in rows] + + @classmethod + async def count_clients(cls, db: AsyncSession, query: ClientPageQueryModel | None = None) -> int: + """ + 统计 OAuth Client 数量 + + :param db: 异步数据库会话 + :param query: Client 分页统计参数 + :return: Client 数量 + """ + + page = query or ClientPageQueryModel() + + return await OAuthClientDao.count_clients(db, page) + + @classmethod + async def _change_status( + cls, db: AsyncSession, payload: ClientStatusModel, *, actor: str, now: datetime | None = None + ) -> ClientViewModel: + """ + 在事务中变更 Client 状态 + + :param db: 异步数据库会话 + :param payload: Client 状态变更参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ClientViewModel 详情 + :raises OAuthClientManagementError: Client 不存在或状态值不合法时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + client = await OAuthClientDao.get_by_client_id(db, payload.client_id, active_only=False, for_update=True) + if client is None: + raise OAuthClientManagementError('客户端不存在') + if client.status != payload.status: + if payload.status == '1': + await OAuthClientDao.revoke_client_credentials(db, client.client_pk, current) + client.status = payload.status + client.policy_version = int(client.policy_version or 0) + 1 + client.update_by, client.update_time = actor_value, current + await OAuthClientDao.persist_client_policy_change(db, client) + await cls._record_audit( + db, + OidcAuditEvent.CLIENT_DISABLED if payload.status == '1' else OidcAuditEvent.SECURITY_VERSION_CHANGED, + actor_value, + client_id=client.client_id, + detail={'action': 'status_changed'}, + ) + return await cls.detail(db, client.client_id) + + @classmethod + async def _soft_disable( + cls, db: AsyncSession, client_id: str, *, actor: str, now: datetime | None = None + ) -> ClientViewModel: + """ + 在事务中禁用 Client 并撤销凭据 + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ClientViewModel 详情 + """ + + return await cls._change_status(db, ClientStatusModel(client_id=client_id, status='1'), actor=actor, now=now) + + @classmethod + async def _rotate_secret( + cls, + db: AsyncSession, + client_id: str, + *, + actor: str, + now: datetime | None = None, + not_before: datetime | None = None, + expires_at: datetime | None = None, + retirement_seconds: int | None = None, + ) -> ClientSecretResponseModel: + """ + 在事务中创建新的 Client Secret + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :param actor: 操作人标识 + :param now: 当前时间 + :param not_before: 生效时间 + :param expires_at: 过期时间 + :param retirement_seconds: 退役等待秒数 + :return: ClientSecretResponseModel(含一次性明文) + :raises OAuthClientManagementError: Client 不存在或 Secret 时效不合法时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + retirement_ttl = retirement_seconds + if retirement_ttl is None: + retirement_ttl = getattr( + OidcConfig, + 'oidc_client_secret_retirement_seconds', + OidcConfig.oidc_key_rotation_overlap_seconds, + ) + if isinstance(retirement_ttl, bool) or not isinstance(retirement_ttl, int) or retirement_ttl <= 0: + raise OAuthClientManagementError('客户端旧密钥退役过渡期必须为正整数') + effective = cls._now(not_before) if not_before is not None else current + expiry = cls._now(expires_at) if expires_at is not None else None + if expiry is not None and expiry <= effective: + raise OAuthClientManagementError('过期时间必须晚于生效时间') + client = await OAuthClientDao.get_by_client_id(db, client_id, active_only=False, for_update=True) + if client is None: + raise OAuthClientManagementError('客户端不存在') + if client.client_type != 'confidential' or client.token_endpoint_auth_method != 'client_secret_basic': + raise OAuthClientManagementError('公开客户端不能生成或轮换客户端密钥') + old_secrets = await OAuthClientDao.list_secrets(db, client.client_pk, active_only=False, for_update=True) + old_secrets = [secret for secret in old_secrets if secret.status == 'active'] + for old_secret in old_secrets: + if old_secret.expires_at is not None and effective >= cls._now(old_secret.expires_at): + raise OAuthClientManagementError('密钥生效时间会造成认证凭据空档,请保留新旧密钥的重叠有效期') + retirement_end = max(current, effective) + timedelta(seconds=retirement_ttl) + for old_secret in old_secrets: + old_secret.status = 'retiring' + old_expiry = cls._now(old_secret.expires_at) if old_secret.expires_at is not None else None + if old_expiry is None or old_expiry > retirement_end: + old_secret.expires_at = retirement_end + secret = await cls._new_secret(db, client, actor_value, current, not_before=effective, expires_at=expiry) + # 凭据轮换只影响客户端认证,保留现有用户授权。 + client.update_by, client.update_time = actor_value, current + await OAuthClientDao.persist_client_policy_change(db, client) + await cls._record_audit(db, OidcAuditEvent.CLIENT_SECRET_ROTATED, actor_value, client_id=client.client_id) + + return secret + + @classmethod + async def _revoke_secret( + cls, db: AsyncSession, client_id: str, secret_id: str, *, actor: str, now: datetime | None = None + ) -> bool: + """ + 在事务中撤销 Client Secret + + :param db: 异步数据库会话 + :param client_id: 客户端标识 + :param secret_id: 待撤销的 Secret 主键 + :param actor: 操作人标识 + :param now: 当前时间 + :return: 本次是否撤销 Client Secret + :raises OAuthClientManagementError: Client Secret 不存在或已经撤销时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + client = await OAuthClientDao.get_by_client_id(db, client_id, active_only=False, for_update=True) + if client is None: + raise OAuthClientManagementError('客户端不存在') + secret = await OAuthClientDao.get_secret_for_update(db, secret_id) + if secret is None: + raise OAuthClientManagementError('客户端密钥不存在') + if secret.client_pk != client.client_pk: + raise OAuthClientManagementError('该密钥不属于当前客户端') + if secret.status == 'revoked': + return False + secret.status, secret.revoked_by, secret.revoked_at = 'revoked', actor_value, current + # 撤销单个凭据不改变用户授权策略版本。 + client.update_by, client.update_time = actor_value, current + await OAuthClientDao.persist_client_policy_change(db, client) + await cls._record_audit( + db, + OidcAuditEvent.CLIENT_SECRET_ROTATED, + actor_value, + client_id=client.client_id, + detail={'action': 'secret_revoked'}, + ) + + return True + + @staticmethod + async def list_active_cors_origins(db: AsyncSession) -> tuple[str, ...]: + """ + 查询当前启用的 CORS Origin + + :param db: 异步数据库会话 + :return: 启用的 CORS Origin 元组 + """ + + return await OAuthClientDao.list_cors_origins(db) + + +class OAuthResourceManagementService(OAuthManagementBaseService): + """ + OAuth Resource 和 Scope 管理模块服务层 + """ + + @classmethod + async def create_resource( + cls, + db: AsyncSession, + payload: ResourceCreateModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ResourceViewModel: + """ + 创建 OAuth Resource 及其 Claim 配置 + + :param db: 异步数据库会话 + :param payload: Resource 创建参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 创建后的 Resource 详情 + """ + + return await cls._transaction( + db, + lambda: cls._create_resource(db, payload, actor=actor, now=now), + after_commit, + ) + + @classmethod + async def update_resource( + cls, + db: AsyncSession, + payload: ResourceUpdateModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ResourceViewModel: + """ + 更新 OAuth Resource 及其 Claim 配置 + + :param db: 异步数据库会话 + :param payload: Resource 更新参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 更新后的 Resource 详情 + """ + + return await cls._transaction( + db, + lambda: cls._update_resource(db, payload, actor=actor, now=now), + after_commit, + ) + + @classmethod + async def disable_resources( + cls, + db: AsyncSession, + identifiers: list[str], + actor: str, + *, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> None: + """ + 禁用指定 OAuth Resource 并失效关联授权 + + :param db: 异步数据库会话 + :param identifiers: Resource 公开标识列表 + :param actor: 操作者标识 + :param after_commit: 事务提交成功后执行的异步回调 + :return: None + """ + + async def operation() -> None: + """ + 在同一事务中批量禁用指定 Resource + + :return: None + """ + + for identifier in identifiers: + await cls._soft_disable_resource(db, identifier, actor=actor) + + await cls._transaction(db, operation, after_commit) + + @classmethod + async def change_resource_status( + cls, + db: AsyncSession, + payload: ResourceStatusModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ResourceViewModel: + """ + 变更 OAuth Resource 启用状态 + + :param db: 异步数据库会话 + :param payload: Resource 状态变更参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 更新后的 Resource 详情 + """ + + return await cls._transaction( + db, lambda: cls._change_resource_status(db, payload, actor=actor, now=now), after_commit + ) + + @classmethod + async def create_scope( + cls, + db: AsyncSession, + payload: ScopeModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ScopeModel: + """ + 创建 OAuth Scope 及其 Resource 关联 + + :param db: 异步数据库会话 + :param payload: Scope 创建参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 创建后的 Scope 详情 + """ + + return await cls._transaction(db, lambda: cls._create_scope(db, payload, actor=actor, now=now), after_commit) + + @classmethod + async def update_scope( + cls, + db: AsyncSession, + payload: ScopeModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ScopeModel: + """ + 更新 OAuth Scope 及其 Resource 关联 + + :param db: 异步数据库会话 + :param payload: Scope 更新参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 更新后的 Scope 详情 + """ + + return await cls._transaction(db, lambda: cls._update_scope(db, payload, actor=actor, now=now), after_commit) + + @classmethod + async def disable_scopes( + cls, + db: AsyncSession, + identifiers: list[str], + actor: str, + *, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> None: + """ + 禁用指定 OAuth Scope 并失效关联授权 + + :param db: 异步数据库会话 + :param identifiers: Scope 编码列表 + :param actor: 操作者标识 + :param after_commit: 事务提交成功后执行的异步回调 + :return: None + """ + + async def operation() -> None: + """ + 在同一事务中批量禁用指定 Scope + + :return: None + """ + + for identifier in identifiers: + await cls._soft_disable_scope(db, identifier, actor=actor) + + await cls._transaction(db, operation, after_commit) + + @classmethod + async def change_scope_status( + cls, + db: AsyncSession, + payload: ScopeStatusModel, + actor: str, + *, + now: datetime | None = None, + after_commit: Callable[[], Awaitable[None]] | None = None, + ) -> ScopeModel: + """ + 变更 OAuth Scope 启用状态 + + :param db: 异步数据库会话 + :param payload: Scope 状态变更参数 + :param actor: 操作者标识 + :param now: 当前时间 + :param after_commit: 事务提交成功后执行的异步回调 + :return: 更新后的 Scope 详情 + """ + + return await cls._transaction( + db, lambda: cls._change_scope_status(db, payload, actor=actor, now=now), after_commit + ) + + @classmethod + async def _create_resource( + cls, + db: AsyncSession, + payload: ResourceCreateModel, + *, + actor: str, + now: datetime | None = None, + ) -> ResourceViewModel: + """ + 在事务中创建 Resource ORM 记录 + + :param db: 异步数据库会话 + :param payload: Resource 创建参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ResourceViewModel 详情 + :raises OAuthClientManagementError: Resource DTO 校验失败或 resource_id/audience 已存在时抛出 + """ + + if not isinstance(payload, ResourceCreateModel) or isinstance(payload, ResourceUpdateModel): + raise OAuthClientManagementError('资源创建参数必须为 ResourceCreateModel') + actor_value, current = cls._actor(actor), cls._now(now) + cls._validate_resource_ttls(payload) + try: + OidcUtil.validate_resource_audience(payload.audience, max_length=cls._MAX_AUDIENCE_LENGTH) + except ValueError as exc: + raise OAuthClientManagementError(str(exc)) from exc + allowed_claims = cls._validate_claims(payload.allowed_claims, 'allowed_claims') + if await OAuthResourceDao.find_resource_duplicate(db, payload.resource_id, payload.audience): + raise OAuthClientManagementError('资源标识或资源受众已注册') + client = await cls._resolve_introspection_client(db, payload.introspection_client_id) + resource = SysOAuthResource( + resource_id=payload.resource_id, + create_by=actor_value, + create_time=current, + update_by=actor_value, + update_time=current, + introspection_client_pk=client.client_pk if client else None, + status='0', + **{**cls._resource_values(payload), 'allowed_claims': allowed_claims}, + ) + await OAuthResourceDao.add_resource(db, resource) + await cls._record_audit( + db, + OidcAuditEvent.RESOURCE_POLICY_CHANGED, + actor_value, + resource_id=resource.resource_id, + detail={'action': 'resource_created'}, + ) + + return cls._resource_view(resource, payload.introspection_client_id) + + @classmethod + async def _update_resource( + cls, + db: AsyncSession, + payload: ResourceUpdateModel, + *, + actor: str, + now: datetime | None = None, + ) -> ResourceViewModel: + """ + 在事务中更新 Resource ORM 记录 + + :param db: 异步数据库会话 + :param payload: Resource 更新参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ResourceViewModel 详情 + :raises OAuthClientManagementError: Resource 不存在、DTO 校验失败或关联 Client 不存在时抛出 + """ + + if not isinstance(payload, ResourceUpdateModel): + raise OAuthClientManagementError('资源更新参数必须为 ResourceUpdateModel') + actor_value, current = cls._actor(actor), cls._now(now) + cls._validate_resource_ttls(payload) + try: + OidcUtil.validate_resource_audience(payload.audience, max_length=cls._MAX_AUDIENCE_LENGTH) + except ValueError as exc: + raise OAuthClientManagementError(str(exc)) from exc + allowed_claims = cls._validate_claims(payload.allowed_claims, 'allowed_claims') + resource = await OAuthResourceDao.get_resource(db, payload.resource_id, for_update=True) + if resource is None: + raise OAuthClientManagementError('资源不存在') + if resource.audience != payload.audience: + raise OAuthClientManagementError('资源受众不可修改,请创建新资源') + client = await cls._resolve_introspection_client(db, payload.introspection_client_id) + policy_changed = ( + resource.token_format != payload.token_format + or resource.signing_alg != payload.signing_alg + or resource.access_token_ttl_seconds != payload.access_token_ttl_seconds + or set(resource.allowed_claims or ()) != set(allowed_claims) + or resource.introspection_client_pk != (client.client_pk if client else None) + or resource.status != payload.status + ) + clients = await cls._lock_resource_clients(db, resource.resource_pk) if policy_changed else [] + resource.resource_name = payload.resource_name + resource.token_format = payload.token_format + resource.signing_alg = payload.signing_alg + resource.access_token_ttl_seconds = payload.access_token_ttl_seconds + resource.allowed_claims = allowed_claims + resource.introspection_client_pk = client.client_pk if client else None + resource.status = payload.status + resource.remark = payload.remark + resource.update_by, resource.update_time = actor_value, current + if policy_changed: + cls._bump_clients(clients, actor_value, current) + await OAuthResourceDao.persist_resource_change(db, resource) + if policy_changed: + await cls._record_audit( + db, + OidcAuditEvent.RESOURCE_POLICY_CHANGED, + actor_value, + resource_id=resource.resource_id, + detail={'action': 'resource_updated'}, + ) + return cls._resource_view(resource, payload.introspection_client_id) + + @classmethod + async def detail_resource(cls, db: AsyncSession, resource_id: str) -> ResourceViewModel: + """ + 查询 OAuth Resource 详情 + + :param db: 异步数据库会话 + :param resource_id: 资源标识 + :return: ResourceViewModel 详情 + :raises OAuthClientManagementError: Resource 不存在时抛出 + """ + + resource = await OAuthResourceDao.get_resource(db, resource_id) + if resource is None: + raise OAuthClientManagementError('资源不存在') + view = cls._resource_view(resource) + if resource.introspection_client_pk is None: + return view + client_id = await OAuthClientDao.id_for_pk(db, resource.introspection_client_pk) + + return view.model_copy(update={'introspection_client_id': client_id}) + + @classmethod + async def list_resources( + cls, db: AsyncSession, query: ResourcePageQueryModel | None = None + ) -> list[ResourceViewModel]: + """ + 分页查询 OAuth Resource + + :param db: 异步数据库会话 + :param query: Resource 分页查询参数 + :return: ResourceViewModel 列表 + """ + + page = query or ResourcePageQueryModel() + rows = await OAuthResourceDao.list_resources_page(db, page) + + return [await cls.detail_resource(db, row.resource_id) for row in rows] + + @classmethod + async def count_resources(cls, db: AsyncSession, query: ResourcePageQueryModel | None = None) -> int: + """ + 统计 OAuth Resource 数量 + + :param db: 异步数据库会话 + :param query: Resource 分页统计参数 + :return: Resource 数量 + """ + + page = query or ResourcePageQueryModel() + + return await OAuthResourceDao.count_resources(db, page) + + @classmethod + async def _change_resource_status( + cls, db: AsyncSession, payload: ResourceStatusModel, *, actor: str, now: datetime | None = None + ) -> ResourceViewModel: + """ + 在事务中变更 Resource 状态 + + :param db: 异步数据库会话 + :param payload: Resource 状态变更参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ResourceViewModel 详情 + :raises OAuthClientManagementError: Resource 不存在或状态值不合法时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + resource = await OAuthResourceDao.get_resource(db, payload.resource_id, for_update=True) + if resource is None: + raise OAuthClientManagementError('资源不存在') + if resource.status != payload.status: + clients = await cls._lock_resource_clients(db, resource.resource_pk) + resource.status = payload.status + resource.update_by, resource.update_time = actor_value, current + cls._bump_clients(clients, actor_value, current) + await OAuthResourceDao.persist_resource_change(db, resource) + await cls._record_audit( + db, + OidcAuditEvent.RESOURCE_POLICY_CHANGED, + actor_value, + resource_id=resource.resource_id, + detail={'action': 'resource_status_changed'}, + ) + return await cls.detail_resource(db, resource.resource_id) + + @classmethod + async def _soft_disable_resource( + cls, db: AsyncSession, resource_id: str, *, actor: str, now: datetime | None = None + ) -> ResourceViewModel: + """ + 在事务中禁用 Resource 并失效关联授权 + + :param db: 异步数据库会话 + :param resource_id: 资源标识 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ResourceViewModel 详情 + """ + + return await cls._change_resource_status( + db, ResourceStatusModel(resource_id=resource_id, status='1'), actor=actor, now=now + ) + + @classmethod + async def _create_scope( + cls, db: AsyncSession, payload: ScopeModel, *, actor: str, now: datetime | None = None + ) -> ScopeModel: + """ + 在事务中创建 Scope ORM 记录 + + :param db: 异步数据库会话 + :param payload: Scope 创建参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ScopeModel 详情 + :raises OAuthClientManagementError: Scope DTO 校验失败、编码重复或关联 Resource 不存在时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + claims = cls._validate_claims(payload.claims) + if payload.scope_code == 'openid' and ( + payload.scope_type != 'identity' + or payload.resource_id is not None + or payload.status != '0' + or claims != ['sub'] + ): + raise OAuthClientManagementError('openid 权限范围使用固定的身份声明策略') + if await OAuthResourceDao.find_scope_duplicate(db, payload.scope_code): + raise OAuthClientManagementError('权限标识 scope_code 已注册') + resource = None + if payload.scope_type == 'resource': + resource = await cls._get_active_resource(db, payload.resource_id) + scope = SysOAuthScope( + scope_code=payload.scope_code, + scope_name=payload.scope_name, + scope_type=payload.scope_type, + resource_pk=resource.resource_pk if resource else None, + claims=claims, + consent_required=int(payload.consent_required), + sensitive=int(payload.sensitive), + status=payload.status, + create_by=actor_value, + create_time=current, + update_by=actor_value, + update_time=current, + remark=payload.remark, + ) + await OAuthResourceDao.add_scope(db, scope) + await cls._record_audit( + db, OidcAuditEvent.SCOPE_POLICY_CHANGED, actor_value, detail={'action': 'scope_created'} + ) + + return cls._scope_view(scope, resource.resource_id if resource else None) + + @classmethod + async def _get_active_resource(cls, db: AsyncSession, resource_id: str | None) -> SysOAuthResource: + """ + 查询启用的 Resource ORM 记录 + + :param db: 异步数据库会话 + :param resource_id: 资源标识 + :return: 启用的 Resource ORM 记录 + :raises OAuthClientManagementError: Resource 标识为空、未知或 Resource 未启用时抛出 + """ + + if not resource_id: + raise OAuthClientManagementError('资源权限必须绑定已启用的资源') + resource = await OAuthResourceDao.active_resource(db, resource_id) + if resource is None: + raise OAuthClientManagementError('权限范围关联的资源不存在或已停用') + return resource + + @classmethod + async def _update_scope( + cls, db: AsyncSession, payload: ScopeModel, *, actor: str, now: datetime | None = None + ) -> ScopeModel: + """ + 在事务中更新 Scope ORM 记录 + + :param db: 异步数据库会话 + :param payload: Scope 更新参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ScopeModel 详情 + :raises OAuthClientManagementError: Scope 不存在、DTO 校验失败或关联 Resource 不存在时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + claims = cls._validate_claims(payload.claims) + scope = await OAuthResourceDao.get_scope(db, payload.scope_code, for_update=True) + if scope is None: + raise OAuthClientManagementError('权限范围不存在') + if scope.scope_code == 'openid' and ( + payload.scope_type != 'identity' + or payload.resource_id is not None + or payload.status != '0' + or claims != ['sub'] + ): + raise OAuthClientManagementError('openid 权限范围不能停用或更改归属') + resource = None + if payload.scope_type == 'resource': + resource = await cls._get_active_resource(db, payload.resource_id) + resource_pk = resource.resource_pk if resource else None + policy_changed = ( + scope.scope_type != payload.scope_type + or scope.resource_pk != resource_pk + or set(scope.claims or ()) != set(claims) + or bool(scope.consent_required) != payload.consent_required + or bool(scope.sensitive) != payload.sensitive + or scope.status != payload.status + ) + clients = await cls._lock_scope_clients(db, scope.scope_pk) if policy_changed else [] + scope.scope_name = payload.scope_name + scope.scope_type = payload.scope_type + scope.resource_pk = resource_pk + scope.claims = claims + scope.consent_required = int(payload.consent_required) + scope.sensitive = int(payload.sensitive) + scope.status = payload.status + scope.update_by, scope.update_time = actor_value, current + scope.remark = payload.remark + if policy_changed: + cls._bump_clients(clients, actor_value, current) + await OAuthResourceDao.persist_scope_change(db, scope) + if policy_changed: + await cls._record_audit( + db, OidcAuditEvent.SCOPE_POLICY_CHANGED, actor_value, detail={'action': 'scope_updated'} + ) + return cls._scope_view(scope, resource.resource_id if resource else None) + + @classmethod + async def _lock_scope_clients(cls, db: AsyncSession, scope_pk: int) -> list[SysOAuthClient]: + """ + 锁定关联 Scope 的 Client 记录 + + :param db: 异步数据库会话 + :param scope_pk: Scope 主键 + :return: 已锁定的 Client ORM 记录列表 + """ + + return list(await OAuthResourceDao.lock_scope_clients(db, scope_pk)) + + @classmethod + async def detail_scope(cls, db: AsyncSession, scope_code: str) -> ScopeModel: + """ + 查询 OAuth Scope 详情 + + :param db: 异步数据库会话 + :param scope_code: Scope 编码 + :return: ScopeModel 详情 + :raises OAuthClientManagementError: Scope 不存在时抛出 + """ + + scope = await OAuthResourceDao.get_scope(db, scope_code) + if scope is None: + raise OAuthClientManagementError('权限范围不存在') + resource_id = None + if scope.resource_pk is not None: + resource_id = await OAuthResourceDao.resource_id_for_scope(db, scope.resource_pk) + return cls._scope_view(scope, resource_id) + + @classmethod + async def list_scopes(cls, db: AsyncSession, query: ScopePageQueryModel | None = None) -> list[ScopeModel]: + """ + 分页查询 OAuth Scope + + :param db: 异步数据库会话 + :param query: Scope 分页查询参数 + :return: ScopeModel 列表 + """ + + page = query or ScopePageQueryModel() + rows = await OAuthResourceDao.list_scopes_page(db, page) + + return [await cls.detail_scope(db, row.scope_code) for row in rows] + + @classmethod + async def count_scopes(cls, db: AsyncSession, query: ScopePageQueryModel | None = None) -> int: + """ + 统计 OAuth Scope 数量 + + :param db: 异步数据库会话 + :param query: Scope 分页统计参数 + :return: Scope 数量 + """ + + page = query or ScopePageQueryModel() + + return await OAuthResourceDao.count_scopes(db, page) + + @classmethod + async def _change_scope_status( + cls, db: AsyncSession, payload: ScopeStatusModel, *, actor: str, now: datetime | None = None + ) -> ScopeModel: + """ + 在事务中变更 Scope 状态 + + :param db: 异步数据库会话 + :param payload: Scope 状态变更参数 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ScopeModel 详情 + :raises OAuthClientManagementError: Scope 不存在或状态值不合法时抛出 + """ + + actor_value, current = cls._actor(actor), cls._now(now) + scope = await OAuthResourceDao.get_scope(db, payload.scope_code, for_update=True) + if scope is None: + raise OAuthClientManagementError('权限范围不存在') + if scope.scope_code == 'openid' and payload.status != '0': + raise OAuthClientManagementError('openid 权限范围不能停用') + if scope.status != payload.status: + clients = await cls._lock_scope_clients(db, scope.scope_pk) + scope.status = payload.status + scope.update_by, scope.update_time = actor_value, current + cls._bump_clients(clients, actor_value, current) + await OAuthResourceDao.persist_scope_change(db, scope) + await cls._record_audit( + db, OidcAuditEvent.SCOPE_POLICY_CHANGED, actor_value, detail={'action': 'scope_status_changed'} + ) + return await cls.detail_scope(db, scope.scope_code) + + @classmethod + async def _soft_disable_scope( + cls, db: AsyncSession, scope_code: str, *, actor: str, now: datetime | None = None + ) -> ScopeModel: + """ + 在事务中禁用 Scope 并失效关联授权 + + :param db: 异步数据库会话 + :param scope_code: Scope 编码 + :param actor: 操作人标识 + :param now: 当前时间 + :return: ScopeModel 详情 + """ + + return await cls._change_scope_status( + db, ScopeStatusModel(scope_code=scope_code, status='1'), actor=actor, now=now + ) + + @classmethod + async def collect_resource_invalidation_targets( + cls, db: AsyncSession, resource_id: str + ) -> ResourceInvalidationTargets: + """ + 收集 Resource 变更后需撤销的授权记录 + + :param db: 异步数据库会话 + :param resource_id: 资源标识 + :return: Resource 失效目标集合 + :raises OAuthClientManagementError: Resource 不存在时抛出 + """ + + resource = await OAuthResourceDao.get_resource(db, resource_id) + if resource is None: + raise OAuthClientManagementError('资源不存在') + client_rows = [ + *await OAuthResourceDao.client_ids_for_resource(db, resource.resource_pk), + *await OAuthResourceDao.client_ids_for_resource_scope(db, resource.resource_pk), + ] + client_pks = {row[0] for row in client_rows} + client_ids = tuple(sorted({row[1] for row in client_rows})) + if not client_pks: + return ResourceInvalidationTargets(resource.resource_id, (), (), ()) + grants = await OAuthResourceDao.active_grants(db, list(client_pks)) + grant_ids = tuple( + grant.grant_id + for grant in grants + if resource.resource_id in (grant.granted_resources or []) + or resource.audience in (grant.granted_resources or []) + ) + refresh_tokens = await OAuthResourceDao.active_refresh_tokens(db, list(client_pks)) + refresh_ids = tuple( + token.token_id + for token in refresh_tokens + if resource.resource_id in (token.resources or []) or resource.audience in (token.resources or []) + ) + + return ResourceInvalidationTargets(resource.resource_id, client_ids, grant_ids, refresh_ids) diff --git a/ruoyi-fastapi-backend/module_identity/service/oauth_session_management_service.py b/ruoyi-fastapi-backend/module_identity/service/oauth_session_management_service.py new file mode 100644 index 000000000..7b6315bf2 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/oauth_session_management_service.py @@ -0,0 +1,388 @@ +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from exceptions.exception import ServiceException +from module_identity.dao.identity_user_dao import IdentityUserDao +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.oauth_grant_do import SysOAuthAccessPolicy, SysOAuthGrant, SysSsoSession +from module_identity.entity.vo.oauth_session_vo import ( + AccessPolicyModel, + AccessPolicyPageQueryModel, + GrantModel, + GrantPageQueryModel, + SessionPageQueryModel, + SsoSessionModel, +) +from module_identity.service.audit_service import AuditService +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.session_service import SsoSessionService +from utils.time_util import TimezoneUtil + + +class OAuthSessionManagementService: + """ + OAuth Session 管理模块服务层 + """ + + @staticmethod + def session_view(row: SysSsoSession) -> SsoSessionModel: + """ + 将 SSO Session ORM 记录转换为 SsoSessionModel + + :param row: SSO Session ORM 记录 + :return: Session 管理视图 + """ + + status = row.status + if ( + status == 'active' + and min(TimezoneUtil.to_utc(row.idle_expires_at), TimezoneUtil.to_utc(row.absolute_expires_at)) + <= TimezoneUtil.utc_now() + ): + status = 'expired' + return SsoSessionModel( + sid=row.sid, + user_id=row.user_id, + subject_id=row.subject_id, + auth_version=row.auth_version, + auth_time=row.auth_time, + last_seen_at=row.last_seen_at, + idle_expires_at=row.idle_expires_at, + absolute_expires_at=row.absolute_expires_at, + acr=row.acr, + amr=list(row.amr or []), + remember_me=bool(getattr(row, 'remember_me', False)), + ip_address=getattr(row, 'ip_address', None), + status=status, + revoked_at=getattr(row, 'revoked_at', None), + revoke_reason=getattr(row, 'revoke_reason', None), + create_time=getattr(row, 'create_time', None), + ) + + @staticmethod + def grant_view(row: SysOAuthGrant, client_id: str | None, policy: SysOAuthAccessPolicy | None = None) -> GrantModel: + """ + 将 OAuth Grant ORM 记录转换为 GrantModel + + :param row: OAuth Grant ORM 记录 + :param client_id: 客户端标识 + :param policy: 当前用户对应用的独立访问策略 + :return: Grant 管理视图 + """ + + status = row.status + if ( + status == 'active' + and row.expires_at is not None + and TimezoneUtil.to_utc(row.expires_at) <= TimezoneUtil.utc_now() + ): + status = 'expired' + return GrantModel( + grant_id=row.grant_id, + user_id=row.user_id, + subject_id=row.subject_id, + client_id=client_id or '', + granted_scopes=list(row.granted_scopes or []), + granted_resources=list(row.granted_resources or []), + remembered_scopes=list(getattr(row, 'remembered_scopes', None) or []), + remembered_resources=list(getattr(row, 'remembered_resources', None) or []), + access_status=policy.access_status if policy else 'allowed', + access_reason=policy.reason if policy else None, + status=status, + client_policy_version=getattr(row, 'client_policy_version', None), + consented_at=row.consented_at, + last_used_at=getattr(row, 'last_used_at', None), + expires_at=row.expires_at, + revoke_reason=getattr(row, 'revoke_reason', None), + ) + + @staticmethod + async def list_sessions(db: AsyncSession, query: SessionPageQueryModel) -> tuple[list[SsoSessionModel], int]: + """ + 分页查询 SSO Session + + :param db: 异步数据库会话 + :param query: Session 分页查询参数 + :return: Session 管理视图列表和总数 + """ + + rows = await SsoSessionDao.list_page( + db, + user_id=query.user_id, + ip_address=query.ip_address, + status=query.status, + start_time=query.start_time, + end_time=query.end_time, + offset=(query.page_num - 1) * query.page_size, + limit=query.page_size, + ) + total = await SsoSessionDao.count( + db, + user_id=query.user_id, + ip_address=query.ip_address, + status=query.status, + start_time=query.start_time, + end_time=query.end_time, + ) + + client_ids = await SsoSessionDao.client_ids_for_sids(db, [row.sid for row in rows]) + return [ + OAuthSessionManagementService.session_view(row).model_copy( + update={'client_ids': client_ids.get(row.sid, [])} + ) + for row in rows + ], total + + @staticmethod + async def get_session(db: AsyncSession, sid: str) -> SsoSessionModel | None: + """ + 按 sid 查询 SSO Session + + :param db: 异步数据库会话 + :param sid: Session 标识 + :return: Session 管理视图,不存在时返回 None + """ + + row = await SsoSessionDao.get_by_sid(db, sid) + if row is None: + return None + return OAuthSessionManagementService.session_view(row).model_copy( + update={'client_ids': list(await SsoSessionDao.client_ids_for_sid(db, sid))} + ) + + @staticmethod + async def list_grants(db: AsyncSession, query: GrantPageQueryModel) -> tuple[list[GrantModel], int]: + """ + 分页查询 OAuth Grant + + :param db: 异步数据库会话 + :param query: Grant 分页查询参数 + :return: Grant 管理视图列表和总数 + """ + + rows = await OAuthGrantDao.list_page( + db, + user_id=query.user_id, + client_id=query.client_id, + status=query.status, + access_status=query.access_status, + offset=(query.page_num - 1) * query.page_size, + limit=query.page_size, + ) + total = await OAuthGrantDao.count( + db, user_id=query.user_id, client_id=query.client_id, status=query.status, access_status=query.access_status + ) + keys = {int(row.client_pk) for row in rows} + mapping = await OAuthClientDao.id_map(db, list(keys)) + + policies = { + (row.user_id, row.client_pk): row + for row in await OAuthAccessPolicyDao.list_for_grants(db, [(row.user_id, row.client_pk) for row in rows]) + } + return [ + OAuthSessionManagementService.grant_view( + row, mapping.get(int(row.client_pk)), policies.get((row.user_id, row.client_pk)) + ) + for row in rows + ], total + + @staticmethod + async def list_access_policies( + db: AsyncSession, query: AccessPolicyPageQueryModel + ) -> tuple[list[AccessPolicyModel], int]: + """ + 查询独立访问策略,包含尚未授权过的用户与应用 + + :param db: 异步数据库会话 + :param query: 访问策略分页查询参数 + :return: 访问策略管理视图列表和总数 + """ + + rows = await OAuthAccessPolicyDao.list_page(db, query) + total = await OAuthAccessPolicyDao.count(db, query) + return [ + AccessPolicyModel( + user_id=row.user_id, + user_name=user_name, + client_id=client_id, + client_name=client_name, + access_status=row.access_status, + reason=row.reason, + update_by=row.update_by, + update_time=row.update_time, + ) + for row, user_name, client_id, client_name in rows + ], total + + @staticmethod + async def get_grant(db: AsyncSession, grant_id: str) -> GrantModel | None: + """ + 按 grant_id 查询 OAuth Grant + + :param db: 异步数据库会话 + :param grant_id: Grant 标识 + :return: Grant 管理视图,不存在时返回 None + """ + + row = await OAuthGrantDao.get_by_grant_id(db, grant_id) + if row is None: + return None + client_id = await OAuthClientDao.id_for_pk(db, row.client_pk) + + policy = await OAuthAccessPolicyDao.get(db, row.user_id, row.client_pk) + return OAuthSessionManagementService.grant_view(row, client_id, policy) + + @staticmethod + async def revoke_user(db: AsyncSession, redis: Redis, user_id: int, actor: str, reason: str) -> int: + """ + 撤销用户全部在线及自然过期的 SSO Session + + :param db: 异步数据库会话 + :param redis: SSO Session 热缓存客户端 + :param user_id: 用户内部标识 + :param actor: 管理操作者标识 + :param reason: 撤销原因 + :return: 实际发生状态变更的 Session 数量 + :raises ServiceException: 撤销用户 Session 事务失败 + """ + + coordinator = AfterCommitCoordinator() + try: + count = await SsoSessionService.revoke_user(db, redis, user_id, reason=reason, coordinator=coordinator) + await AuditService.record( + db, + OidcAuditEvent.SESSION_REVOKED, + 'success', + user_id=user_id, + detail={'actor': actor, 'reason': reason}, + ) + await coordinator.commit(db) + return count + except Exception as exc: + await coordinator.rollback(db) + raise ServiceException(message='批量撤销会话失败') from exc + + @staticmethod + async def revoke_sessions(db: AsyncSession, redis: Redis, sids: list[str], actor: str, reason: str) -> int: + """ + 按 sid 集合撤销 SSO Session + + :param db: 异步数据库会话 + :param redis: SSO Session 热缓存客户端 + :param sids: Session 标识列表 + :param actor: 管理操作者标识 + :param reason: 撤销原因 + :return: 实际发生状态变更的 Session 数量 + :raises ServiceException: 批量撤销指定 Session 事务失败 + """ + + coordinator, count = AfterCommitCoordinator(), 0 + try: + for sid in sids: + changed = await SsoSessionService.revoke(db, redis, sid, reason=reason, coordinator=coordinator) + count += int(changed) + if changed: + await AuditService.record( + db, + OidcAuditEvent.SESSION_REVOKED, + 'success', + sid=sid, + detail={'actor': actor, 'reason': reason}, + ) + await coordinator.commit(db) + return count + except Exception as exc: + await coordinator.rollback(db) + raise ServiceException(message='批量撤销会话失败') from exc + + @staticmethod + async def revoke_grants(db: AsyncSession, grant_ids: list[str], actor: str, reason: str) -> int: + """ + 撤销选中记录所属用户对应用的全部现有授权 + + :param db: 异步数据库会话 + :param grant_ids: 用于确定用户和应用的 Grant 标识集合 + :param actor: 操作人标识 + :param reason: 撤销原因 + :return: 实际撤销的 Grant 数量 + :raises ServiceException: 批量撤销 Grant 事务失败 + """ + + try: + count = 0 + for client_pk, user_id in await OAuthGrantDao.targets(db, grant_ids): + await OAuthAccessPolicyDao.lock_client(db, client_pk) + revoked = await OAuthGrantDao.revoke_for_user_client(db, user_id, client_pk, reason) + client_id = await OAuthClientDao.id_for_pk(db, client_pk) + for grant_id in revoked: + await AuditService.record( + db, + OidcAuditEvent.GRANT_REVOKED, + 'success', + client_id=client_id, + user_id=user_id, + grant_id=grant_id, + detail={'actor': actor, 'reason': reason}, + ) + count += len(revoked) + await db.commit() + return count + except Exception as exc: + await db.rollback() + raise ServiceException(message='批量撤销授权失败') from exc + + @staticmethod + async def set_access(db: AsyncSession, user_id: int, client_id: str, blocked: bool, actor: str, reason: str) -> int: + """ + 更新用户对应用的访问策略,禁止时撤销全部授权,解除时保留撤销状态 + + :param db: 异步数据库会话 + :param user_id: 用户编号 + :param client_id: Client 公开标识 + :param blocked: 是否禁止访问 + :param actor: 操作人标识 + :param reason: 操作原因 + :return: 本次撤销的授权数量 + :raises ServiceException: 用户或应用不存在,或事务失败 + """ + + try: + client = await OAuthClientDao.get_by_client_id(db, client_id) + user = await IdentityUserDao.get_user(db, user_id) + if client is None or user is None or user.del_flag != '0': + raise ServiceException(message='用户或应用不存在') + await OAuthAccessPolicyDao.lock_client(db, client.client_pk) + await OAuthAccessPolicyDao.set_status(db, user_id, client.client_pk, blocked, actor, reason) + revoked = ( + await OAuthGrantDao.revoke_for_user_client(db, user_id, client.client_pk, reason) if blocked else [] + ) + for grant_id in revoked: + await AuditService.record( + db, + OidcAuditEvent.GRANT_REVOKED, + 'success', + client_id=client_id, + user_id=user_id, + grant_id=grant_id, + detail={'actor': actor, 'reason': reason}, + ) + await AuditService.record( + db, + OidcAuditEvent.CLIENT_ACCESS_BLOCKED if blocked else OidcAuditEvent.CLIENT_ACCESS_ALLOWED, + 'success', + client_id=client_id, + user_id=user_id, + detail={'actor': actor, 'reason': reason}, + ) + await db.commit() + return len(revoked) + except ServiceException: + await db.rollback() + raise + except Exception as exc: + await db.rollback() + raise ServiceException(message='更新用户应用访问策略失败') from exc diff --git a/ruoyi-fastapi-backend/module_identity/service/runtime_service.py b/ruoyi-fastapi-backend/module_identity/service/runtime_service.py new file mode 100644 index 000000000..9f16029bf --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/runtime_service.py @@ -0,0 +1,382 @@ +import asyncio +from collections.abc import Awaitable, Callable +from dataclasses import dataclass +from datetime import datetime, timedelta +from time import monotonic + +from fastapi import FastAPI +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.database import DataSourceRegistry +from config.env import OidcConfig +from config.scheduler.manager import SchedulerManager +from module_identity.dao.oauth_audit_dao import OAuthAuditDao +from module_identity.service.audit_service import AuditService +from module_identity.service.key_service import KeyService, KeyServiceError +from module_identity.service.oauth_management_service import OAuthClientManagementService +from module_identity.service.session_service import LogoutService +from utils.log_util import logger +from utils.time_util import TimezoneUtil + + +@dataclass(frozen=True, slots=True) +class OidcReadiness: + """ + OIDC 协议运行时就绪快照 + + ``enabled`` 只表示部署配置已开启协议,``ready`` 还要求数据库中 + 存在可加载的 active 签名密钥。两者分离后,管理后台可以在协议 + 尚未就绪时继续启动并完成首把密钥初始化。 + """ + + enabled: bool + ready: bool + reason: str + checked_at: datetime + + def as_dict(self) -> dict[str, object]: + """ + 转换为可安全返回给管理端的状态字段 + + :return: OIDC 运行时状态字段映射 + """ + + return { + 'enabled': self.enabled, + 'ready': self.ready, + 'readinessReason': self.reason, + 'readinessCheckedAt': self.checked_at, + } + + +class OidcRuntimeService: + """ + OIDC 运行时模块服务层 + """ + + _AUDIT_ARCHIVE_INTERVAL_SECONDS = 86400 + _READINESS_CACHE_SECONDS = 5 + _CORS_CACHE_SECONDS = 5 + _CORS_VERSION_KEY = 'oidc:runtime:cors_version' + + @staticmethod + def _disabled_readiness() -> OidcReadiness: + """ + 构造 OIDC 协议关闭状态 + + :return: OIDC 协议关闭时的就绪快照 + """ + + return OidcReadiness( + enabled=False, + ready=False, + reason='disabled', + checked_at=TimezoneUtil.utc_now(), + ) + + @classmethod + async def inspect_readiness(cls, db: AsyncSession) -> OidcReadiness: + """ + 检查 OIDC 是否具备对外提供协议服务的签名能力 + + 密钥不存在或密钥材料不可用都会被收敛为未就绪状态,不会影响 + 管理后台启动;数据库基础设施异常仍向上抛出,由应用通用健康 + 机制处理。 + + :param db: 异步数据库会话 + :return: OIDC 运行时就绪快照 + """ + + if not OidcConfig.oidc_enabled: + return cls._disabled_readiness() + try: + await KeyService.get_signing_key(db) + except KeyServiceError: + return OidcReadiness( + enabled=True, + ready=False, + reason='signing_key_unavailable', + checked_at=TimezoneUtil.utc_now(), + ) + return OidcReadiness( + enabled=True, + ready=True, + reason='ready', + checked_at=TimezoneUtil.utc_now(), + ) + + @classmethod + async def refresh_readiness(cls, app: FastAPI, db: AsyncSession | None = None) -> OidcReadiness: + """ + 刷新并保存当前 Worker 的 OIDC 就绪快照 + + :param app: FastAPI 应用 + :param db: 可复用的异步数据库会话 + :return: 刷新后的 OIDC 运行时就绪快照 + """ + + if not OidcConfig.oidc_enabled: + readiness = cls._disabled_readiness() + elif db is not None: + readiness = await cls.inspect_readiness(db) + else: + async with DataSourceRegistry.session() as current_db: + readiness = await cls.inspect_readiness(current_db) + app.state.oidc_readiness = readiness + + return readiness + + @classmethod + async def cached_readiness(cls, app: FastAPI, db: AsyncSession) -> OidcReadiness: + """ + 读取带短时缓存的 OIDC 就绪状态 + + 缓存避免每个协议请求都重复解密私钥;未就绪状态也会很快重新 + 检查,因此其他 Worker 在首把密钥激活后无需重启即可恢复。 + + :param app: FastAPI 应用 + :param db: 异步数据库会话 + :return: 当前 Worker 的 OIDC 运行时就绪快照 + """ + + current = TimezoneUtil.utc_now() + readiness = getattr(app.state, 'oidc_readiness', None) + if isinstance(readiness, OidcReadiness): + age = (current - readiness.checked_at).total_seconds() + if age < cls._READINESS_CACHE_SECONDS: + return readiness + lock = getattr(app.state, 'oidc_readiness_lock', None) + if lock is None: + lock = asyncio.Lock() + app.state.oidc_readiness_lock = lock + async with lock: + readiness = getattr(app.state, 'oidc_readiness', None) + current = TimezoneUtil.utc_now() + if isinstance(readiness, OidcReadiness): + age = (current - readiness.checked_at).total_seconds() + if age < cls._READINESS_CACHE_SECONDS: + return readiness + return await cls.refresh_readiness(app, db) + + @classmethod + async def validate_runtime(cls, app: FastAPI | None = None) -> OidcReadiness: + """ + 检查启动阶段的 OIDC 就绪状态,但不阻断管理后台启动 + + 协议启用但尚无 active 密钥时记录明确告警,公共 OIDC 端点随后 + 返回 503;管理员仍可登录后台创建并激活首把密钥。 + + :param app: 需要保存就绪快照的 FastAPI 应用 + :return: 启动阶段的 OIDC 运行时就绪快照 + """ + + if not OidcConfig.oidc_enabled: + readiness = cls._disabled_readiness() + else: + async with DataSourceRegistry.session() as db: + readiness = await cls.inspect_readiness(db) + if app is not None: + app.state.oidc_readiness = readiness + app.state.oidc_readiness_lock = asyncio.Lock() + if readiness.enabled and not readiness.ready: + logger.warning('统一认证中心已启用但尚未就绪;请创建并激活有效签名密钥,就绪前协议接口返回 503') + return readiness + + @staticmethod + async def load_cors_origins() -> tuple[str, ...]: + """ + 加载启用的 CORS Origin + + :return: 启用的 CORS Origin 元组 + """ + + async with DataSourceRegistry.session() as db: + origins = await OAuthClientManagementService.list_active_cors_origins(db) + return tuple(origins) + + @classmethod + async def refresh_cors_snapshot(cls, app: FastAPI) -> tuple[str, ...]: + """ + 刷新应用的 CORS Origin 快照 + + :param app: FastAPI 应用 + :return: 刷新后的 CORS Origin 元组 + """ + + if not OidcConfig.oidc_enabled: + app.state.oidc_registered_cors_origins = () + return () + origins = await cls.load_cors_origins() + app.state.oidc_registered_cors_origins = origins + app.state.oidc_cors_loaded_at = monotonic() + + return origins + + @classmethod + async def ensure_cors_snapshot(cls, app: FastAPI) -> None: + """ + 刷新当前Worker的客户端跨域策略快照 + + 通过共享版本号发现变更,并以有时限的数据库缓存兜底。 + + :param app: FastAPI应用对象 + :return: 无 + """ + + if not OidcConfig.oidc_enabled: + app.state.oidc_registered_cors_origins = () + return + lock = getattr(app.state, 'oidc_cors_lock', None) + if lock is None: + lock = asyncio.Lock() + app.state.oidc_cors_lock = lock + async with lock: + redis = getattr(app.state, 'redis', None) + revision = None + try: + if redis is not None: + revision = await redis.get(cls._CORS_VERSION_KEY) + except Exception: + # 共享版本源不可用时,不延长旧跨域白名单的缓存期限 + app.state.oidc_cors_loaded_at = None + if isinstance(revision, bytes): + revision = revision.decode('ascii') + loaded_at = getattr(app.state, 'oidc_cors_loaded_at', None) + fresh = loaded_at is not None and monotonic() - loaded_at < cls._CORS_CACHE_SECONDS + if fresh and revision == getattr(app.state, 'oidc_cors_revision', None): + return + try: + await cls.refresh_cors_snapshot(app) + app.state.oidc_cors_revision = revision + except Exception: + app.state.oidc_registered_cors_origins = () + app.state.oidc_cors_loaded_at = None + logger.error('OAuth 已注册跨域来源快照刷新失败') + + @classmethod + def cors_snapshot_callback(cls, app: FastAPI) -> Callable[[], Awaitable[None]]: + """ + 构造提交后刷新 CORS 快照的回调 + + :param app: 需要更新 CORS Origin 快照的 FastAPI 应用 + :return: 无参数异步 CORS 快照刷新回调 + """ + + async def refresh() -> None: + """ + 刷新应用 CORS Origin 快照 + + :return: None + """ + + redis = getattr(app.state, 'redis', None) + try: + if redis is not None: + await redis.incr(cls._CORS_VERSION_KEY) + except Exception: + logger.error('OAuth 跨域配置失效通知发布失败,各工作进程将在缓存有效期内刷新') + app.state.oidc_cors_loaded_at = None + await cls.ensure_cors_snapshot(app) + + return refresh + + @classmethod + async def start_background_tasks(cls, app: FastAPI) -> None: + """ + 每个 Worker 启动待命循环,仅当前 Leader 执行业务;重新获租后自动恢复 + + :param app: FastAPI 应用 + :return: None + """ + + redis = app.state.redis + app.state.oidc_key_lifecycle_task = ( + asyncio.create_task(cls.key_lifecycle_loop(app)) if OidcConfig.oidc_enabled else None + ) + app.state.oidc_backchannel_retry_task = ( + asyncio.create_task(cls.backchannel_retry_loop(redis)) if OidcConfig.oidc_enabled else None + ) + + @classmethod + async def key_lifecycle_loop(cls, app: FastAPI) -> None: + """ + 运行签名密钥生命周期后台循环 + + :param app: FastAPI 应用,包含 Redis 与 OIDC 就绪状态 + :return: None + """ + + last_archive_at: datetime | None = None + redis = app.state.redis + while True: + await asyncio.sleep(60) + if not SchedulerManager.is_application_leader(): + continue + try: + async with DataSourceRegistry.session() as db: + activated = await KeyService.activate_due(db, redis=redis) + if activated: + await cls.refresh_readiness(app, db) + changed = await KeyService.retire_due(db) + if changed: + await AuditService.record( + db, + OidcAuditEvent.SIGNING_KEY_ROTATED, + 'success', + detail={'action': 'retired_due', 'count': changed, 'actor': 'system:lifecycle'}, + ) + await db.commit() + current = TimezoneUtil.utc_now() + if ( + last_archive_at is None + or (current - last_archive_at).total_seconds() >= cls._AUDIT_ARCHIVE_INTERVAL_SECONDS + ): + before = current - timedelta(days=OidcConfig.oidc_audit_retention_days) + while await OAuthAuditDao.archive_before(db, before): + await db.commit() + last_archive_at = current + except asyncio.CancelledError: + raise + except Exception: + logger.error('OIDC 签名密钥生命周期维护任务执行失败') + + @staticmethod + async def backchannel_retry_loop(redis: object) -> None: + """ + 运行 Back-Channel 重试后台循环 + + :param redis: 异步 Redis 客户端 + :return: None + """ + + while True: + await asyncio.sleep(1) + if not SchedulerManager.is_application_leader(): + continue + try: + async with DataSourceRegistry.session() as db: + await LogoutService.consume_backchannel_retry(db, redis, max_items=32) + except asyncio.CancelledError: + raise + except Exception: + logger.error('OIDC 后端退出通知重试任务执行失败') + + @staticmethod + async def stop_background_tasks(app: FastAPI) -> None: + """ + 停止 OIDC 后台任务 + + :param app: FastAPI 应用 + :return: None + """ + + for name in ('oidc_key_lifecycle_task', 'oidc_backchannel_retry_task'): + task = getattr(app.state, name, None) + if task is None: + continue + task.cancel() + try: + await task + except asyncio.CancelledError: + pass diff --git a/ruoyi-fastapi-backend/module_identity/service/session_service.py b/ruoyi-fastapi-backend/module_identity/service/session_service.py new file mode 100644 index 000000000..0398bd7ea --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/session_service.py @@ -0,0 +1,1895 @@ +import asyncio +import hmac +import json +import logging +import secrets +from collections.abc import Awaitable, Callable, Sequence +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Any +from urllib.parse import urlsplit +from uuid import uuid4 + +import jwt +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from module_identity.dao.identity_subject_dao import IdentitySubjectDao +from module_identity.dao.identity_user_dao import IdentityUserDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_token_dao import OAuthTokenDao +from module_identity.dao.oidc_key_dao import OidcKeyDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.oauth_grant_do import SysSsoSession +from module_identity.redis_keys import OidcRedisKey +from module_identity.security.backchannel_transport import ( + BackchannelNotifier, + PermanentBackchannelError, + send_backchannel_once, +) +from module_identity.security.jwt_profile import ( + BACKCHANNEL_LOGOUT_EVENT, + JwtProfileError, + decode_id_token_hint, + encode_logout_token, +) +from module_identity.security.uri_validator import ( + is_safe_backchannel_uri, + public_dns_addresses, + public_dns_only, +) +from module_identity.service.audit_service import AuditService +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.key_service import KeyService, KeyServiceError +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + +_COOKIE_PREFIX = 'ss1' +_COOKIE_SECRET_BYTES = 32 +_MAX_IP_ADDRESS_LENGTH = 128 +_MAX_ACR_LENGTH = 100 +_MAX_AMR_LENGTH = 50 +_MAX_REMEMBER_SECONDS = 7 * 24 * 60 * 60 +_SESSION_TTL_FIELDS = ('oidc_sso_idle_seconds', 'oidc_sso_absolute_seconds', 'oidc_sso_remember_absolute_seconds') +_SESSION_CACHE_TTL_FLOOR = 1 +_SESSION_STATUS_ACTIVE = 'active' +_SESSION_STATUS_REVOKED = 'revoked' +_SESSION_STATUS_EXPIRED = 'expired' + + +class SsoSessionError(ValueError): + """ + SSO Session 创建、校验和撤销错误类型 + """ + + +@dataclass(frozen=True) +class SsoSessionSnapshot: + """ + 记录 SSO Session 缓存所需的不可变字段 + """ + + sid: str + user_id: int + subject_id: str + auth_version: int + auth_time: datetime | None + last_seen_at: datetime | None + idle_expires_at: datetime | None + absolute_expires_at: datetime | None + acr: str + amr: tuple[str, ...] + remember_me: bool + status: str + session_secret_hash: str + + +class SsoSessionService: + """ + OIDC SSO Session 模块服务层 + """ + + @staticmethod + def _require_coordinator(coordinator: AfterCommitCoordinator | None) -> AfterCommitCoordinator: + """ + 确认提交后副作用协调器可用 + + :param coordinator: 只登记提交后副作用的协调器 + :return: 通过校验的协调器 + :raises SsoSessionError: 未提供有效协调器 + """ + + if not isinstance(coordinator, AfterCommitCoordinator): + raise SsoSessionError('缺少事务提交后副作用协调器') + return coordinator + + @staticmethod + def _validate_config_ttls() -> None: + """ + 检查会话相关 TTL 配置均为正整数 + + :return: None + :raises SsoSessionError: TTL 不是正整数 + """ + + for field_name in _SESSION_TTL_FIELDS: + value = getattr(OidcConfig, field_name, None) + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise SsoSessionError(f'{field_name} 必须为正整数') + + @staticmethod + def _remaining_ttl(absolute_expires_at: datetime, now: datetime) -> int: + """ + 根据绝对过期时间计算 Redis 缓存剩余 TTL + + :param absolute_expires_at: Session 绝对过期时间 + :param now: 当前项目时间 + :return: 至少一秒的 Redis TTL + :raises SsoSessionError: Session 已过期 + """ + + seconds = int((absolute_expires_at - now).total_seconds()) + if seconds < _SESSION_CACHE_TTL_FLOOR: + raise SsoSessionError('会话已过期') + return seconds + + @classmethod + def _snapshot(cls, row: SysSsoSession) -> SsoSessionSnapshot: + """ + 从 ORM 行提取不依赖数据库会话的 Session 快照 + + :param row: 数据库 Session ORM 行 + :return: 不再依赖 ORM 生命周期的 Session 快照 + """ + + return SsoSessionSnapshot( + sid=row.sid, + user_id=row.user_id, + subject_id=row.subject_id, + auth_version=row.auth_version, + auth_time=TimezoneUtil.to_optional_utc(row.auth_time), + last_seen_at=TimezoneUtil.to_optional_utc(row.last_seen_at), + idle_expires_at=TimezoneUtil.to_optional_utc(row.idle_expires_at), + absolute_expires_at=TimezoneUtil.to_optional_utc(row.absolute_expires_at), + acr=row.acr, + amr=tuple(row.amr or ()), + remember_me=bool(row.remember_me), + status=row.status, + session_secret_hash=row.session_secret_hash, + ) + + @classmethod + def _cache_payload(cls, row: SsoSessionSnapshot) -> dict[str, Any]: + """ + 将 Session 快照编码为 Redis 缓存载荷 + + :param row: 数据库 Session + :return: 可 JSON 序列化的缓存摘要 + """ + + return { + 'sid': row.sid, + 'user_id': row.user_id, + 'subject_id': row.subject_id, + 'auth_version': row.auth_version, + 'auth_time': TimezoneUtil.to_optional_utc(row.auth_time).isoformat() + if TimezoneUtil.to_optional_utc(row.auth_time) + else None, + 'last_seen_at': TimezoneUtil.to_optional_utc(row.last_seen_at).isoformat() + if TimezoneUtil.to_optional_utc(row.last_seen_at) + else None, + 'idle_expires_at': TimezoneUtil.to_optional_utc(row.idle_expires_at).isoformat() + if TimezoneUtil.to_optional_utc(row.idle_expires_at) + else None, + 'absolute_expires_at': TimezoneUtil.to_optional_utc(row.absolute_expires_at).isoformat() + if TimezoneUtil.to_optional_utc(row.absolute_expires_at) + else None, + 'acr': row.acr, + 'amr': row.amr, + 'remember_me': row.remember_me, + 'status': row.status, + } + + @classmethod + async def _cache_row(cls, redis: Redis, row: SsoSessionSnapshot, now: datetime) -> None: + """ + 将 Session 快照写入会话、Cookie 和用户索引缓存 + + :param redis: 异步 Redis 客户端 + :param row: SSO Session 快照 + :param now: 当前时间 + :return: None + """ + + ttl = cls._remaining_ttl(TimezoneUtil.to_optional_utc(row.absolute_expires_at) or now, now) + payload = json.dumps(cls._cache_payload(row), ensure_ascii=False, separators=(',', ':')) + await redis.set(OidcRedisKey.sso_session(row.sid), payload, ex=ttl) + await redis.set(OidcRedisKey.sso_cookie(row.session_secret_hash), row.sid, ex=ttl) + user_sessions_key = OidcRedisKey.user_sessions(row.user_id) + await redis.sadd(user_sessions_key, row.sid) + current_ttl = await redis.ttl(user_sessions_key) + if current_ttl < ttl: + await redis.expire(user_sessions_key, ttl) + + @classmethod + async def _clear_cache(cls, redis: Redis, row: SsoSessionSnapshot, extra_secret_hash: str | None = None) -> None: + """ + 删除 Session、Cookie 和用户索引缓存 + + :param redis: 异步 Redis 客户端 + :param row: SSO Session 快照 + :param extra_secret_hash: 需要额外清理的 Cookie Secret 摘要 + :return: None + """ + + hashes = {row.session_secret_hash} + if extra_secret_hash: + hashes.add(extra_secret_hash) + await redis.delete(OidcRedisKey.sso_session(row.sid)) + await redis.srem(OidcRedisKey.user_sessions(row.user_id), row.sid) + for digest in hashes: + await redis.delete(OidcRedisKey.sso_cookie(digest)) + + @classmethod + async def _publish_revoked(cls, redis: Redis, sid: str, reason: str) -> None: + """ + 发布 Session 已撤销事件 + + :param redis: 异步 Redis 客户端 + :param sid: Session 标识 + :param reason: 撤销原因 + :return: None + """ + + await redis.publish( + OidcRedisKey.event_session_revoked(), + json.dumps({'event': 'session_revoked', 'sid': sid, 'reason': reason}, separators=(',', ':')), + ) + + @classmethod + async def _best_effort_cleanup( + cls, + redis: Redis, + row: SsoSessionSnapshot, + *, + presented_digest: str | None = None, + reason: str | None = None, + ) -> None: + """ + 尽力清理 Session 缓存并发布撤销事件 + + :param redis: 异步 Redis 客户端 + :param row: SSO Session 快照 + :param presented_digest: 当前 Cookie 摘要 + :param reason: 撤销原因 + :return: None + """ + + try: + await cls._clear_cache(redis, row, presented_digest) + except Exception: + pass + if reason is not None: + try: + await cls._publish_revoked(redis, row.sid, reason) + except Exception: + pass + + @classmethod + async def create( # noqa: PLR0913 + cls, + db: AsyncSession, + redis: Redis, + user_id: int, + subject_id: str, + auth_version: int, + acr: str, + amr: Sequence[str], + *, + ip_address: str | None = None, + user_agent: str | None = None, + remember_me: bool = False, + pepper: str | bytes = OidcConfig.oidc_token_hash_pepper, + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> tuple[str, SysSsoSession]: + """ + 校验身份状态并创建 SSO Session 与 Cookie + + :param db: 异步数据库会话,调用方负责提交事务 + :param redis: Redis 热缓存客户端 + :param user_id: 本地用户 ID + :param subject_id: 稳定 OIDC Subject + :param auth_version: 创建时的身份安全版本 + :param acr: 认证上下文 + :param amr: 认证方式列表 + :param ip_address: 登录来源 IP + :param user_agent: 登录 User-Agent,仅保存摘要 + :param remember_me: 是否使用 remember-me 绝对 TTL + :param pepper: OIDC Token Pepper + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行缓存写入 + :return: ``(Cookie 原文, 数据库 Session)`` + :raises SsoSessionError: Session 状态或事务操作不符合要求 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + cls._validate_config_ttls() + if not isinstance(user_id, int) or isinstance(user_id, bool) or user_id <= 0: + raise SsoSessionError('用户编号 user_id 无效') + if not isinstance(subject_id, str) or not OidcUtil.is_rfc4122_uuid(subject_id): + raise SsoSessionError('用户主体标识 subject_id 无效') + if not isinstance(auth_version, int) or isinstance(auth_version, bool) or auth_version < 1: + raise SsoSessionError('身份安全版本 auth_version 无效') + if not isinstance(acr, str) or not acr.strip() or len(acr) > _MAX_ACR_LENGTH: + raise SsoSessionError('认证上下文 acr 无效') + if not isinstance(amr, Sequence) or isinstance(amr, (str, bytes)) or not amr: + raise SsoSessionError('认证方式列表 amr 无效') + if any(not isinstance(item, str) or not item.strip() or len(item) > _MAX_AMR_LENGTH for item in amr): + raise SsoSessionError('认证方式列表 amr 无效') + if not isinstance(remember_me, bool): + raise SsoSessionError('保持登录标识 remember_me 无效') + if ip_address is not None and (not isinstance(ip_address, str) or len(ip_address) > _MAX_IP_ADDRESS_LENGTH): + raise SsoSessionError('IP 地址无效') + try: + OidcUtil.session_pepper_bytes(pepper) + except ValueError as exc: + raise SsoSessionError(str(exc)) from exc + user = await IdentityUserDao.get_user(db, user_id) + subject = await IdentitySubjectDao.get_by_user_id(db, user_id) + if ( + user is None + or user.status != '0' + or user.del_flag != '0' + or subject is None + or subject.subject_id != subject_id + or subject.auth_version != auth_version + ): + raise SsoSessionError('用户身份安全状态无效') + absolute_seconds = OidcConfig.oidc_sso_absolute_seconds + if remember_me: + absolute_seconds = min(OidcConfig.oidc_sso_remember_absolute_seconds, _MAX_REMEMBER_SECONDS) + absolute = current + timedelta(seconds=absolute_seconds) + idle = min(current + timedelta(seconds=OidcConfig.oidc_sso_idle_seconds), absolute) + sid = str(uuid4()) + secret = secrets.token_urlsafe(_COOKIE_SECRET_BYTES) + try: + secret_hash = OidcUtil.session_secret_digest(secret, pepper) + user_agent_hash = OidcUtil.session_user_agent_digest(user_agent, pepper) + except ValueError as exc: + raise SsoSessionError(str(exc)) from exc + row = SysSsoSession( + sid=sid, + session_secret_hash=secret_hash, + user_id=user_id, + subject_id=subject_id, + auth_version=auth_version, + auth_time=current, + last_seen_at=current, + idle_expires_at=idle, + absolute_expires_at=absolute, + acr=acr, + amr=list(amr), + remember_me=int(remember_me), + ip_address=ip_address, + user_agent_hash=user_agent_hash, + status=_SESSION_STATUS_ACTIVE, + ) + row = await SsoSessionDao.create(db, row) + snapshot = cls._snapshot(row) + + async def cache_after_commit() -> None: + """ + 在事务提交后写入新建 Session 的缓存 + + :return: None + """ + + try: + await cls._cache_row(redis, snapshot, current) + except Exception: + pass + + await coordinator.register(cache_after_commit) + + return f'{_COOKIE_PREFIX}.{sid}.{secret}', row + + @classmethod + async def validate_logout_cookie(cls, db: AsyncSession, cookie: str) -> SysSsoSession: + """ + 校验退出确认时的浏览器会话归属 + + 允许用户显式退出自然过期的会话,并继续撤销其关联离线授权。 + + :param db: orm对象 + :param cookie: 浏览器提交的SSO Cookie + :return: 已校验归属的SSO会话 + :raises SsoSessionError: 会话不存在、已撤销或凭据不匹配 + """ + + try: + sid, secret = OidcUtil.parse_sso_cookie(cookie) + digest = OidcUtil.session_secret_digest(secret, OidcConfig.oidc_token_hash_pepper) + except ValueError as exc: + raise SsoSessionError(str(exc)) from exc + row = await SsoSessionDao.get_by_sid(db, sid, for_update=True) + if ( + row is None + or row.status not in {'active', 'expired'} + or not hmac.compare_digest(row.session_secret_hash, digest) + ): + raise SsoSessionError('会话 Cookie 无效') + return row + + @classmethod + async def validate( # noqa: PLR0915 + cls, + db: AsyncSession, + redis: Redis, + cookie: str, + *, + pepper: str | bytes = OidcConfig.oidc_token_hash_pepper, + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> SysSsoSession: + """ + 验证 Cookie、Session 有效期和用户身份安全版本 + + :param db: 异步数据库会话 + :param redis: Redis 热缓存客户端 + :param cookie: 客户端 Cookie 原文 + :param pepper: OIDC Token Pepper + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行缓存清理或失效事件 + :return: 已验证的数据库 Session + :raises SsoSessionError: Cookie 或任何安全状态校验失败 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + try: + sid, secret = OidcUtil.parse_sso_cookie(cookie) + digest = OidcUtil.session_secret_digest(secret, pepper) + except ValueError as exc: + raise SsoSessionError(str(exc)) from exc + try: + await redis.get(OidcRedisKey.sso_cookie(digest)) + except Exception: + pass + row = await SsoSessionDao.get_by_sid(db, sid) + if row is None: + + async def clear_missing_cookie() -> None: + """ + 在 Session 不存在时删除 Cookie 索引 + + :return: None + """ + + try: + await redis.delete(OidcRedisKey.sso_cookie(digest)) + except Exception: + pass + + await coordinator.register(clear_missing_cookie) + raise SsoSessionError('会话不存在') + if not hmac.compare_digest(row.session_secret_hash, digest): + + async def clear_presented_cookie() -> None: + """ + 在 Cookie Secret 不匹配时删除当前索引 + + :return: None + """ + + try: + await redis.delete(OidcRedisKey.sso_cookie(digest)) + except Exception: + pass + + await coordinator.register(clear_presented_cookie) + raise SsoSessionError('会话密钥无效') + absolute = TimezoneUtil.to_optional_utc(row.absolute_expires_at) + idle = TimezoneUtil.to_optional_utc(row.idle_expires_at) + if ( + row.status != _SESSION_STATUS_ACTIVE + or absolute is None + or idle is None + or idle <= current + or absolute <= current + ): + if row.status == _SESSION_STATUS_ACTIVE and absolute is not None and idle is not None: + await SsoSessionDao.expire(db, row.sid, now=current) + row.status = _SESSION_STATUS_EXPIRED + snapshot = cls._snapshot(row) + + async def clear_expired() -> None: + """ + 在 Session 到期后清理缓存 + + :return: None + """ + + await cls._best_effort_cleanup(redis, snapshot, presented_digest=digest) + + await coordinator.register(clear_expired) + else: + await cls._invalidate(db, redis, row, digest, 'expired_or_inactive', current, coordinator) + raise SsoSessionError('会话已过期或已失效') + user = await IdentityUserDao.get_user(db, row.user_id) + subject = await IdentitySubjectDao.get_by_user_id(db, row.user_id) + if ( + user is None + or user.status != '0' + or user.del_flag != '0' + or subject is None + or subject.subject_id != row.subject_id + or subject.auth_version != row.auth_version + ): + await cls._invalidate(db, redis, row, digest, 'identity_security_changed', current, coordinator) + raise SsoSessionError('用户身份安全状态无效') + + snapshot = cls._snapshot(row) + + async def refresh_cache() -> None: + """ + 在校验成功后刷新 Session 缓存 + + :return: None + """ + + try: + await cls._cache_row(redis, snapshot, current) + except Exception: + pass + + await coordinator.register(refresh_cache) + + return row + + @classmethod + async def _invalidate( + cls, + db: AsyncSession, + redis: Redis, + row: SysSsoSession, + presented_digest: str, + reason: str, + now: datetime, + coordinator: AfterCommitCoordinator, + ) -> None: + """ + 将无效 Session 标记为撤销并登记提交后清理动作 + + :param db: 异步数据库会话 + :param redis: 异步 Redis 客户端 + :param row: SSO Session ORM 记录 + :param presented_digest: 当前 Cookie 摘要 + :param reason: 撤销原因 + :param now: 当前时间 + :param coordinator: 事务提交协调器 + :return: None + """ + + changed = False + if row.status == _SESSION_STATUS_ACTIVE: + changed = await SsoSessionDao.revoke(db, row.sid, reason=reason, now=now) + if changed: + await AuditService.record( + db, + OidcAuditEvent.SESSION_REVOKED, + 'success', + risk_level='high', + sid=row.sid, + user_id=row.user_id, + subject_id=row.subject_id, + detail={'reason': reason}, + ) + snapshot = cls._snapshot(row) + + async def cleanup() -> None: + """ + 在撤销或过期后清理 Session 缓存 + + :return: None + """ + + await cls._best_effort_cleanup( + redis, + snapshot, + presented_digest=presented_digest, + reason=reason if changed else None, + ) + + await coordinator.register(cleanup) + + @classmethod + async def touch( + cls, + db: AsyncSession, + redis: Redis, + cookie: str, + *, + pepper: str | bytes = OidcConfig.oidc_token_hash_pepper, + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> SysSsoSession: + """ + 刷新活动 Session 的最近访问时间和空闲过期时间 + + :param db: 异步数据库会话 + :param redis: Redis 热缓存客户端 + :param cookie: 当前 Cookie 原文 + :param pepper: OIDC Token Pepper + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行缓存更新 + :return: 更新后的数据库 Session + :raises SsoSessionError: Session 状态或事务操作不符合要求 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + cls._validate_config_ttls() + row = await cls.validate(db, redis, cookie, pepper=pepper, now=current, coordinator=coordinator) + absolute = TimezoneUtil.to_optional_utc(row.absolute_expires_at) + if absolute is None: + raise SsoSessionError('会话缺少绝对到期时间') + idle = min(current + timedelta(seconds=OidcConfig.oidc_sso_idle_seconds), absolute) + if not await SsoSessionDao.touch(db, row.sid, idle, now=current): + raise SsoSessionError('会话续期被拒绝') + refreshed = await SsoSessionDao.get_by_sid(db, row.sid) + if refreshed is None: + raise SsoSessionError('续期后的会话记录不存在') + snapshot = cls._snapshot(refreshed) + + async def refresh_cache() -> None: + """ + 在校验成功后刷新 Session 缓存 + + :return: None + """ + + try: + await cls._cache_row(redis, snapshot, current) + except Exception: + pass + + await coordinator.register(refresh_cache) + + return refreshed + + @classmethod + async def rotate_cookie( + cls, + db: AsyncSession, + redis: Redis, + cookie: str, + *, + pepper: str | bytes = OidcConfig.oidc_token_hash_pepper, + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> str: + """ + 轮换 Session Cookie Secret 并更新缓存索引 + + :param db: 异步数据库会话 + :param redis: Redis 热缓存客户端 + :param cookie: 旧 Cookie 原文 + :param pepper: OIDC Token Pepper + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行缓存更新 + :return: 新 Cookie 原文 + :raises SsoSessionError: Session 状态或事务操作不符合要求 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + row = await cls.validate(db, redis, cookie, pepper=pepper, now=current, coordinator=coordinator) + try: + sid, old_secret = OidcUtil.parse_sso_cookie(cookie) + new_secret = secrets.token_urlsafe(_COOKIE_SECRET_BYTES) + old_digest = OidcUtil.session_secret_digest(old_secret, pepper) + new_digest = OidcUtil.session_secret_digest(new_secret, pepper) + except ValueError as exc: + raise SsoSessionError(str(exc)) from exc + if not await SsoSessionDao.rotate_secret(db, sid, old_digest, new_digest, now=current): + raise SsoSessionError('会话 Cookie 轮换被拒绝') + row.session_secret_hash = new_digest + snapshot = cls._snapshot(row) + + async def refresh_rotated_cache() -> None: + """ + 在 Cookie 轮换后更新新旧缓存索引 + + :return: None + """ + + try: + await redis.delete(OidcRedisKey.sso_cookie(old_digest)) + await cls._cache_row(redis, snapshot, current) + except Exception: + pass + + await coordinator.register(refresh_rotated_cache) + + return f'{_COOKIE_PREFIX}.{sid}.{new_secret}' + + @classmethod + async def revoke( + cls, + db: AsyncSession, + redis: Redis, + sid: str, + *, + reason: str = 'logout', + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> bool: + """ + 撤销指定 Session 并清理其 Cookie 凭据 + + :param db: 异步数据库会话 + :param redis: Redis 热缓存客户端 + :param sid: 要撤销的 Session 标识 + :param reason: 撤销原因 + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行清理闭包 + :return: 数据库状态实际改变时返回 True + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + row = await SsoSessionDao.get_by_sid(db, sid, for_update=True) + if row is None: + return False + changed = await SsoSessionDao.revoke(db, sid, reason=reason, now=current) + if changed: + await AuditService.record( + db, + OidcAuditEvent.SESSION_REVOKED, + 'success', + risk_level='high', + sid=row.sid, + user_id=row.user_id, + subject_id=row.subject_id, + detail={'reason': reason}, + ) + snapshot = cls._snapshot(row) + + async def cleanup() -> None: + """ + 在撤销或过期后清理 Session 缓存 + + :return: None + """ + + await cls._best_effort_cleanup(redis, snapshot, reason=reason if changed else None) + + await coordinator.register(cleanup) + + return changed + + @classmethod + async def revoke_user( + cls, + db: AsyncSession, + redis: Redis, + user_id: int, + *, + reason: str = 'logout_all', + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> int: + """ + 撤销指定用户全部在线及自然过期的 Session,终止其离线访问 + + :param db: 异步数据库会话 + :param redis: Redis 热缓存客户端 + :param user_id: 本地用户 ID + :param reason: 撤销原因 + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行清理闭包 + :return: 数据库更新的 Session 数量 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + rows = list(await SsoSessionDao.list_for_user(db, user_id, for_update=True, revocable_only=True)) + changed_snapshots: list[SsoSessionSnapshot] = [] + for row in rows: + if await SsoSessionDao.revoke(db, row.sid, reason=reason, now=current): + changed_snapshots.append(cls._snapshot(row)) + await AuditService.record( + db, + OidcAuditEvent.SESSION_REVOKED, + 'success', + risk_level='high', + sid=row.sid, + user_id=row.user_id, + subject_id=row.subject_id, + detail={'reason': reason}, + ) + + async def cleanup() -> None: + """ + 在撤销或过期后清理 Session 缓存 + + :return: None + """ + + for snapshot in changed_snapshots: + await cls._best_effort_cleanup(redis, snapshot, reason=reason) + + await coordinator.register(cleanup) + + return len(changed_snapshots) + + @classmethod + async def expire_due( + cls, + db: AsyncSession, + redis: Redis, + *, + now: datetime | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> int: + """ + 批量标记已到期的 Session 并登记缓存清理 + + :param db: 异步数据库会话 + :param redis: Redis 热缓存客户端 + :param now: 可注入当前项目时间 + :param coordinator: 必须由调用方在数据库提交成功后执行缓存清理 + :return: 数据库更新数量 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + coordinator = cls._require_coordinator(coordinator) + rows = list(await SsoSessionDao.list_due(db, now=current, for_update=True)) + changed_snapshots: list[SsoSessionSnapshot] = [] + for row in rows: + if await SsoSessionDao.expire(db, row.sid, now=current): + changed_snapshots.append(cls._snapshot(row)) # noqa: PERF401 + + async def cleanup() -> None: + """ + 在撤销或过期后清理 Session 缓存 + + :return: None + """ + + for snapshot in changed_snapshots: + await cls._best_effort_cleanup(redis, snapshot) + + await coordinator.register(cleanup) + + return len(changed_snapshots) + + @classmethod + def cookie_max_age(cls, session: SysSsoSession, *, now: datetime | None = None) -> int | None: + """ + 根据 Session 的绝对过期时间计算保持登录 Cookie 的剩余寿命 + + :param session: 当前 SSO Session ORM + :param now: 可注入当前项目时间 + :return: 保持登录 Cookie 的剩余秒数,普通登录返回 None + :raises SsoSessionError: Session 缺少绝对过期时间或已过期 + """ + + if not session.remember_me: + return None + absolute = TimezoneUtil.to_optional_utc(session.absolute_expires_at) + if absolute is None: + raise SsoSessionError('会话缺少绝对到期时间') + return cls._remaining_ttl(absolute, TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now()) + + @staticmethod + def cookie_parameters(*, max_age: int | None = None) -> dict[str, str | bool | int]: + """ + 根据配置构造 SSO Cookie 的 HTTP 属性 + + :param max_age: 可选 Cookie 最大寿命秒数 + :return: 由 Controller 传给 HTTP Response 的 Cookie 参数 + :raises SsoSessionError: Logout Session 校验失败 + """ + + if not OidcConfig.oidc_sso_cookie_name.startswith('__Host-') or not OidcConfig.oidc_sso_cookie_secure: + raise SsoSessionError('使用 __Host- 前缀的单点登录 Cookie 必须启用 Secure') + if OidcConfig.oidc_sso_cookie_domain or OidcConfig.oidc_sso_cookie_samesite != 'lax': + raise SsoSessionError('单点登录 Cookie 的安全配置无效') + if max_age is not None and (isinstance(max_age, bool) or not isinstance(max_age, int) or max_age <= 0): + raise SsoSessionError('单点登录 Cookie 的有效期无效') + parameters: dict[str, str | bool | int] = { + 'key': OidcConfig.oidc_sso_cookie_name, + 'path': '/', + 'secure': True, + 'httponly': True, + 'samesite': 'lax', + } + if max_age is not None: + parameters['max_age'] = max_age + return parameters + + +LOGGER = logging.getLogger(__name__) +_RS256 = 'RS256' +_ID_TOKEN_TYPE = 'JWT' +_REMOTE_HEADERS = frozenset({'jku', 'x5u', 'jwk', 'x5c', 'crit'}) +_VERIFYING_KEY_STATUSES = ('active', 'retiring') +_URI_BACKCHANNEL = 'backchannel_logout' +_URI_POST_LOGOUT = 'post_logout' +_MAX_STATE_LENGTH = 2048 +_LOGOUT_TOKEN_TTL_SECONDS = 120 +_BACKCHANNEL_TIMEOUT_SECONDS = 5.0 +_MAX_BACKCHANNEL_TIMEOUT_SECONDS = 30.0 +_BACKCHANNEL_MAX_ATTEMPTS = 3 +_BACKCHANNEL_RETRY_DELAY_SECONDS = 0.05 +_BACKCHANNEL_MAX_CONCURRENCY = 8 +_JWT_DOT_COUNT = 2 +_BACKCHANNEL_MAX_QUEUE_ATTEMPTS = 5 +_BACKCHANNEL_MAX_QUEUE_ITEMS = 32 +_BACKCHANNEL_QUEUE_MAX_LENGTH = 1000 + + +@dataclass(frozen=True, slots=True) +class LogoutResult: + """ + 记录 RP-Initiated Logout 的重定向结果和错误信息 + """ + + redirect_uri: str | None + state: str | None + session_revoked: bool + + @property + def is_local(self) -> bool: + """ + 判断退出回调是否指向本地地址 + + :return: LogoutResult 未提供重定向 URI 时是否在本地完成退出 + """ + + return self.redirect_uri is None + + +class LogoutServiceError(ValueError): + """ + RP-Initiated Logout 参数或回调校验错误类型 + """ + + +BackchannelAuditWriter = Callable[..., Awaitable[None]] +BackchannelRetryQueue = Callable[..., Awaitable[None]] + + +class LogoutService: + """ + OIDC Logout 模块服务层 + """ + + @classmethod + async def execute_logout( # noqa: PLR0913 + cls, + db: AsyncSession, + redis: Redis, + *, + id_token_hint: str | None = None, + cookie: str | None = None, + post_logout_redirect_uri: str | None = None, + state: str | None = None, + confirmed: bool = False, + now: datetime | None = None, + signing_key: RSAPrivateKey | None = None, + signing_kid: str | None = None, + notifier: BackchannelNotifier | None = None, + audit_writer: BackchannelAuditWriter | None = None, + retry_queue: BackchannelRetryQueue | None = None, + ) -> LogoutResult: + """ + 编排 RP-Initiated Logout 的校验、撤销和响应 + + :param db: 异步数据库会话 + :param redis: SSO 热缓存客户端 + :param id_token_hint: OIDC ID Token Hint + :param cookie: 当前 SSO Cookie + :param post_logout_redirect_uri: 请求的退出后重定向 URI + :param state: 仅在重定向 URI 验证成功后返回的状态值 + :param now: 可注入的 项目当前时间 + :param signing_key: 测试或受控部署注入的 RSA 私钥 + :param signing_kid: 注入私钥对应的 kid + :param notifier: 提交后的 Back-Channel 通知回调 + :param audit_writer: Back-Channel 失败审计回调 + :param retry_queue: Back-Channel 重试队列回调 + :param confirmed: 是否已通过当前浏览器的一次性退出确认 + :return: 不包含敏感材料的退出结果 + :raises Exception: 核心流程或数据库提交失败时原样抛出异常 + """ + + coordinator = AfterCommitCoordinator() + try: + result = await cls._logout( + db, + redis, + id_token_hint=id_token_hint, + cookie=cookie, + post_logout_redirect_uri=post_logout_redirect_uri, + state=state, + confirmed=confirmed, + now=now, + coordinator=coordinator, + signing_key=signing_key, + signing_kid=signing_kid, + notifier=notifier, + audit_writer=audit_writer, + retry_queue=retry_queue, + ) + await coordinator.commit(db) + return result + except Exception: + try: + await coordinator.rollback(db) + except Exception: + pass + raise + + @classmethod + async def _logout( # noqa: PLR0912, PLR0913, PLR0915 + cls, + db: AsyncSession, + redis: Redis, + *, + id_token_hint: str | None = None, + cookie: str | None = None, + post_logout_redirect_uri: str | None = None, + state: str | None = None, + confirmed: bool = False, + now: datetime | None = None, + coordinator: AfterCommitCoordinator, + signing_key: RSAPrivateKey | None = None, + signing_kid: str | None = None, + notifier: BackchannelNotifier | None = None, + audit_writer: BackchannelAuditWriter | None = None, + retry_queue: BackchannelRetryQueue | None = None, + ) -> LogoutResult: + """ + 在事务中处理 Logout Hint、Session 撤销和通知登记 + + :param db: 异步数据库会话;事务由公共编排入口提交或回滚 + :param redis: SSO 热缓存客户端,仅传递给 Session 服务 + :param id_token_hint: OIDC ID Token Hint,不会写入日志或响应 + :param cookie: 当前 ``__Host-`` SSO Cookie + :param post_logout_redirect_uri: 待精确匹配的注册退出 URI + :param state: 仅在 URI 已验证后附加的原始状态值 + :param now: 可注入的当前项目时间 + :param coordinator: 提交后副作用协调器 + :param signing_key: 测试或受控部署注入的 RSA 私钥 + :param signing_kid: 注入私钥对应的 kid + :param notifier: 提交后的 Back-Channel 通知回调 + :param audit_writer: 审计写入回调 + :param retry_queue: 重试队列回调 + :param confirmed: 是否已通过当前浏览器的一次性退出确认 + :return: 不包含敏感材料的退出结果 + :raises LogoutServiceError: 请求违反安全边界 + """ + + if not OidcConfig.oidc_enabled: + raise LogoutServiceError('统一认证中心未启用') + if not confirmed: + raise LogoutServiceError('请在浏览器中确认退出操作') + commit_coordinator = coordinator + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + if state is not None and (not isinstance(state, str) or len(state) > _MAX_STATE_LENGTH): + state = None + session = None + client = None + validated_redirect: str | None = None + hint_valid = False + if id_token_hint: + try: + claims, client = await cls._validate_id_token_hint(db, id_token_hint, current) + hint_valid = True + except (JwtProfileError, KeyServiceError, LogoutServiceError, TypeError, ValueError, jwt.PyJWTError): + client = None + if hint_valid: + sid = claims['sid'] + session = await SsoSessionDao.get_by_sid(db, sid, for_update=True) + if hint_valid and ( + session is None or session.status not in {'active', 'expired'} or session.subject_id != claims['sub'] + ): + session = None + hint_valid = False + if hint_valid and post_logout_redirect_uri and OidcUtil.is_safe_post_logout_uri(post_logout_redirect_uri): + registered = await OAuthClientDao.find_exact_uri( + db, client.client_pk, _URI_POST_LOGOUT, post_logout_redirect_uri + ) + if registered is not None and registered.status == '0': + validated_redirect = post_logout_redirect_uri + else: + state = None + elif hint_valid and post_logout_redirect_uri: + state = None + browser_session = None + if cookie: + try: + browser_session = await SsoSessionService.validate_logout_cookie(db, cookie) + except (SsoSessionError, TypeError, ValueError): + browser_session = None + if browser_session is not None: + if session is None or session.sid != browser_session.sid: + # 其他账号的退出提示不得用于终止该账号会话 + validated_redirect = None + client = None + session = browser_session + if session is None: + return LogoutResult(None, None, False) + + refresh_rows = await cls._lock_session_refresh_tokens(db, session.sid) + client_ids = {row.client_pk for row in refresh_rows} + if client is not None: + client_ids.add(client.client_pk) + for participant_id in await SsoSessionDao.client_ids_for_sid(db, session.sid): + participant = await OAuthClientDao.get_by_client_id(db, participant_id, active_only=True) + if participant is not None: + client_ids.add(participant.client_pk) + await cls._revoke_session_state(db, redis, session.sid, refresh_rows, current, commit_coordinator) + await cls._register_backchannel( + db, + redis, + session.sid, + client_ids, + now=current, + coordinator=commit_coordinator, + signing_key=signing_key, + signing_kid=signing_kid, + notifier=notifier, + audit_writer=audit_writer, + retry_queue=retry_queue, + ) + + return LogoutResult(validated_redirect, state if validated_redirect else None, True) + + @classmethod + async def _validate_id_token_hint(cls, db: AsyncSession, token: str, now: datetime) -> tuple[dict[str, Any], Any]: + """ + 验证 ID Token Hint 的签名、发行者和 Session 绑定 + + :param db: 异步数据库会话 + :param token: 令牌值 + :param now: 当前时间 + :return: 包含协议字段的字典 + :raises LogoutServiceError: Logout 参数或回调地址不符合要求 + """ + + if not isinstance(token, str) or not token or token.count('.') != _JWT_DOT_COUNT: + raise LogoutServiceError('退出请求中的身份令牌提示 id_token_hint 无效') + try: + header = jwt.get_unverified_header(token) + if ( + header.get('alg') != _RS256 + or header.get('typ') != _ID_TOKEN_TYPE + or _REMOTE_HEADERS.intersection(header) + or not isinstance(header.get('kid'), str) + or not header['kid'] + ): + raise LogoutServiceError('身份令牌提示 id_token_hint 的头部无效') + unverified = jwt.decode(token, options={'verify_signature': False, 'verify_aud': False}) + audience = unverified.get('aud') + if not isinstance(audience, str) or not audience: + raise LogoutServiceError('身份令牌的受众无效') + except (jwt.PyJWTError, TypeError, ValueError) as exc: + raise LogoutServiceError('退出请求中的身份令牌提示 id_token_hint 无效') from exc + client = await OAuthClientDao.get_by_client_id(db, audience, active_only=True) + if client is None or client.status != '0': + raise LogoutServiceError('身份令牌的受众不是有效客户端') + key_record = await OidcKeyDao.get_verifying(db, header['kid'], now) + if key_record is None or key_record.status not in _VERIFYING_KEY_STATUSES: + raise LogoutServiceError('签名密钥不存在或不可用') + public_jwk = OidcUtil.normalize_public_jwk(key_record) + verification_key = jwt.algorithms.RSAAlgorithm.from_jwk(json.dumps(public_jwk, separators=(',', ':'))) + claims = decode_id_token_hint( + token, + verification_key=verification_key, + issuer=OidcConfig.oidc_issuer.rstrip('/'), + audience=client.client_id, + clock_skew=OidcConfig.oidc_allowed_clock_skew_seconds, + ) + + return claims, client + + @classmethod + async def _lock_session_refresh_tokens(cls, db: AsyncSession, sid: str) -> list[Any]: + """ + 锁定指定 Session 关联的 Refresh Token 记录 + + :param db: 异步数据库会话 + :param sid: Session 标识 + :return: 规范化后的列表 + """ + + return list(await OAuthTokenDao.list_for_sid_for_update(db, sid)) + + @classmethod + async def _revoke_session_state( + cls, + db: AsyncSession, + redis: Redis, + sid: str, + refresh_rows: list[Any], + now: datetime, + coordinator: AfterCommitCoordinator, + ) -> None: + """ + 标记 Session 及其 Refresh Token 为撤销并登记缓存清理 + + :param db: 异步数据库会话 + :param redis: 异步 Redis 客户端 + :param sid: Session 标识 + :param refresh_rows: Refresh 记录集合 + :param now: 当前时间 + :param coordinator: 事务提交协调器 + :return: None + """ + + await SsoSessionService.revoke(db, redis, sid, reason='rp_initiated_logout', now=now, coordinator=coordinator) + for family_id in {row.family_id for row in refresh_rows}: + await OAuthTokenDao.revoke_family(db, family_id, reason='rp_initiated_logout') + + @classmethod + async def _register_backchannel( # noqa: PLR0913 + cls, + db: AsyncSession, + redis: Redis, + sid: str, + client_pks: set[int], + *, + now: datetime, + coordinator: AfterCommitCoordinator, + signing_key: RSAPrivateKey | None, + signing_kid: str | None, + notifier: BackchannelNotifier | None, + audit_writer: BackchannelAuditWriter | None, + retry_queue: BackchannelRetryQueue | None, + ) -> None: + """ + 为已撤销 Session 登记各 Client 的 Back-Channel 通知 + + :param db: 异步数据库会话 + :param redis: 异步 Redis 客户端 + :param sid: Session 标识 + :param client_pks: 客户端主键集合 + :param now: 当前时间 + :param coordinator: 事务提交协调器 + :param signing_key: 签名密钥 + :param signing_kid: 签名密钥标识 + :param notifier: 通知器 + :param audit_writer: 审计写入回调 + :param retry_queue: 重试队列回调 + :return: None + """ + + effective_retry_queue = retry_queue or cls._redis_retry_queue(redis) + effective_audit_writer = audit_writer or cls._audit_writer(db) + callbacks: list[Callable[[], Awaitable[None]]] = [] + for client_pk in client_pks: + client = await OAuthClientDao.get_by_pk(db, client_pk, active_only=True) + if client is None: + continue + uris = await OAuthClientDao.list_uris(db, client_pk, uri_type=_URI_BACKCHANNEL, active_only=True) + session = await SsoSessionDao.get_by_sid(db, sid) + subject_id = getattr(session, 'subject_id', None) + include_sid = bool(getattr(client, 'backchannel_logout_session_required', 1)) + event_jti = str(uuid4()) + try: + token = await cls._make_logout_token( + client.client_id, + sid, + subject_id=subject_id, + include_sid=include_sid, + event_jti=event_jti, + now=now, + signing_key=signing_key, + signing_kid=signing_kid, + db=db, + ) + except (KeyServiceError, JwtProfileError, TypeError, ValueError): + token = None + timeout_seconds = cls._backchannel_timeout() + for uri_row in uris: + uri = uri_row.uri + safe = await cls._safe_backchannel_uri(uri) + if not safe: + + async def permanent_failure(uri_value: str = uri, client_value: str = client.client_id) -> None: + """ + 记录 Back-Channel 通知的永久失败 + + :param uri_value: URI 值 + :param client_value: 客户端值 + :return: None + """ + + await cls._write_backchannel_audit( + effective_audit_writer, + OidcAuditEvent.BACKCHANNEL_LOGOUT_FAILED, + uri_value, + client_value, + sid, + ) + + await coordinator.register(permanent_failure) + continue + if token is None: + + async def signing_failure( + uri_value: str = uri, + client_value: str = client.client_id, + event_jti_value: str = event_jti, + subject_id_value: str | None = subject_id, + include_sid_value: bool = include_sid, + ) -> None: + """ + 记录 Logout Token 签名失败 + + :param uri_value: URI 值 + :param client_value: 客户端值 + :param event_jti_value: 事件 JTI + :param subject_id_value: Subject 标识 + :param include_sid_value: 是否包含 Session 标识 + :return: None + """ + + await cls._write_backchannel_audit( + effective_audit_writer, + OidcAuditEvent.BACKCHANNEL_LOGOUT_FAILED, + uri_value, + client_value, + sid, + ) + await cls._enqueue_backchannel_retry( + effective_retry_queue, + uri_value, + client_value, + sid, + event_jti=event_jti_value, + subject_id=subject_id_value, + include_sid=include_sid_value, + ) + + await coordinator.register(signing_failure) + continue + callback = cls._notification_callback( + uri, + token, + notifier, + timeout_seconds=timeout_seconds, + client_id=client.client_id, + sid=sid, + subject_id=subject_id, + include_sid=include_sid, + event_jti=event_jti, + audit_writer=effective_audit_writer, + retry_queue=effective_retry_queue, + ) + callbacks.append(callback) + if callbacks: + semaphore = asyncio.Semaphore(_BACKCHANNEL_MAX_CONCURRENCY) + + async def notify_all() -> None: + """ + 向已登记的 Client 发送退出通知 + + :return: None + """ + + async def run(callback: Callable[[], Awaitable[None]]) -> None: + """ + 执行一次通知回调并统一处理发送异常 + + :param callback: 提交后回调 + :return: None + """ + + async with semaphore: + await callback() + + await asyncio.gather(*(run(callback) for callback in callbacks)) + + await coordinator.register(notify_all) + + @classmethod + async def consume_backchannel_retry( # noqa: PLR0912, PLR0915 + cls, + db: AsyncSession, + redis: Redis, + *, + now: datetime | None = None, + notifier: BackchannelNotifier | None = None, + max_items: int = _BACKCHANNEL_MAX_QUEUE_ITEMS, + ) -> int: + """ + 读取并处理 Back-Channel Logout 重试队列 + + :param db: 异步数据库会话 + :param redis: Redis 队列客户端 + :param now: 可注入当前时间 + :param notifier: 测试或受控网络发送器 + :param max_items: 单轮最大任务数 + :return: 已取出的任务数 + :raises PermanentBackchannelError: Back-Channel 通知达到永久失败条件 + :raises ValueError: 输入值不符合约束 + """ + + limit = max(1, min(int(max_items), _BACKCHANNEL_MAX_QUEUE_ITEMS)) + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + processed = 0 + for _ in range(limit): + raw = await redis.lpop(OidcRedisKey.backchannel_retry_queue()) + if raw is None: + break + processed += 1 + item: dict[str, Any] | None = None + try: + try: + item = json.loads(raw) + except (TypeError, UnicodeDecodeError, json.JSONDecodeError) as exc: + raise PermanentBackchannelError('invalid retry payload', '后端退出通知重试载荷无效') from exc + if not isinstance(item, dict): + raise PermanentBackchannelError('invalid retry payload', '后端退出通知重试载荷无效') + uri = item['uri'] + client_id = item['client_id'] + sid = item['sid'] + event_jti = item['event_jti'] + subject_id = item.get('subject_id') + include_sid = bool(item.get('include_sid', True)) + attempt = int(item.get('attempt', 1)) + next_attempt = int(item.get('nextAttemptAt', 0)) + if next_attempt > int(current.timestamp()): + await redis.rpush(OidcRedisKey.backchannel_retry_queue(), json.dumps(item, separators=(',', ':'))) + break + if ( + not isinstance(uri, str) + or not isinstance(client_id, str) + or not isinstance(sid, str) + or not isinstance(event_jti, str) + or not event_jti + or attempt < 1 + ): + raise PermanentBackchannelError('invalid retry task', '后端退出通知重试任务无效') + client = await OAuthClientDao.get_by_client_id(db, client_id, active_only=True) + registered = ( + await OAuthClientDao.find_exact_uri(db, client.client_pk, _URI_BACKCHANNEL, uri) if client else None + ) + if ( + client is None + or registered is None + or registered.status != '0' + or not await cls._safe_backchannel_uri(uri) + ): + raise PermanentBackchannelError('stale retry registration', '后端退出通知的重试注册信息已失效') + token = await cls._make_logout_token( + client_id, + sid, + subject_id=subject_id, + include_sid=include_sid, + event_jti=event_jti, + now=current, + signing_key=None, + signing_kid=None, + db=db, + ) + if token is None: + raise ValueError('签名密钥不可用') + await cls._send_backchannel_once(uri, token, notifier, cls._backchannel_timeout()) + await cls._write_backchannel_audit( + cls._audit_writer(db), OidcAuditEvent.BACKCHANNEL_LOGOUT_SUCCEEDED, uri, client_id, sid + ) + except PermanentBackchannelError as exc: + dead_key = OidcRedisKey.backchannel_retry_queue() + ':dead' + dead_item = item if isinstance(item, dict) else {'error': 'invalid_retry_payload'} + await redis.rpush(dead_key, json.dumps(dead_item, separators=(',', ':'))) + await redis.ltrim(dead_key, -_BACKCHANNEL_QUEUE_MAX_LENGTH, -1) + await cls._write_backchannel_audit( + cls._audit_writer(db), + OidcAuditEvent.BACKCHANNEL_LOGOUT_FAILED, + str(dead_item.get('uri', '')), + str(dead_item.get('client_id', '')), + str(dead_item.get('sid', '')), + failure_code=exc.failure_code, + ) + except Exception: + if isinstance(item, dict): + attempt = int(item.get('attempt', 1)) if str(item.get('attempt', 1)).isdigit() else 1 + if attempt < _BACKCHANNEL_MAX_QUEUE_ATTEMPTS: + item['attempt'] = attempt + 1 + item['nextAttemptAt'] = int(current.timestamp()) + min(300, 2**attempt) + await redis.rpush( + OidcRedisKey.backchannel_retry_queue(), json.dumps(item, separators=(',', ':')) + ) + await redis.ltrim(OidcRedisKey.backchannel_retry_queue(), -_BACKCHANNEL_QUEUE_MAX_LENGTH, -1) + else: + dead_key = OidcRedisKey.backchannel_retry_queue() + ':dead' + await redis.rpush(dead_key, json.dumps(item, separators=(',', ':'))) + await redis.ltrim(dead_key, -_BACKCHANNEL_QUEUE_MAX_LENGTH, -1) + await cls._write_backchannel_audit( + cls._audit_writer(db), + OidcAuditEvent.BACKCHANNEL_LOGOUT_FAILED, + str(item.get('uri', '')), + str(item.get('client_id', '')), + str(item.get('sid', '')), + ) + return processed + + @staticmethod + async def _safe_backchannel_uri(uri: Any) -> bool: + """ + 验证 Back-Channel URI 使用允许的网络地址 + + :param uri: 回调 URI + :return: URI 是否通过 Back-Channel 网络地址安全检查 + """ + + return await is_safe_backchannel_uri(uri) + + @staticmethod + async def _public_dns_only(hostname: str, port: int) -> bool: + """ + 解析主机名并确认所有地址均为公网地址 + + :param hostname: 主机名 + :param port: 端口 + :return: 主机解析得到的所有地址是否均为公网地址 + """ + + return await public_dns_only(hostname, port) + + @staticmethod + async def _verified_addresses(hostname: str, port: int) -> set[str]: + """ + 获取通过公网地址校验的主机地址集合 + + :param hostname: 主机名 + :param port: 端口 + :return: 通过校验的地址集合 + """ + + return await public_dns_addresses(hostname, port) + + @classmethod + async def _make_logout_token( + cls, + client_id: str, + sid: str, + *, + subject_id: str | None = None, + include_sid: bool = True, + event_jti: str | None = None, + now: datetime, + signing_key: RSAPrivateKey | None, + signing_kid: str | None, + db: AsyncSession, + ) -> str | None: + """ + 为指定 Client 和 Session 签发 Back-Channel Logout Token + + :param client_id: 客户端标识 + :param sid: Session 标识 + :param subject_id: Subject 标识 + :param include_sid: 是否包含 Session 标识 + :param event_jti: 事件 JTI + :param now: 当前时间 + :param signing_key: 签名密钥 + :param signing_kid: 签名密钥标识 + :param db: 异步数据库会话 + :return: 已签名的 Back-Channel Logout JWT,签名材料不可用时返回 None + """ + + if signing_key is None or signing_kid is None: + record = await OidcKeyDao.get_active(db, alg=_RS256) + if record is None: + return None + signing_key = await KeyService.load_private_key_async(record, now=now) + signing_kid = record.kid + try: + return encode_logout_token( + { + 'iss': OidcConfig.oidc_issuer.rstrip('/'), + 'aud': client_id, + 'iat': int(now.timestamp()), + 'exp': int(now.timestamp()) + _LOGOUT_TOKEN_TTL_SECONDS, + 'events': {BACKCHANNEL_LOGOUT_EVENT: {}}, + 'jti': event_jti or str(uuid4()), + **({'sid': sid} if include_sid else {'sub': subject_id}), + }, + signing_key, + signing_kid, + ) + except (JwtProfileError, TypeError, ValueError): + return None + + @staticmethod + def _notification_callback( # noqa: PLR0913 + uri: str, + token: str, + notifier: BackchannelNotifier | None, + *, + timeout_seconds: float = _BACKCHANNEL_TIMEOUT_SECONDS, + client_id: str = '', + sid: str = '', + subject_id: str | None = None, + include_sid: bool = True, + event_jti: str | None = None, + audit_writer: BackchannelAuditWriter | None = None, + retry_queue: BackchannelRetryQueue | None = None, + ) -> Callable[[], Awaitable[None]]: + """ + 创建发送 Back-Channel Logout 通知并处理失败的回调 + + :param uri: 回调 URI + :param token: 令牌值 + :param notifier: 通知器 + :param timeout_seconds: Back-Channel 请求超时时间 + :param client_id: 客户端标识 + :param sid: Session 标识 + :param subject_id: Subject 标识 + :param include_sid: 是否包含 Session 标识 + :param event_jti: 事件 JTI + :param audit_writer: 审计写入回调 + :param retry_queue: 重试队列回调 + :return: 异步回调 + """ + + async def notify() -> None: + """ + 发送一次 Back-Channel Logout 通知 + + :return: None + """ + + last_error: Exception | None = None + for attempt in range(_BACKCHANNEL_MAX_ATTEMPTS): + try: + await LogoutService._send_backchannel_once(uri, token, notifier, timeout_seconds) + await LogoutService._write_backchannel_audit( + audit_writer, 'backchannel_logout_succeeded', uri, client_id, sid + ) + return + except PermanentBackchannelError as exc: # noqa: PERF203 + await LogoutService._write_backchannel_audit( + audit_writer, + OidcAuditEvent.BACKCHANNEL_LOGOUT_FAILED, + uri, + client_id, + sid, + failure_code=exc.failure_code, + ) + return + except Exception as exc: + last_error = exc + if attempt + 1 < _BACKCHANNEL_MAX_ATTEMPTS: + await asyncio.sleep(_BACKCHANNEL_RETRY_DELAY_SECONDS * (attempt + 1)) + await LogoutService._enqueue_backchannel_retry( + retry_queue, + uri, + client_id, + sid, + event_jti=event_jti, + subject_id=subject_id, + include_sid=include_sid, + ) + await LogoutService._write_backchannel_audit(audit_writer, 'backchannel_logout_failed', uri, client_id, sid) + LOGGER.warning( + '后端退出通知发送失败,事件=%s,客户端=%s,会话=%s,异常类型=%s', + 'backchannel_logout_failed', + client_id, + sid, + type(last_error).__name__ if last_error else '未知', + ) + + return notify + + @staticmethod + async def _send_backchannel_once( + uri: str, token: str, notifier: BackchannelNotifier | None, timeout_seconds: float + ) -> None: + """ + 向 Client 的 Back-Channel URI 发送 Logout Token + + :param uri: 已注册的 Back-Channel URI + :param token: 仅存在于本次内存请求中的 Logout Token + :param notifier: 可注入发送器 + :param timeout_seconds: 有界网络超时 + :return: None + :raises LogoutServiceError: Logout 参数或回调地址不符合要求 + """ + + addresses: set[str] | None = None + if notifier is None: + parsed = urlsplit(uri) + hostname = parsed.hostname or '' + addresses = await LogoutService._verified_addresses(hostname, parsed.port or 443) + if not addresses: + raise LogoutServiceError('后端退出通知目标不是公网地址') + await send_backchannel_once( + uri, + token, + notifier, + timeout_seconds, + addresses=addresses, + ) + + @staticmethod + def _audit_writer(db: AsyncSession) -> BackchannelAuditWriter: + """ + 创建绑定数据库会话的 Back-Channel 审计写入器 + + :param db: 异步数据库会话 + :return: 绑定当前数据库会话的 Back-Channel 审计写入器 + """ + + async def write( + event: str, + uri: str, + client_id: str, + sid: str, + *, + failure_code: str | None = None, + ) -> None: + """ + 写入一次 Back-Channel 审计事件 + + :param event: 事件类型 + :param uri: 回调 URI + :param client_id: 客户端标识 + :param sid: Session 标识 + :param failure_code: 失败代码 + :return: None + """ + + try: + await AuditService.record_independent( + db, + event, + 'success' if event == OidcAuditEvent.BACKCHANNEL_LOGOUT_SUCCEEDED else 'failure', + client_id=client_id, + sid=sid, + failure_code=failure_code, + detail={'uri': uri}, + ) + except Exception: + LOGGER.warning('后端退出通知审计写入失败,事件=%s,客户端=%s,会话=%s', event, client_id, sid) + + return write + + @staticmethod + async def _write_backchannel_audit( + audit_writer: BackchannelAuditWriter | None, + event: str, + uri: str, + client_id: str, + sid: str, + failure_code: str | None = None, + ) -> None: + """ + 记录一次 Back-Channel 通知结果 + + :param audit_writer: 审计写入回调 + :param event: 事件类型 + :param uri: 回调 URI + :param client_id: 客户端标识 + :param sid: Session 标识 + :param failure_code: 失败代码 + :return: None + """ + + if audit_writer is not None: + try: + await audit_writer(event, uri, client_id, sid, failure_code=failure_code) + except Exception: + LOGGER.warning('后端退出通知审计写入失败,事件=%s', event) + + @staticmethod + async def _enqueue_backchannel_retry( + retry_queue: BackchannelRetryQueue | None, + uri: str, + client_id: str, + sid: str, + *, + event_jti: str | None = None, + subject_id: str | None = None, + include_sid: bool = True, + ) -> None: + """ + 将失败的 Back-Channel 通知加入重试队列 + + :param retry_queue: 重试队列回调 + :param uri: 回调 URI + :param client_id: 客户端标识 + :param sid: Session 标识 + :param event_jti: 事件 JTI + :param subject_id: Subject 标识 + :param include_sid: 是否包含 Session 标识 + :return: None + """ + + if retry_queue is not None: + try: + await retry_queue( + uri, + client_id, + sid, + event_jti=event_jti, + subject_id=subject_id, + include_sid=include_sid, + ) + except Exception: + LOGGER.warning('后端退出通知重试任务入队失败') + + @staticmethod + def _redis_retry_queue(redis: Redis) -> BackchannelRetryQueue: + """ + 创建基于 Redis 的 Back-Channel 重试队列适配器 + + :param redis: 异步 Redis 客户端 + :return: 基于 Redis 的 Back-Channel 重试队列适配器 + """ + + async def enqueue( + uri: str, + client_id: str, + sid: str, + *, + event_jti: str | None = None, + subject_id: str | None = None, + include_sid: bool = True, + ) -> None: + """ + 加入一条 Back-Channel 重试任务 + + :param uri: 回调 URI + :param client_id: 客户端标识 + :param sid: Session 标识 + :param event_jti: 事件 JTI + :param subject_id: Subject 标识 + :param include_sid: 是否包含 Session 标识 + :return: None + """ + + payload = json.dumps( + { + 'uri': uri, + 'client_id': client_id, + 'sid': sid, + 'event_jti': event_jti or str(uuid4()), + 'subject_id': subject_id, + 'include_sid': include_sid, + 'attempt': 1, + 'nextAttemptAt': int(TimezoneUtil.utc_now().timestamp()), + }, + separators=(',', ':'), + ) + await redis.rpush(OidcRedisKey.backchannel_retry_queue(), payload) + await redis.ltrim(OidcRedisKey.backchannel_retry_queue(), -_BACKCHANNEL_QUEUE_MAX_LENGTH, -1) + + return enqueue + + @staticmethod + def _backchannel_timeout() -> float: + """ + 读取 Back-Channel 请求超时配置 + + :return: 超时时间(秒) + """ + + value = getattr(OidcConfig, 'oidc_backchannel_logout_timeout_seconds', _BACKCHANNEL_TIMEOUT_SECONDS) + if ( + isinstance(value, bool) + or not isinstance(value, (int, float)) + or value <= 0 + or value > _MAX_BACKCHANNEL_TIMEOUT_SECONDS + ): + return _BACKCHANNEL_TIMEOUT_SECONDS + return float(value) diff --git a/ruoyi-fastapi-backend/module_identity/service/token_protocol_service.py b/ruoyi-fastapi-backend/module_identity/service/token_protocol_service.py new file mode 100644 index 000000000..9588656b2 --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/token_protocol_service.py @@ -0,0 +1,967 @@ +import hmac +import math +from collections.abc import Awaitable, Callable +from datetime import datetime +from typing import Any + +from redis.asyncio import Redis +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.dao.identity_subject_dao import IdentitySubjectDao +from module_identity.dao.identity_user_dao import IdentityUserDao +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.dao.oauth_resource_dao import OAuthResourceDao +from module_identity.dao.oauth_token_dao import OAuthTokenDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.redis_keys import OidcRedisKey +from module_identity.security.jwt_profile import JwtProfileError, decode_access_token +from module_identity.security.opaque_token import OpaqueTokenError, parse_opaque_token, token_digest +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.audit_service import AuditService +from module_identity.service.identity_service import ClaimService +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.token_service import TokenService +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + + +class IntrospectionService: + """ + OAuth Token 内省模块服务层 + """ + + _USERINFO_SUFFIX = '/oauth2/userinfo' + _ACTIVE_CLIENT_STATUS = '0' + _ACTIVE_USER_STATUS = '0' + _ACTIVE_GRANT_STATUS = 'active' + _ACTIVE_SESSION_STATUS = 'active' + _ACTIVE_TOKEN_STATUS = 'active' + _REFRESH_PREFIX = 'rt1' + _STANDARD_CLAIMS = frozenset( + {'scope', 'client_id', 'token_type', 'exp', 'iat', 'nbf', 'sub', 'aud', 'iss', 'jti', 'sid'} + ) + + @classmethod + async def introspect( + cls, + db: AsyncSession, + redis: Redis, + token: str, + caller: OAuthClientPrincipal, + *, + verification_keys: Any = None, + verification_key: Any = None, + now: datetime | None = None, + token_type_hint: str | None = None, + pepper: str | bytes | None = None, + ) -> dict[str, Any]: + """ + 验证 Token 并生成 OAuth Introspection 响应 + + :param db: 异步数据库会话 + :param redis: 统一认证中心 Redis 客户端 + :param token: 待内省的 OAuth Access Token 或 Refresh Token + :param caller: 已完成认证的 OAuth Client 主体 + :param verification_keys: 按 kid 索引的 Access Token 公钥 + :param verification_key: 单个 Access Token 公钥 + :param now: 可注入的当前项目时间 + :param token_type_hint: 可选的 Token 类型提示 + :param pepper: 可选 Token HMAC Pepper 覆盖值 + :return: 有效 Token 的声明,或严格的 ``{'active': False}`` + """ + + try: + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + client = await cls._resolve_caller(db, caller) + if client is None or not isinstance(token, str) or not token: + return {'active': False} + # 本系统可按令牌格式识别类型,忽略提示以避免错误提示改变有效性。 + if token.startswith('rt1.'): + return await cls._introspect_refresh(db, token, client, current, pepper) + return await cls._introspect_access( + db, + redis, + token, + client, + current, + verification_keys, + verification_key, + ) + except (JwtProfileError, OpaqueTokenError, TypeError, ValueError, KeyError): + # 格式和令牌状态错误统一返回 inactive,不向调用方暴露校验细节 + return {'active': False} + + @classmethod + async def _introspect_transaction( + cls, + db: AsyncSession, + redis: Redis, + token: str, + caller: OAuthClientPrincipal, + *, + verification_keys: Any = None, + verification_key: Any = None, + now: datetime | None = None, + token_type_hint: str | None = None, + pepper: str | bytes | None = None, + ) -> dict[str, Any]: + """ + 在事务边界内执行 Token 内省 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: 认证中心 Redis 客户端 + :param token: 待内省的 OAuth Access Token 或 Refresh Token + :param caller: 已认证的 OAuth Client 主体 + :param verification_keys: Access Token 公钥集合 + :param verification_key: 当前 Token 对应的公钥 + :param now: 可注入的 项目当前时间 + :param token_type_hint: Token 类型提示 + :param pepper: Refresh Token HMAC Pepper + :return: OAuth Introspection 响应 + """ + + try: + result = await cls.introspect( + db, + redis, + token, + caller, + verification_keys=verification_keys, + verification_key=verification_key, + now=now, + token_type_hint=token_type_hint, + pepper=pepper, + ) + await db.commit() + return result + except Exception: + await db.rollback() + raise + + @classmethod + async def introspect_request( + cls, + db: AsyncSession, + redis: Redis, + token: str, + *, + authorization: str | None = None, + client_id: str | None = None, + client_secret: str | None = None, + verification_key: Any = None, + verification_key_loader: Callable[[AsyncSession, str], Awaitable[Any]] | None = None, + token_type_hint: str | None = None, + ) -> dict[str, Any]: + """ + 认证调用 Client 并处理内省请求 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: 认证中心 Redis 客户端 + :param token: 待内省的 OAuth Access Token 或 Refresh Token + :param authorization: RFC 7617 Basic Header + :param client_id: 表单 Client ID + :param client_secret: 表单 Client Secret + :param verification_key: 当前 Access Token 的本地公钥 + :param verification_key_loader: 验证密钥加载回调 + :param token_type_hint: Token 类型提示 + :return: OAuth Introspection 响应 + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + try: + client_row, principal = await TokenService.authenticate_client( + db, + authorization=authorization, + client_id=client_id, + client_secret=client_secret, + ) + if client_row.client_type != 'confidential': + raise OAuthProtocolException('invalid_client', 'Client authentication failed', 401) + if verification_key_loader is not None: + verification_key = await verification_key_loader(db, token) + except Exception: + await db.rollback() + raise + return await cls._introspect_transaction( + db, + redis, + token, + principal, + verification_key=verification_key, + token_type_hint=token_type_hint, + ) + + @classmethod + async def _resolve_caller(cls, db: AsyncSession, caller: OAuthClientPrincipal) -> Any: + """ + 确认内省调用方 Client 处于有效状态 + + :param db: 异步数据库会话 + :param caller: 已认证的 OAuth Client 主体 + :return: 已启用且认证方式匹配的 OAuth Client ORM 对象,不满足条件时返回 None + """ + + if not isinstance(caller, OAuthClientPrincipal) or caller.auth_method != 'client_secret_basic': + return None + client = await OAuthClientDao.get_by_client_id(db, caller.client_id, active_only=True) + if ( + client is None + or client.status != cls._ACTIVE_CLIENT_STATUS + or client.client_type != 'confidential' + or client.client_id != caller.client_id + or client.client_type != caller.client_type + or client.token_endpoint_auth_method != 'client_secret_basic' + ): + return None + return client + + @classmethod + async def _introspect_access( # noqa: PLR0912 + cls, + db: AsyncSession, + redis: Redis, + token: str, + caller: Any, + now: datetime, + verification_keys: Any, + verification_key: Any, + ) -> dict[str, Any]: + """ + 验证 Access Token 并生成内省声明 + + :param db: 异步数据库会话 + :param redis: 异步 Redis 客户端 + :param token: 待内省的 Access Token + :param caller: 已认证且拥有 Resource 内省权限的 OAuth Client ORM + :param now: 当前时间 + :param verification_keys: 验证密钥集合 + :param verification_key: 验证密钥 + :return: 有效 Access Token 的 Introspection 声明,或严格的 ``{'active': False}`` + """ + + claims = decode_access_token( + token, + verification_keys, + issuer=OidcConfig.oidc_issuer, + clock_skew=OidcConfig.oidc_allowed_clock_skew_seconds, + verification_key=verification_key, + ) + jti = claims.get('jti') + if not isinstance(jti, str) or await cls._redis_key_exists(redis, OidcRedisKey.revoked_jti(jti)): + return {'active': False} + audiences = OidcUtil.normalize_audiences(claims.get('aud')) + resources = cls._resource_audiences(audiences) + if not resources or not await cls._resources_owned_by_caller(db, audiences, caller.client_pk): + return {'active': False} + issuer_client_id = claims.get('client_id') + if not isinstance(issuer_client_id, str) or not issuer_client_id: + return {'active': False} + issuer_client = await OAuthClientDao.get_by_client_id(db, issuer_client_id, active_only=True) + if issuer_client is None or issuer_client.status != cls._ACTIVE_CLIENT_STATUS: + return {'active': False} + if not await cls._client_allows_access(db, issuer_client, claims, resources): + return {'active': False} + grant_type = claims.get('gty') + user_state = None + if grant_type == 'client_credentials': + if ( + issuer_client.client_type != 'confidential' + or 'client_credentials' not in OidcUtil.json_list(issuer_client.grant_types) + or claims.get('sub') != f'client:{issuer_client.client_id}' + ): + return {'active': False} + elif grant_type not in {'authorization_code', 'refresh_token'}: + return {'active': False} + else: + user_state = await cls._access_user_state(db, claims, issuer_client, resources, now) + if user_state is None: + return {'active': False} + result: dict[str, Any] = {'active': True, 'token_type': 'Bearer'} + for name in cls._STANDARD_CLAIMS: + if name in claims: + result[name] = claims[name] + result['token_type'] = 'Bearer' + result['gty'] = grant_type + if user_state is not None: + result['username'] = user_state.user_name + result['ver'] = claims['ver'] + return result + + @classmethod + async def _access_user_state( + cls, db: AsyncSession, claims: dict[str, Any], client: Any, resources: list[str], now: datetime + ) -> Any: + """ + 验证 Access Token 对应的用户、Session 和授权来源 + + :param db: 异步数据库会话 + :param claims: 令牌声明 + :param client: 签发该 Access Token 的 OAuth Client ORM + :param resources: Access Token 声明中的 Resource Audience 列表 + :param now: 当前时间 + :return: 当前身份对应的启用用户 ORM 对象,安全状态或授权来源不匹配时返回 None + """ + + subject_id = claims.get('sub') + version = claims.get('ver') + sid = claims.get('sid') + scope = claims.get('scope') + if ( + not isinstance(subject_id, str) + or not subject_id + or isinstance(version, bool) + or not isinstance(version, int) + or not isinstance(sid, str) + or not sid + or not isinstance(scope, str) + or not scope.strip() + ): + return None + subject = await IdentitySubjectDao.get_by_subject_id(db, subject_id) + if subject is None or subject.auth_version != version: + return None + user = await IdentityUserDao.get_user(db, subject.user_id) + if ( + user is None + or user.status != cls._ACTIVE_USER_STATUS + or user.del_flag != cls._ACTIVE_USER_STATUS + or not isinstance(user.user_name, str) + or not user.user_name + ): + return None + session = await SsoSessionDao.get_for_token( + db, + sid, + now=now, + allow_offline=claims.get('gty') == 'refresh_token' or 'offline_access' in scope.split(), + ) + if ( + session is None + or session.user_id != subject.user_id + or session.subject_id != subject_id + or session.auth_version != version + ): + return None + scopes = scope.split() + grant_id = claims.get('grant_id') + if not isinstance(grant_id, str) or not grant_id: + # 升级前未绑定具体授权的用户令牌须重新授权,避免撤销范围不明确。 + return None + if await OAuthAccessPolicyDao.is_blocked(db, subject.user_id, client.client_pk): + return None + grant = await OAuthGrantDao.get_by_grant_id(db, grant_id) + if not cls._grant_active(grant, subject.user_id, subject_id, client, scopes, resources, now): + return None + return user + + @classmethod + async def _client_allows_access( + cls, db: AsyncSession, client: Any, claims: dict[str, Any], resources: list[str] + ) -> bool: + """ + 检查当前 Client 策略是否仍允许令牌中的 Scope 和 Resource + + :param db: 异步数据库会话 + :param client: 签发该 Access Token 的 OAuth Client ORM + :param claims: 已验证签名的 Access Token 声明 + :param resources: 待检查的 Resource Audience 列表 + :return: 策略版本及 Scope、Resource 绑定均有效时为 True + """ + + scope = claims.get('scope') + if not isinstance(scope, str) or not scope.strip(): + return False + if 'client_policy_version' in claims: + version = claims['client_policy_version'] + if isinstance(version, bool) or not isinstance(version, int) or version != client.policy_version: + return False + try: + await TokenService._validate_client_scope_resource( + db, client, scope.split(), resources, machine_only=claims.get('gty') == 'client_credentials' + ) + except OAuthProtocolException: + return False + return True + + @classmethod + async def _introspect_refresh( + cls, + db: AsyncSession, + token: str, + caller: Any, + now: datetime, + pepper: str | bytes | None, + ) -> dict[str, Any]: + """ + 验证 Refresh Token 并生成内省声明 + + :param db: 异步数据库会话 + :param token: 待内省的 Refresh Token + :param caller: 已认证且拥有 Resource 内省权限的 OAuth Client ORM + :param now: 当前时间 + :param pepper: 摘要 Pepper + :return: 有效 Refresh Token 的 Introspection 声明,或严格的 ``{'active': False}`` + """ + + parsed = parse_opaque_token(token, cls._REFRESH_PREFIX) + secret_pepper = OidcConfig.oidc_token_hash_pepper if pepper is None else pepper + digest = token_digest(token, secret_pepper) + row = await OAuthTokenDao.get_by_token_id(db, parsed.token_id, for_update=False) + if row is None or not hmac.compare_digest(row.token_hash, digest) or row.status != cls._ACTIVE_TOKEN_STATUS: + return {'active': False} + issuer_client = await OAuthClientDao.get_by_pk(db, row.client_pk, active_only=True) + if issuer_client is None or issuer_client.status != cls._ACTIVE_CLIENT_STATUS: + return {'active': False} + if ( + TimezoneUtil.to_optional_utc(row.idle_expires_at) <= now + or TimezoneUtil.to_optional_utc(row.absolute_expires_at) <= now + ): + return {'active': False} + if not await cls._refresh_family_active(db, row.family_id): + return {'active': False} + user = await IdentityUserDao.get_user(db, row.user_id) + subject = await IdentitySubjectDao.get_by_user_id(db, row.user_id) + if ( + user is None + or user.status != cls._ACTIVE_USER_STATUS + or user.del_flag != cls._ACTIVE_USER_STATUS + or subject is None + or subject.subject_id != row.subject_id + or subject.auth_version != row.auth_version + ): + return {'active': False} + session = await SsoSessionDao.get_for_token(db, row.sid, now=now, allow_offline=True) + if ( + session is None + or session.subject_id != row.subject_id + or session.user_id != row.user_id + or session.auth_version != row.auth_version + ): + return {'active': False} + if await OAuthAccessPolicyDao.is_blocked(db, row.user_id, row.client_pk): + return {'active': False} + grant = await OAuthGrantDao.get_by_grant_id(db, row.grant_id) + if not cls._grant_active( + grant, + row.user_id, + row.subject_id, + issuer_client, + OidcUtil.json_list(row.scopes), + OidcUtil.json_list(row.resources), + now, + ): + return {'active': False} + resources = OidcUtil.normalize_audiences(row.resources) + if not resources or not await cls._resources_owned_by_caller(db, resources, caller.client_pk): + return {'active': False} + result: dict[str, Any] = { + 'active': True, + 'client_id': issuer_client.client_id, + 'token_type': 'refresh_token', + 'scope': ' '.join(str(item) for item in OidcUtil.json_list(row.scopes)), + 'sub': row.subject_id, + 'username': user.user_name, + 'aud': resources, + 'jti': parsed.token_id, + 'sid': row.sid, + 'iat': OidcUtil.numeric_date(row.issued_at), + 'exp': OidcUtil.numeric_date(row.absolute_expires_at), + } + + return result + + @classmethod + async def _refresh_family_active(cls, db: AsyncSession, family_id: str) -> bool: + """ + 检查 Refresh Token Family 是否仍处于活动状态 + + :param db: 异步数据库会话 + :param family_id: Refresh Family 标识 + :return: 仅在存在记录且没有全族终止状态时返回真 + """ + + return await OAuthTokenDao.family_is_active(db, family_id) + + @classmethod + def _grant_active( + cls, + grant: Any, + user_id: int, + subject_id: str, + client: Any, + scopes: list[Any], + resources: list[Any], + now: datetime, + ) -> bool: + """ + 检查 Grant 是否匹配当前身份、Client 和授权范围 + + :param grant: Grant 对象 + :param user_id: 用户标识 + :param subject_id: Subject 标识 + :param client: 签发该 Token 的 OAuth Client ORM + :param scopes: Token 授权的 Scope 列表 + :param resources: Token 授权的 Resource Audience 列表 + :param now: 当前时间 + :return: Grant 匹配用户、Subject、OAuth Client、Scope、Resource 且未过期时为 True + """ + + return bool( + grant is not None + and grant.user_id == user_id + and grant.subject_id == subject_id + and grant.client_pk == client.client_pk + and grant.status == cls._ACTIVE_GRANT_STATUS + and grant.client_policy_version == client.policy_version + and (grant.expires_at is None or TimezoneUtil.to_optional_utc(grant.expires_at) > now) + and set(scopes).issubset(set(OidcUtil.json_list(grant.granted_scopes))) + and set(resources).issubset(set(OidcUtil.json_list(grant.granted_resources))) + ) + + @classmethod + async def _resources_owned_by_caller(cls, db: AsyncSession, audiences: list[str], client_pk: int) -> bool: + """ + 检查调用方是否拥有资源的内省权限 + + :param db: 异步数据库会话 + :param audiences: Token 声明中的 Audience 列表 + :param client_pk: 已认证 OAuth Client 的主键 + :return: 所有业务 Resource Audience 均归属于该 OAuth Client 的内省权限时为 True + """ + + userinfo = f'{OidcConfig.oidc_issuer.rstrip("/")}{cls._USERINFO_SUFFIX}' + resource_audiences = [value for value in audiences if value != userinfo] + if not resource_audiences: + return False + rows = {row.audience: row for row in await OAuthResourceDao.active_by_audiences(db, resource_audiences)} + + return len(rows) == len(set(resource_audiences)) and all( + rows[value].introspection_client_pk == client_pk for value in resource_audiences + ) + + @classmethod + def _resource_audiences(cls, audiences: list[str]) -> list[str]: + """ + 移除 UserInfo audience 并保留业务资源 + + :param audiences: Token 声明中的 Audience 列表 + :return: 移除 UserInfo Audience 后的业务 Resource Audience 列表 + """ + + userinfo = f'{OidcConfig.oidc_issuer.rstrip("/")}{cls._USERINFO_SUFFIX}' + + return [value for value in audiences if value != userinfo] + + @staticmethod + async def _redis_key_exists(redis: Redis, key: str) -> bool: + """ + 检查 Redis 撤销键是否存在 + + :param redis: Redis 异步客户端 + :param key: 待检查的 Redis 撤销键 + :return: 键存在时返回 True + """ + + return bool(await redis.exists(key)) + + +class RevocationError(ValueError): + """ + 令牌撤销参数错误类型 + """ + + def __init__(self, error: str, description: str) -> None: + """ + 保存 OAuth 撤销错误代码和描述 + + :param error: OAuth 错误代码 + :param description: 可向上层映射的安全错误描述 + :return: None + """ + + self.message = OidcUtil.localized_oauth_message(error, description) + super().__init__(self.message) + self.error = error + self.description = OidcUtil.protocol_error_description(error, description) or '' + + +class RevocationService: + """ + OAuth Token 撤销模块服务层 + """ + + _REFRESH_PREFIX = 'rt1' + + @classmethod + async def revoke( + cls, + db: AsyncSession, + redis: Redis, + token: str, + caller: OAuthClientPrincipal, + *, + verification_keys: Any = None, + verification_key: Any = None, + now: datetime | None = None, + token_type_hint: str | None = None, + pepper: str | bytes | None = None, + coordinator: AfterCommitCoordinator | None = None, + ) -> None: + """ + 按 Token 类型撤销 Access Token 或 Refresh Token + + :param db: 异步数据库会话;事务由提交协调器或调用方管理 + :param redis: 统一认证中心 Redis 客户端 + :param token: 待撤销的 OAuth Access Token 或 Refresh Token;未知 Token 幂等成功 + :param caller: 已认证的 OAuth Client 主体 + :param verification_keys: 按 kid 索引的 Access Token 公钥 + :param verification_key: 单个 Access Token 公钥 + :param now: 可注入的当前项目时间 + :param token_type_hint: 可选 Token 类型提示 + :param pepper: 可选 Token HMAC Pepper 覆盖值 + :param coordinator: 提交后副作用协调器;Access 撤销必须提供 + :return: None + :raises RevocationError: Client 未认证或副作用边界不可用 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + client = await cls._resolve_caller(db, caller) + if client is None: + raise RevocationError('invalid_client', 'Client authentication failed') + if not isinstance(token, str) or not token: + return + # 类型提示只用于优化查询,未知或错误提示不得阻止实际撤销。 + if token.startswith('rt1.'): + await cls._revoke_refresh(db, token, client, current, pepper) + return + await cls._revoke_access( + db, + redis, + token, + client, + current, + verification_keys, + verification_key, + coordinator, + ) + + return + + @classmethod + async def _revoke_transaction( + cls, + db: AsyncSession, + redis: Redis, + token: str, + caller: OAuthClientPrincipal, + *, + verification_keys: Any = None, + verification_key: Any = None, + now: datetime | None = None, + token_type_hint: str | None = None, + pepper: str | bytes | None = None, + ) -> bool: + """ + 在事务边界内执行 Token 撤销 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: 认证中心 Redis 客户端 + :param token: 待撤销的完整 Token + :param caller: 已认证的 OAuth Client 主体 + :param verification_keys: Access Token 公钥集合 + :param verification_key: 当前 Token 对应的公钥 + :param now: 可注入的 项目当前时间 + :param token_type_hint: Token 类型提示 + :param pepper: Refresh Token HMAC Pepper + :return: 提交后副作用全部成功时为 True + :raises RevocationError: 令牌撤销失败且事务已回滚 + """ + + coordinator = AfterCommitCoordinator() + try: + await cls.revoke( + db, + redis, + token, + caller, + verification_keys=verification_keys, + verification_key=verification_key, + now=now, + token_type_hint=token_type_hint, + pepper=pepper, + coordinator=coordinator, + ) + await coordinator.commit(db) + except Exception: + await coordinator.rollback(db) + raise + return not coordinator.callback_errors + + @classmethod + async def revoke_request( + cls, + db: AsyncSession, + redis: Redis, + token: str, + *, + authorization: str | None = None, + client_id: str | None = None, + client_secret: str | None = None, + verification_key: Any = None, + verification_key_loader: Callable[[AsyncSession, str], Awaitable[Any]] | None = None, + token_type_hint: str | None = None, + ) -> bool: + """ + 认证调用 Client 并处理撤销请求 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: 认证中心 Redis 客户端 + :param token: 待撤销的完整 Token + :param authorization: RFC 7617 Basic Header + :param client_id: 表单 Client ID + :param client_secret: 表单 Client Secret + :param verification_key: 当前 Access Token 的本地公钥 + :param verification_key_loader: 验证密钥加载回调 + :param token_type_hint: Token 类型提示 + :return: 提交后副作用全部成功时为 True + """ + + try: + _, principal = await TokenService.authenticate_client( + db, + authorization=authorization, + client_id=client_id, + client_secret=client_secret, + ) + if verification_key_loader is not None: + verification_key = await verification_key_loader(db, token) + except Exception: + await db.rollback() + raise + return await cls._revoke_transaction( + db, + redis, + token, + principal, + verification_key=verification_key, + token_type_hint=token_type_hint, + ) + + @staticmethod + async def _resolve_caller(db: AsyncSession, caller: OAuthClientPrincipal) -> Any: + """ + 确认内省调用方 Client 处于有效状态 + + :param db: 异步数据库会话 + :param caller: 已认证的 OAuth Client 主体 + :return: 已通过认证和资源校验的 OAuth Client ORM 对象,不满足条件时返回 None + """ + + if not isinstance(caller, OAuthClientPrincipal): + return None + client = await OAuthClientDao.get_by_client_id(db, caller.client_id, active_only=True) + if ( + client is None + or client.status != '0' + or client.client_id != caller.client_id + or client.client_type != caller.client_type + or ( + client.client_type == 'public' + and (caller.auth_method != 'none' or client.token_endpoint_auth_method != 'none') + ) + or ( + client.client_type == 'confidential' + and ( + caller.auth_method != 'client_secret_basic' + or client.token_endpoint_auth_method != 'client_secret_basic' + ) + ) + ): + return None + if client.client_type not in {'public', 'confidential'}: + return None + return client + + @classmethod + async def _revoke_refresh( + cls, + db: AsyncSession, + token: str, + client: Any, + now: datetime, + pepper: str | bytes | None, + ) -> None: + """ + 撤销 Refresh Token 及其 Family + + :param db: 异步数据库会话 + :param token: 待撤销的 Refresh Token + :param client: 发起撤销请求的 OAuth Client ORM + :param now: 当前时间 + :param pepper: 摘要 Pepper + :return: None + """ + + try: + parsed = parse_opaque_token(token, cls._REFRESH_PREFIX) + secret_pepper = OidcConfig.oidc_token_hash_pepper if pepper is None else pepper + digest = token_digest(token, secret_pepper) + row = await OAuthTokenDao.get_by_token_id(db, parsed.token_id, for_update=True) + except (OpaqueTokenError, TypeError, ValueError): + return + if row is None or not hmac.compare_digest(row.token_hash, digest) or row.client_pk != client.client_pk: + return + if row.status in {'revoked', 'expired', 'reuse_detected'}: + return + await OAuthTokenDao.revoke_family(db, row.family_id, reason='client_revocation') + await AuditService.record( + db, + OidcAuditEvent.TOKEN_REVOKED, + 'success', + client_id=client.client_id, + subject_id=getattr(row, 'subject_id', None), + sid=getattr(row, 'sid', None), + grant_id=getattr(row, 'grant_id', None), + token_id=parsed.token_id, + ) + + return + + @classmethod + async def _revoke_access( + cls, + db: AsyncSession, + redis: Redis, + token: str, + client: Any, + now: datetime, + verification_keys: Any, + verification_key: Any, + coordinator: AfterCommitCoordinator | None, + ) -> None: + """ + 验证并撤销 Access Token + + :param db: 异步数据库会话 + :param redis: 异步 Redis 客户端 + :param token: 待撤销的 Access Token + :param client: 发起撤销请求的 OAuth Client ORM + :param now: 当前时间 + :param verification_keys: 验证密钥集合 + :param verification_key: 验证密钥 + :param coordinator: 事务提交协调器 + :return: None + :raises RevocationError: 令牌撤销操作失败 + """ + + try: + claims = decode_access_token( + token, + verification_keys, + issuer=OidcConfig.oidc_issuer, + clock_skew=OidcConfig.oidc_allowed_clock_skew_seconds, + verification_key=verification_key, + ) + except Exception: + return + if claims.get('client_id') != client.client_id: + return + jti = claims.get('jti') + exp = claims.get('exp') + if not isinstance(jti, str) or not jti or isinstance(exp, bool) or not isinstance(exp, (int, float)): + return + remaining = exp - now.timestamp() + OidcConfig.oidc_allowed_clock_skew_seconds + if not math.isfinite(remaining) or remaining <= 0: + return + if coordinator is None or not isinstance(coordinator, AfterCommitCoordinator): + raise RevocationError('server_error', 'A transaction coordinator is required') + ttl = max(1, math.ceil(remaining)) + key = OidcRedisKey.revoked_jti(jti) + + async def write_revocation() -> None: + """ + 执行 write_revocation 的校验和转换 + + :return: None + """ + + await redis.set(key, '1', ex=ttl) + await AuditService.record_independent( + db, + OidcAuditEvent.TOKEN_REVOKED, + 'success', + client_id=client.client_id, + token_id=jti, + ) + + await coordinator.register(write_revocation) + + return + + +class UserInfoService: + """ + OIDC UserInfo 模块服务层 + """ + + @classmethod + async def build(cls, db: AsyncSession, claims: dict[str, Any], redis: Redis | None) -> dict[str, Any]: + """ + 验证 Access Token 声明并构造 UserInfo 响应 + + :param db: 异步数据库会话 + :param claims: 令牌声明 + :param redis: 异步 Redis 客户端 + :return: 包含协议字段的字典 + :raises ValueError: 输入值不符合约束 + """ + + client = await OAuthClientDao.get_by_client_id(db, claims['client_id'], active_only=True) + if client is None or client.status != '0': + raise ValueError('客户端已停用') + resources = IntrospectionService._resource_audiences(OidcUtil.normalize_audiences(claims.get('aud'))) + if not await IntrospectionService._client_allows_access(db, client, claims, resources): + raise ValueError('客户端访问策略已失效') + user = await IntrospectionService._access_user_state(db, claims, client, resources, TimezoneUtil.utc_now()) + if user is None: + raise ValueError('授权已失效') + await cls._check_revocation(redis, claims['jti']) + + scopes = str(claims.get('scope', '')).split() + policy, allowed = await ClaimService.resolve_scope_policy(db, client.client_pk, scopes) + roles, department = await ClaimService.load_roles_and_department(db, user.user_id) + + return ClaimService.build_claims( + user, + scopes, + policy, + allowed, + subject_id=claims['sub'], + roles=roles, + department=department, + ) + + @staticmethod + async def _check_revocation(redis: Redis | None, jti: str) -> None: + """ + 检查 Access Token 是否已被撤销 + + :param redis: 异步 Redis 客户端 + :param jti: 令牌 JTI + :return: None + :raises ValueError: 输入值不符合约束 + """ + + if redis is None: + raise ValueError('令牌已撤销') + try: + revoked = bool(await redis.exists(OidcRedisKey.revoked_jti(jti))) + except (ConnectionError, TimeoutError, OSError): + raise ValueError('令牌状态暂不可用') from None + if revoked: + raise ValueError('令牌已撤销') diff --git a/ruoyi-fastapi-backend/module_identity/service/token_service.py b/ruoyi-fastapi-backend/module_identity/service/token_service.py new file mode 100644 index 000000000..409e9a17e --- /dev/null +++ b/ruoyi-fastapi-backend/module_identity/service/token_service.py @@ -0,0 +1,1409 @@ +import hmac +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timedelta +from typing import Any +from uuid import uuid4 + +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey +from pydantic import ValidationError +from redis.asyncio import Redis +from redis.exceptions import RedisError +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_admin.entity.do.user_do import SysUser +from module_identity.dao.identity_user_dao import IdentityUserDao +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.dao.oauth_token_dao import OAuthTokenDao +from module_identity.dao.oidc_key_dao import OidcKeyDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken, SysSsoSession +from module_identity.entity.do.oauth_resource_do import SysOAuthResource, SysOAuthScope +from module_identity.entity.vo.protocol_vo import TokenRequest +from module_identity.security.client_auth import ClientAuthenticationError, authenticate_client +from module_identity.security.jwt_profile import JwtProfileError, encode_access_token, encode_id_token +from module_identity.security.opaque_token import ( + OpaqueTokenError, + generate_refresh_token, + parse_opaque_token, + token_digest, +) +from module_identity.security.pkce import PkceError, verify_code_challenge +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import AuthorizationCodeReuseError, AuthorizationCodeService +from module_identity.service.identity_service import ClaimService, IdentitySubjectService +from module_identity.service.key_service import KeyService, KeyServiceError +from utils.oidc_util import OidcUtil +from utils.time_util import TimezoneUtil + + +@dataclass(frozen=True, slots=True) +class TokenResult: + """ + 记录 Token Endpoint 生成的各类令牌 + """ + + access_token: str + expires_in: int + refresh_token: str | None = None + scope: str | None = None + id_token: str | None = None + + @property + def token_type(self) -> str: + """ + 返回 OAuth Bearer 令牌类型 + + :return: 固定为 OAuth 2.0 Bearer token 类型的字符串 + """ + + return 'Bearer' + + def as_dict(self) -> dict[str, Any]: + """ + 将令牌结果转换为 Token Endpoint 响应字典 + + :return: 包含协议字段的字典 + """ + + result: dict[str, Any] = { + 'access_token': self.access_token, + 'token_type': 'Bearer', + 'expires_in': self.expires_in, + } + if self.refresh_token is not None: + result['refresh_token'] = self.refresh_token + if self.scope is not None: + result['scope'] = self.scope + if self.id_token is not None: + result['id_token'] = self.id_token + return result + + def __getitem__(self, key: str) -> Any: + """ + 按字段名读取 Token Endpoint 响应值 + + :param key: 要读取的 Token Endpoint 响应字段名 + :return: 对应字段的 Access Token、有效期或可选令牌值 + """ + + return self.as_dict()[key] + + +class RefreshTokenReuseDetected(OAuthProtocolException): + """ + Refresh Token 重放检测错误类型 + """ + + must_commit = True + + def __init__(self) -> None: + """ + 保存 OAuth 撤销错误代码和描述 + + :return: None + """ + + super().__init__('invalid_grant', 'The authorization grant is invalid or expired', 400) + + +class TokenService: + """ + OAuth Token 模块服务层 + """ + + _USERINFO_AUDIENCE_SUFFIX = '/oauth2/userinfo' + _MACHINE_SCOPE_TYPE = 'resource' + _ACTIVE_CLIENT_STATUSES = frozenset({'0', 0, 'active'}) + _ACTIVE_GRANT_STATUS = 'active' + _ACTIVE_SESSION_STATUS = 'active' + _MAX_RESOURCE_COUNT = 1 + _MIN_TOKEN_PEPPER_BYTES = 32 + + @classmethod + async def authenticate_client( + cls, + db: AsyncSession, + *, + authorization: str | None = None, + client_id: str | None = None, + client_secret: str | None = None, + ) -> tuple[SysOAuthClient, OAuthClientPrincipal]: + """ + 解析凭据并验证 OAuth Client 身份 + + :param db: 异步数据库会话 + :param authorization: RFC 7617 Basic Header;机密 Client 必须使用该方式 + :param client_id: 公共 Client 的表单 Client ID + :param client_secret: 仅用于拒绝公共 Client 的 body Secret + :return: 已启用 Client 行和不可伪造的 Client Principal + :raises OAuthProtocolException: Client 凭据无效时抛出 invalid_client + """ + + if authorization is not None and client_secret is not None: + cls._invalid_client() + lookup_id = client_id + if authorization is not None: + try: + lookup_id, _ = OidcUtil.parse_basic_credentials(authorization) + except (ClientAuthenticationError, TypeError, ValueError): + cls._invalid_client() + if not isinstance(lookup_id, str) or not lookup_id: + cls._invalid_client() + row = await OAuthClientDao.get_by_client_id(db, lookup_id, active_only=True) + if row is None: + cls._invalid_client() + secrets = await OAuthClientDao.list_secrets(db, row.client_pk, active_only=True) + matched_secret_hash: str | None = None + + def capture_secret_match(secret_hash: str) -> None: + """ + 保存认证过程中匹配的 Client Secret 哈希 + + :param secret_hash: 已匹配的持久化 Secret 哈希 + :return: None + """ + + nonlocal matched_secret_hash + matched_secret_hash = secret_hash + + try: + principal = authenticate_client( + row, + authorization, + client_id=client_id, + client_secret=client_secret, + secret_hashes=secrets, + secret_match_callback=capture_secret_match, + ) + except (ClientAuthenticationError, TypeError, ValueError): + cls._invalid_client() + if principal.client_type == 'confidential' and matched_secret_hash is not None: + matched_secret_id = next( + ( + getattr(secret_row, 'secret_id', None) + for secret_row in secrets + if getattr(secret_row, 'secret_hash', None) == matched_secret_hash + ), + None, + ) + if isinstance(matched_secret_id, str): + await OAuthClientDao.mark_secret_used(db, matched_secret_id) + return row, principal + + @classmethod + async def _handle_authorization_code_reuse( + cls, + db: AsyncSession, + redis: Redis, + code: str, + client_id: str, + pepper: str | bytes, + ) -> None: + """ + 撤销重用授权码关联的 Grant 并记录高风险审计 + + :param db: 异步数据库会话 + :param redis: Authorization Code Redis 客户端 + :param code: 被重复提交的 Authorization Code 明文 + :param client_id: 已认证 Client 公开标识 + :param pepper: Authorization Code 摘要 Pepper + :return: None + :raises OAuthProtocolException: 撤销或审计无法安全持久化时抛出 + """ + + try: + consumed_payload = await AuthorizationCodeService.consumed_payload(redis, code, pepper=pepper) + grant_id = consumed_payload.get('grantId') if consumed_payload else None + if isinstance(grant_id, str) and callable(getattr(db, 'execute', None)): + await OAuthGrantDao.revoke(db, grant_id, reason='authorization_code_reuse') + await AuditService.record_independent( + db, + OidcAuditEvent.AUTHORIZATION_CODE_REUSED, + 'failure', + risk_level='high', + client_id=client_id, + failure_code='authorization_code_reused', + ) + except Exception: + raise OAuthProtocolException('server_error', 'Token endpoint is unavailable', 500) from None + + @classmethod + async def authorization_code( + cls, + db: AsyncSession, + redis: Redis, + request: TokenRequest | Mapping[str, Any], + client: OAuthClientPrincipal, + *, + signing_key: RSAPrivateKey | None = None, + kid: str | None = None, + now: datetime | None = None, + token_pepper: str | bytes | None = None, + ) -> TokenResult: + """ + 消费授权码并签发用户令牌 + + :param db: 异步数据库会话,调用方负责提交或回滚事务 + :param redis: Authorization Code 所在 Redis 客户端 + :param request: Token Endpoint 请求或字段映射 + :param client: 已完成 Client Authentication 的不可伪造主体 + :param signing_key: 可注入的 RSA 私钥;省略时从 KeyService 读取 + :param kid: 注入私钥对应的签名 Key ID + :param now: 可注入 项目当前时间 + :param token_pepper: Refresh Token HMAC Pepper + :return: 标准 Token 结果 + :raises OAuthProtocolException: 任一授权绑定或安全状态无效时抛出统一错误 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + parsed = cls._request(request) + if parsed.get('grant_type') != 'authorization_code': + cls._invalid_request() + client_row = await cls._resolve_client(db, client) + await OAuthAccessPolicyDao.lock_client(db, client_row.client_pk) + if parsed.get('client_id') not in (None, client_row.client_id) or ( + client_row.client_type == 'public' and parsed.get('client_id') != client_row.client_id + ): + cls._invalid_client() + if 'authorization_code' not in OidcUtil.json_list(client_row.grant_types): + cls._invalid_grant() + code = parsed.get('code') + if not isinstance(code, str): + cls._invalid_grant() + try: + code_payload = await AuthorizationCodeService.consume(redis, code, pepper=cls._token_pepper(token_pepper)) + except AuthorizationCodeReuseError: + await cls._handle_authorization_code_reuse( + db, + redis, + code, + client_row.client_id, + cls._token_pepper(token_pepper), + ) + raise + except (OAuthProtocolException, OpaqueTokenError, RedisError) as exc: + cls._map_oauth_failure(exc, 'invalid_grant') + if code_payload.get('clientPk') != client_row.client_pk: + cls._invalid_grant() + if parsed.get('redirect_uri') != code_payload.get('redirectUri'): + cls._invalid_grant() + verifier = parsed.get('code_verifier') + try: + if not verify_code_challenge(verifier, code_payload['codeChallenge'], code_payload['codeChallengeMethod']): + cls._invalid_grant() + except (PkceError, TypeError, ValueError): + cls._invalid_grant() + user, subject = await cls._require_user_identity( + db, + int(code_payload['userId']), + str(code_payload['subjectId']), + int(code_payload['authVersion']), + ) + session = await cls._require_session( + db, + str(code_payload['sid']), + int(code_payload['userId']), + str(code_payload['subjectId']), + int(code_payload['authVersion']), + current, + ) + grant = await cls._require_grant( + db, + code_payload.get('grantId'), + int(code_payload['userId']), + client_row.client_pk, + list(code_payload['scopes']), + list(code_payload['resources']), + current, + ) + grant.last_used_at = current + scopes, resources, resource = await cls._validate_client_scope_resource( + db, client_row, list(code_payload['scopes']), list(code_payload['resources']) + ) + access_token, expires_in = await cls._issue_access_token( + db, + client_row, + user, + subject, + session, + scopes, + resources, + resource, + grant_type='authorization_code', + grant_id=grant.grant_id, + signing_key=signing_key, + kid=kid, + now=current, + ) + refresh_token = None + if 'offline_access' in scopes and 'refresh_token' in OidcUtil.json_list(client_row.grant_types): + refresh_token = await cls._create_refresh_token( + db, + client_row, + grant, + user, + subject, + session, + scopes, + resources, + current, + token_pepper=cls._token_pepper(token_pepper), + ) + id_token = None + if 'openid' in scopes: + id_token = await cls._issue_id_token( + db, + client_row, + user, + subject, + session, + scopes, + code_payload['nonce'], + access_token, + signing_key=signing_key, + kid=kid, + now=current, + ) + await SsoSessionDao.record_client(db, session.sid, client_row.client_pk, current) + await AuditService.record( + db, + OidcAuditEvent.TOKEN_ISSUED, + 'success', + client_id=client_row.client_id, + subject_id=subject.subject_id, + sid=session.sid, + grant_id=getattr(grant, 'grant_id', None), + ) + + return TokenResult(access_token, expires_in, refresh_token, ' '.join(scopes), id_token) + + @classmethod + async def refresh_token( # noqa: PLR0915 + cls, + db: AsyncSession, + request: TokenRequest | Mapping[str, Any], + client: OAuthClientPrincipal, + *, + signing_key: RSAPrivateKey | None = None, + kid: str | None = None, + now: datetime | None = None, + token_pepper: str | bytes | None = None, + ) -> TokenResult: + """ + 校验 Refresh Token 并完成令牌轮换 + + :param db: 异步数据库会话,调用方负责提交或回滚事务 + :param request: Refresh Token 请求或字段映射 + :param client: 已完成 Client Authentication 的不可伪造主体 + :param signing_key: 可注入的 RSA 私钥 + :param kid: 注入私钥对应的 Key ID + :param now: 可注入 项目当前时间 + :param token_pepper: Refresh Token HMAC Pepper + :return: 新 Access Token 与轮换后的 Refresh Token + :raises OAuthProtocolException: Token 无效、过期、重放或绑定状态失效 + :raises RefreshTokenReuseDetected: 检测到 Refresh Token 重放 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + parsed = cls._request(request) + if parsed.get('grant_type') != 'refresh_token' or not isinstance(parsed.get('refresh_token'), str): + cls._invalid_request() + client_row = await cls._resolve_client(db, client) + await OAuthAccessPolicyDao.lock_client(db, client_row.client_pk) + if parsed.get('client_id') not in (None, client_row.client_id) or ( + client_row.client_type == 'public' and parsed.get('client_id') != client_row.client_id + ): + cls._invalid_client() + if 'refresh_token' not in OidcUtil.json_list(client_row.grant_types): + cls._invalid_grant() + pepper = cls._token_pepper(token_pepper) + try: + opaque = parse_opaque_token(parsed['refresh_token'], 'rt1') + digest = token_digest(parsed['refresh_token'], pepper) + except (OpaqueTokenError, TypeError, ValueError): + cls._invalid_grant() + row = await OAuthTokenDao.get_by_token_id(db, opaque.token_id, for_update=True) + if row is None or not hmac.compare_digest(row.token_hash, digest): + cls._invalid_grant() + if row.client_pk != client_row.client_pk: + cls._invalid_grant() + if row.status != 'active': + if row.status == 'used': + await OAuthTokenDao.refresh_token_family_reuse(db, row.family_id, row.token_id, now=current) + await AuditService.record( + db, + OidcAuditEvent.REFRESH_REUSE_DETECTED, + 'failure', + risk_level='high', + client_id=client_row.client_id, + subject_id=row.subject_id, + sid=row.sid, + token_id=row.token_id, + ) + raise RefreshTokenReuseDetected + cls._invalid_grant() + idle_expires = TimezoneUtil.to_optional_utc(row.idle_expires_at) + absolute_expires = TimezoneUtil.to_optional_utc(row.absolute_expires_at) + if idle_expires is None or absolute_expires is None or idle_expires <= current or absolute_expires <= current: + row.status = 'expired' + cls._invalid_grant() + user, subject = await cls._require_user_identity(db, row.user_id, row.subject_id, row.auth_version) + session = await cls._require_session( + db, row.sid, row.user_id, row.subject_id, row.auth_version, current, allow_offline=True + ) + grant = await cls._require_grant( + db, row.grant_id, row.user_id, row.client_pk, row.scopes, row.resources, current + ) + if grant is None: + cls._invalid_grant() + grant.last_used_at = current + requested_scopes = cls._requested_scopes(parsed.get('scope'), row.scopes) + requested_resources = cls._requested_resources(parsed.get('resource'), row.resources) + if not set(requested_scopes).issubset(set(row.scopes)) or requested_resources != list(row.resources): + cls._invalid_grant() + scopes, resources, resource = await cls._validate_client_scope_resource( + db, client_row, requested_scopes, requested_resources + ) + new_token = generate_refresh_token() + new_opaque = parse_opaque_token(new_token, 'rt1') + # 先写入后继令牌,再更新即时校验的自引用外键 + new_refresh = await cls._create_refresh_token( + db, + client_row, + grant, + user, + subject, + session, + scopes, + resources, + current, + token_pepper=pepper, + token=new_token, + family_id=row.family_id, + parent_token_id=row.token_id, + absolute_expires_at=absolute_expires, + ) + if not await OAuthTokenDao.mark_used(db, row.token_id, new_opaque.token_id, now=current): + await OAuthTokenDao.refresh_token_family_reuse(db, row.family_id, row.token_id, now=current) + await AuditService.record( + db, + OidcAuditEvent.REFRESH_REUSE_DETECTED, + 'failure', + risk_level='high', + client_id=client_row.client_id, + subject_id=row.subject_id, + sid=row.sid, + token_id=row.token_id, + ) + raise RefreshTokenReuseDetected + access_token, expires_in = await cls._issue_access_token( + db, + client_row, + user, + subject, + session, + scopes, + resources, + resource, + grant_type='refresh_token', + grant_id=grant.grant_id, + signing_key=signing_key, + kid=kid, + now=current, + ) + await AuditService.record( + db, + OidcAuditEvent.REFRESH_ROTATED, + 'success', + client_id=client_row.client_id, + subject_id=subject.subject_id, + sid=session.sid, + grant_id=grant.grant_id, + ) + + return TokenResult(access_token, expires_in, new_refresh, ' '.join(scopes)) + + @classmethod + async def client_credentials( + cls, + db: AsyncSession, + request: TokenRequest | Mapping[str, Any], + client: OAuthClientPrincipal, + *, + signing_key: RSAPrivateKey | None = None, + kid: str | None = None, + now: datetime | None = None, + ) -> TokenResult: + """ + 按 Client Credentials Grant 签发机器令牌 + + :param db: 异步数据库会话,调用方负责提交事务 + :param request: Client Credentials 请求或字段映射 + :param client: 已完成 Client Authentication 的不可伪造主体 + :param signing_key: 可注入的 RSA 私钥 + :param kid: 注入私钥对应的 Key ID + :param now: 可注入 项目当前时间 + :return: 不包含 Refresh/ID Token 的机器 Token 结果 + :raises OAuthProtocolException: Client、Scope 或 Resource 策略无效时抛出 + """ + + current = TimezoneUtil.to_optional_utc(now) or TimezoneUtil.utc_now() + parsed = cls._request(request) + if parsed.get('grant_type') != 'client_credentials': + cls._invalid_request() + client_row = await cls._resolve_client(db, client) + if parsed.get('client_id') not in (None, client_row.client_id) or ( + client_row.client_type == 'public' and parsed.get('client_id') != client_row.client_id + ): + cls._invalid_client() + if client_row.client_type != 'confidential' or 'client_credentials' not in OidcUtil.json_list( + client_row.grant_types + ): + cls._invalid_client() + if isinstance(client, OAuthClientPrincipal) and client.auth_method == 'none': + cls._invalid_client() + requested = cls._requested_scopes(parsed.get('scope'), []) + if not requested: + cls._invalid_scope() + resource_value = parsed.get('resource') + if not isinstance(resource_value, str) or not resource_value: + cls._invalid_scope() + scopes, resources, resource = await cls._validate_client_scope_resource( + db, client_row, requested, [resource_value], machine_only=True + ) + ttl = cls._access_ttl(client_row, resource) + signer, signing_kid = await cls._resolve_signer(db, signing_key, kid, current) + claims = { + 'iss': OidcConfig.oidc_issuer, + 'sub': f'client:{client_row.client_id}', + 'aud': resources, + 'client_id': client_row.client_id, + 'scope': ' '.join(scopes), + 'gty': 'client_credentials', + 'client_policy_version': client_row.policy_version, + 'iat': int(current.timestamp()), + 'nbf': int(current.timestamp()), + 'exp': int((current + timedelta(seconds=ttl)).timestamp()), + 'jti': str(uuid4()), + } + try: + access = encode_access_token(claims, signer, signing_kid) + except (JwtProfileError, TypeError, ValueError): + cls._server_error() + return TokenResult(access, ttl, scope=' '.join(scopes)) + + @classmethod + async def issue_token( + cls, + db: AsyncSession, + redis: Redis, + request: TokenRequest | Mapping[str, Any], + client: OAuthClientPrincipal, + **kwargs: Any, + ) -> TokenResult: + """ + 根据 Grant Type 调度令牌签发流程 + + :param db: 异步数据库会话 + :param redis: 异步 Redis 客户端 + :param request: 含 grant_type 等字段的 TokenRequest 或请求字段映射 + :param client: 已认证的 OAuthClientPrincipal 主体 + :param kwargs: 按 Grant Type 传给具体令牌签发方法的可选参数 + :return: 包含 Access Token 及有效期的 Token Endpoint 响应对象 + :raises AssertionError: 内部令牌状态不满足签发条件 + """ + + if not isinstance(client, OAuthClientPrincipal): + cls._invalid_client() + grant_type = request.grant_type if isinstance(request, TokenRequest) else request.get('grant_type') + if grant_type == 'authorization_code': + return await cls.authorization_code(db, redis, request, client, **kwargs) + if grant_type == 'refresh_token': + return await cls.refresh_token(db, request, client, **kwargs) + if grant_type == 'client_credentials': + return await cls.client_credentials(db, request, client, **kwargs) + cls._invalid_request() + raise AssertionError('程序进入了不可达分支') + + @classmethod + async def _issue_token_transaction( + cls, + db: AsyncSession, + redis: Redis, + request: TokenRequest | Mapping[str, Any], + client: OAuthClientPrincipal, + **kwargs: Any, + ) -> TokenResult: + """ + 在事务边界内执行令牌签发 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: Authorization Code 所在 Redis 客户端 + :param request: 已解析的 Token Endpoint 字段 + :param client: 已认证的 Client 主体 + :param kwargs: 传给令牌签发流程的可选参数 + :return: 标准 Token 领域结果 + :raises RefreshTokenReuseDetected: Family 撤销已写入但需映射为 invalid_grant + :raises OAuthProtocolException: 业务失败且事务已回滚 + """ + + try: + result = await cls.issue_token(db, redis, request, client, **kwargs) + except (RefreshTokenReuseDetected, AuthorizationCodeReuseError): + # 令牌重放属于安全事件,撤销状态和审计记录必须持久化 + try: + await db.commit() + except Exception: + try: + await db.rollback() + except Exception: + pass + raise OAuthProtocolException('server_error', 'Token endpoint is unavailable', 500) from None + raise + except Exception: + await db.rollback() + raise + await db.commit() + + return result + + @classmethod + async def issue_token_request( + cls, + db: AsyncSession, + redis: Redis, + request: TokenRequest | Mapping[str, Any], + *, + authorization: str | None = None, + client_id: str | None = None, + client_secret: str | None = None, + **kwargs: Any, + ) -> TokenResult: + """ + 认证 Client 并处理 Token Endpoint 请求 + + :param db: 异步数据库会话,由本方法提交或回滚 + :param redis: Authorization Code 所在 Redis 客户端 + :param request: 已解析的 Token Endpoint 字段 + :param authorization: RFC 7617 Basic Header + :param client_id: 表单 Client ID + :param client_secret: 表单 Client Secret + :param kwargs: 传给令牌签发流程的可选参数 + :return: 标准 Token 领域结果 + """ + + try: + _, principal = await cls.authenticate_client( + db, + authorization=authorization, + client_id=client_id, + client_secret=client_secret, + ) + except Exception: + await db.rollback() + raise + return await cls._issue_token_transaction(db, redis, request, principal, **kwargs) + + @classmethod + async def _issue_access_token( # noqa: PLR0913 + cls, + db: AsyncSession, + client: SysOAuthClient, + user: SysUser, + subject: SysIdentitySubject, + session: SysSsoSession, + scopes: list[str], + resources: list[str], + resource: SysOAuthResource | None, + *, + grant_type: str, + grant_id: str | None = None, + signing_key: RSAPrivateKey | None, + kid: str | None, + now: datetime, + ) -> tuple[str, int]: + """ + 构造并签署 Access Token + + :param db: 异步数据库会话 + :param client: OAuth Client ORM(SysOAuthClient) + :param user: 系统用户 ORM(SysUser) + :param subject: 身份 Subject ORM(SysIdentitySubject) + :param session: SSO Session ORM(SysSsoSession) + :param scopes: 已授权且经 Client 绑定校验的 Scope code 列表 + :param resources: 已校验的 Resource audience 列表 + :param resource: 与 Resource audience 对应的 OAuth Resource ORM,或 None + :param grant_type: 产生该 Access Token 的 OAuth Grant Type 字符串 + :param grant_id: 持久授权的 Grant ID,一次性在线授权为 None + :param signing_key: 用于签署 Access Token JWT 的 RSA 私钥,或 None + :param kid: 签名 JWT header 使用的 Key ID,或 None + :param now: 生成 JWT 的 项目当前时间 + :return: 签名后的 JWT Access Token 文本及其有效秒数 + """ + + ttl = cls._access_ttl(client, resource) + claims = await cls._user_claims(db, client, user, subject, scopes, resource) + auth_time = cls._numeric_time(session.auth_time) + session_amr = OidcUtil.json_list(session.amr) + claims.update( + { + 'iss': OidcConfig.oidc_issuer, + 'sub': subject.subject_id, + 'aud': OidcUtil.token_audiences(OidcConfig.oidc_issuer, resources), + 'exp': int((now + timedelta(seconds=ttl)).timestamp()), + 'iat': int(now.timestamp()), + 'nbf': int(now.timestamp()), + 'jti': str(uuid4()), + 'client_id': client.client_id, + 'sid': session.sid, + 'scope': ' '.join(scopes), + 'auth_time': auth_time, + 'acr': session.acr, + 'amr': session_amr, + 'ver': int(subject.auth_version), + 'gty': grant_type, + 'grant_id': grant_id, + 'client_policy_version': client.policy_version, + } + ) + signer, signing_kid = await cls._resolve_signer(db, signing_key, kid, now) + try: + return encode_access_token(claims, signer, signing_kid), ttl + except (JwtProfileError, TypeError, ValueError): + cls._server_error() + + @classmethod + async def _issue_id_token( # noqa: PLR0913 + cls, + db: AsyncSession, + client: SysOAuthClient, + user: SysUser, + subject: SysIdentitySubject, + session: SysSsoSession, + scopes: list[str], + nonce: str, + access_token: str, + *, + signing_key: RSAPrivateKey | None, + kid: str | None, + now: datetime, + ) -> str: + """ + 构造包含用户声明的 ID Token + + :param db: 异步数据库会话 + :param client: OAuth Client ORM(SysOAuthClient) + :param user: 系统用户 ORM(SysUser) + :param subject: 身份 Subject ORM(SysIdentitySubject) + :param session: SSO Session ORM(SysSsoSession) + :param scopes: 已授权并用于筛选 ID Token 用户声明的 Scope code 列表 + :param nonce: 授权请求绑定的 OIDC nonce 字符串 + :param access_token: 同一 Token Endpoint 响应中的 JWT Access Token 文本 + :param signing_key: 用于签署 ID Token JWT 的 RSA 私钥,或 None + :param kid: 签名 JWT header 使用的 Key ID,或 None + :param now: 生成 JWT 的 项目当前时间 + :return: 签名后的 OIDC ID Token JWT 文本 + """ + + ttl = cls._id_ttl(client) + resource = None + claims = await cls._user_claims(db, client, user, subject, scopes, resource) + claims.update( + { + 'iss': OidcConfig.oidc_issuer, + 'sub': subject.subject_id, + 'aud': client.client_id, + 'exp': int((now + timedelta(seconds=ttl)).timestamp()), + 'iat': int(now.timestamp()), + 'auth_time': cls._numeric_time(session.auth_time), + 'nonce': nonce, + 'sid': session.sid, + 'acr': session.acr, + 'amr': OidcUtil.json_list(session.amr), + 'at_hash': OidcUtil.access_token_hash(access_token), + } + ) + signer, signing_kid = await cls._resolve_signer(db, signing_key, kid, now) + try: + return encode_id_token(claims, signer, signing_kid) + except (JwtProfileError, TypeError, ValueError): + cls._server_error() + + @classmethod + async def _user_claims( + cls, + db: AsyncSession, + client: SysOAuthClient, + user: SysUser, + subject: SysIdentitySubject, + scopes: list[str], + resource: SysOAuthResource | None, + ) -> dict[str, Any]: + """ + 收集 ID Token 所需的用户声明 + + :param db: 异步数据库会话 + :param client: OAuth Client ORM(SysOAuthClient) + :param user: 系统用户 ORM(SysUser) + :param subject: 身份 Subject ORM(SysIdentitySubject) + :param scopes: 用于解析声明策略的 Scope code 列表 + :param resource: 目标 OAuth Resource ORM,或 None(ID Token 不绑定资源) + :return: 按 Scope 策略生成的 OIDC 用户声明字典 + """ + + policy, allowed = await ClaimService.resolve_scope_policy( + db, + client.client_pk, + scopes, + resource_allowed_claims=resource.allowed_claims if resource is not None else None, + ) + roles, department = await ClaimService.load_roles_and_department(db, user.user_id) + claims = ClaimService.build_claims( + user, + scopes, + policy, + allowed, + subject_id=subject.subject_id, + roles=roles, + department=department, + ) + + return claims + + @classmethod + async def _create_refresh_token( # noqa: PLR0913 + cls, + db: AsyncSession, + client: SysOAuthClient, + grant: SysOAuthGrant, + user: SysUser, + subject: SysIdentitySubject, + session: SysSsoSession, + scopes: list[str], + resources: list[str], + now: datetime, + *, + token_pepper: str | bytes | None, + token: str | None = None, + family_id: str | None = None, + parent_token_id: str | None = None, + absolute_expires_at: datetime | None = None, + ) -> str: + """ + 创建并持久化 Refresh Token 记录 + + :param db: 异步数据库会话 + :param client: OAuth Client ORM(SysOAuthClient) + :param grant: 用户对该 Client 的 OAuth Grant ORM(SysOAuthGrant) + :param user: 系统用户 ORM(SysUser) + :param subject: 身份 Subject ORM(SysIdentitySubject) + :param session: 绑定用户身份版本的 SSO Session ORM(SysSsoSession) + :param scopes: 写入 Refresh Token 记录的 Scope code 列表 + :param resources: 写入 Refresh Token 记录的 Resource audience 列表 + :param now: 签发 Refresh Token 记录的 项目当前时间 + :param token_pepper: 用于 Refresh Token HMAC 摘要的配置值或覆盖值 + :param token: 可复用的 opaque Refresh Token 文本,或 None 表示新生成 + :param family_id: Refresh Token 轮换族标识,或 None 表示创建新族 + :param parent_token_id: 被轮换的父 Refresh Token ID,或 None + :param absolute_expires_at: 轮换族已有的绝对过期时间,或 None + :return: 新建 opaque Refresh Token 文本 + """ + + raw = token or generate_refresh_token() + parsed = parse_opaque_token(raw, 'rt1') + pepper = cls._token_pepper(token_pepper) + absolute = absolute_expires_at or now + timedelta(seconds=cls._refresh_absolute_ttl(client)) + absolute = min(absolute, now + timedelta(seconds=cls._refresh_absolute_ttl(client))) if absolute else absolute + idle = min(now + timedelta(seconds=cls._refresh_idle_ttl(client)), absolute) + row = SysOAuthRefreshToken( + token_id=parsed.token_id, + token_hash=token_digest(raw, pepper), + family_id=family_id or str(uuid4()), + parent_token_id=parent_token_id, + grant_id=grant.grant_id, + user_id=user.user_id, + subject_id=subject.subject_id, + auth_version=subject.auth_version, + client_pk=client.client_pk, + sid=session.sid, + scopes=list(scopes), + resources=list(resources), + status='active', + issued_at=now, + idle_expires_at=idle, + absolute_expires_at=absolute, + ) + await OAuthTokenDao.create(db, row) + + return raw + + @classmethod + async def _validate_client_scope_resource( + cls, + db: AsyncSession, + client: SysOAuthClient, + scopes: list[str], + resources: list[str], + *, + machine_only: bool = False, + ) -> tuple[list[str], list[str], SysOAuthResource | None]: + """ + 校验 Client 请求的 Scope 和 Resource + + :param db: 异步数据库会话 + :param client: OAuth Client ORM(SysOAuthClient) + :param scopes: 请求或授权记录中的 Scope code 列表 + :param resources: 请求或授权记录中的 Resource audience 列表 + :param machine_only: 是否限制为 resource 类型的机器 Scope + :return: 去重后的 Scope code 列表、Resource audience 列表及匹配的 Resource ORM + """ + + if len(resources) > cls._MAX_RESOURCE_COUNT or len(set(scopes)) != len(scopes): + cls._invalid_scope() + bindings = await OAuthClientDao.list_scope_bindings(db, client.client_pk) + definitions = await OAuthClientDao.list_scope_definitions(db) + by_pk = {item.scope_pk: item for item in definitions} + bound: dict[str, SysOAuthScope] = {} + for binding in bindings: + definition = by_pk.get(binding.scope_pk) + if definition is not None and definition.status == '0': + bound[definition.scope_code] = definition + if any(scope not in bound for scope in scopes): + cls._invalid_scope() + if machine_only and any(bound[scope].scope_type != cls._MACHINE_SCOPE_TYPE for scope in scopes): + cls._invalid_scope() + resource_rows = list(await OAuthClientDao.list_resources(db, client.client_pk)) + resource_by_audience = {row.audience: row for row in resource_rows if row.status == '0'} + resource = resource_by_audience.get(resources[0]) if resources else None + if resources and resource is None: + cls._invalid_scope() + for scope in scopes: + definition = bound[scope] + if definition.scope_type == 'resource' and ( + resource is None or definition.resource_pk != resource.resource_pk + ): + cls._invalid_scope() + if resource is not None and resource.signing_alg != 'RS256': + cls._invalid_scope() + return list(dict.fromkeys(scopes)), list(dict.fromkeys(resources)), resource + + @classmethod + async def _require_grant( + cls, + db: AsyncSession, + grant_id: object, + user_id: int, + client_pk: int, + scopes: list[str], + resources: list[str], + now: datetime, + ) -> SysOAuthGrant: + """ + 加载并验证用户授权 Grant + + :param db: 异步数据库会话 + :param grant_id: 授权记录的非空 grant_id 字符串 + :param user_id: 系统用户主键 + :param client_pk: OAuth Client ORM 主键 + :param scopes: 需要包含在 Grant 授权范围内的 Scope code 列表 + :param resources: 需要包含在 Grant 授权范围内的 Resource audience 列表 + :param now: 检查 Grant 过期状态的 项目当前时间 + :return: 匹配请求范围的有效授权 Grant + :raises OAuthProtocolException: 授权记录或用户访问策略无效 + """ + + if not isinstance(grant_id, str) or not grant_id: + cls._invalid_grant() + if await OAuthAccessPolicyDao.is_blocked(db, user_id, client_pk, for_update=True): + cls._invalid_grant() + grant = await OAuthGrantDao.get_by_grant_id_for_update(db, grant_id, refresh=True) + client = await OAuthClientDao.get_by_pk(db, client_pk, active_only=True) + if ( + grant is None + or client is None + or grant.user_id != user_id + or grant.client_pk != client_pk + or grant.status != cls._ACTIVE_GRANT_STATUS + or grant.client_policy_version != client.policy_version + or (grant.expires_at is not None and TimezoneUtil.to_optional_utc(grant.expires_at) <= now) + or not set(scopes).issubset(set(grant.granted_scopes or [])) + or not set(resources).issubset(set(grant.granted_resources or [])) + ): + cls._invalid_grant() + return grant + + @classmethod + async def _require_user_identity( + cls, db: AsyncSession, user_id: int, subject_id: str, auth_version: int + ) -> tuple[SysUser, SysIdentitySubject]: + """ + 验证用户与 Subject 的身份版本 + + :param db: 异步数据库会话 + :param user_id: 系统用户主键 + :param subject_id: 身份 Subject 的稳定标识 + :param auth_version: 必须与用户 Subject 和 SSO Session 一致的认证版本 + :return: 通过状态、删除标记及认证版本校验的系统用户 ORM 与 Subject ORM + """ + + user = await IdentityUserDao.get_user(db, user_id) + try: + subject = await IdentitySubjectService.require_by_user_id( + db, + user_id, + audit_writer=lambda missing_user_id: AuditService.record_independent( + db, + OidcAuditEvent.IDENTITY_SUBJECT_MISSING, + 'failure', + risk_level='high', + user_id=missing_user_id, + failure_code='identity_integrity', + ), + ) + except OAuthProtocolException: + cls._invalid_grant() + if ( + user is None + or user.status != '0' + or user.del_flag != '0' + or subject.subject_id != subject_id + or subject.auth_version != auth_version + ): + cls._invalid_grant() + return user, subject + + @classmethod + async def _require_session( + cls, + db: AsyncSession, + sid: str, + user_id: int, + subject_id: str, + auth_version: int, + now: datetime, + *, + allow_offline: bool = False, + ) -> SysSsoSession: + """ + 加载并验证关联的 SSO Session + + :param db: 异步数据库会话 + :param sid: SSO Session 的 sid 字符串 + :param user_id: 系统用户主键 + :param subject_id: 必须与 Session 绑定一致的身份 Subject 标识 + :param auth_version: 必须与 Session 绑定一致的认证版本 + :param now: 检查 Session active 和过期状态的 项目当前时间 + :param allow_offline: 是否允许自然过期但未撤销的SSO会话继续离线续期 + :return: 与用户身份和 Token 绑定一致的 SSO Session + """ + + session = await SsoSessionDao.get_for_token(db, sid, now=now, allow_offline=allow_offline) + if ( + session is None + or session.status not in ({'active', 'expired'} if allow_offline else {'active'}) + or session.user_id != user_id + or session.subject_id != subject_id + or session.auth_version != auth_version + ): + cls._invalid_grant() + return session + + @classmethod + async def _resolve_client(cls, db: AsyncSession, client: OAuthClientPrincipal) -> SysOAuthClient: + """ + 确认 OAuth Client 处于可用状态 + + :param db: 异步数据库会话 + :param client: 已认证的 OAuthClientPrincipal 主体 + :return: 通过 active 状态、Client 类型及认证方式校验的 OAuth Client ORM + """ + + if not isinstance(client, OAuthClientPrincipal): + cls._invalid_client() + row = await OAuthClientDao.get_by_client_id(db, client.client_id, active_only=True) + if row is None: + cls._invalid_client() + if client.client_id != row.client_id: + cls._invalid_client() + expected_method = 'none' if row.client_type == 'public' else 'client_secret_basic' + if client.client_type != row.client_type or client.auth_method != expected_method: + cls._invalid_client() + return row + + @classmethod + async def _resolve_signer( + cls, + db: AsyncSession, + signing_key: RSAPrivateKey | None, + kid: str | None, + now: datetime, + ) -> tuple[RSAPrivateKey, str]: + """ + 加载或确认 JWT 签名密钥 + + :param db: 异步数据库会话 + :param signing_key: 可注入的 RSA JWT 签名私钥,或 None + :param kid: 注入私钥对应的 JWT Key ID,或 None + :param now: 读取 active OIDC 签名密钥的 项目当前时间 + :return: 用于 JWT 签名的 RSA 私钥及其 Key ID + """ + + if signing_key is not None: + if not isinstance(kid, str) or not kid: + cls._server_error() + return signing_key, kid + record = await OidcKeyDao.get_active(db, alg='RS256') + if record is None: + cls._server_error() + try: + key = await KeyService.load_private_key_async(record, now=now) + except (KeyServiceError, ValueError, OSError): + cls._server_error() + return key, record.kid + + @classmethod + def _access_ttl( + cls, + client: SysOAuthClient, + resource: SysOAuthResource | None, + ) -> int: + """ + 计算 Access Token 的有效秒数 + + :param client: OAuth Client ORM(SysOAuthClient) + :param resource: 目标 OAuth Resource ORM,或 None + :return: Access Token 有效秒数(取 Client、Resource 与 OIDC 上限的最小值) + """ + + values: list[int] = [] + if client.access_token_ttl_seconds is not None: + values.append(client.access_token_ttl_seconds) + if resource is not None and resource.access_token_ttl_seconds is not None: + values.append(resource.access_token_ttl_seconds) + if any(isinstance(value, bool) or not isinstance(value, int) or value <= 0 for value in values): + cls._server_error() + if not values: + values.append(OidcConfig.oidc_access_token_ttl_seconds) + maximum = OidcConfig.oidc_max_access_token_ttl_seconds + if isinstance(maximum, bool) or not isinstance(maximum, int) or maximum <= 0: + cls._server_error() + return min(*values, maximum) + + @staticmethod + def _id_ttl(client: SysOAuthClient) -> int: + """ + 计算 ID Token 的有效秒数 + + :param client: OAuth Client ORM(此方法仅按统一签名接收) + :return: OIDC ID Token 的有效秒数 + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + value = OidcConfig.oidc_id_token_ttl_seconds + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise OAuthProtocolException('server_error', 'Token policy is unavailable', 500) + return value + + @staticmethod + def _refresh_idle_ttl(client: SysOAuthClient) -> int: + """ + 计算 Refresh Token 空闲有效秒数 + + :param client: OAuth Client ORM(SysOAuthClient) + :return: Refresh Token 空闲有效秒数(不超过 OIDC 配置上限) + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + value = client.refresh_token_idle_seconds + if value is None: + value = OidcConfig.oidc_refresh_token_idle_seconds + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise OAuthProtocolException('server_error', 'Token policy is unavailable', 500) + maximum = OidcConfig.oidc_refresh_token_idle_seconds + if isinstance(maximum, bool) or not isinstance(maximum, int) or maximum <= 0: + raise OAuthProtocolException('server_error', 'Token policy is unavailable', 500) + return min(value, maximum) + + @staticmethod + def _refresh_absolute_ttl(client: SysOAuthClient) -> int: + """ + 计算 Refresh Token 绝对有效秒数 + + :param client: OAuth Client ORM(SysOAuthClient) + :return: Refresh Token 绝对有效秒数(不超过 OIDC 配置上限) + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + value = client.refresh_token_absolute_seconds + if value is None: + value = OidcConfig.oidc_refresh_token_absolute_seconds + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + raise OAuthProtocolException('server_error', 'Token policy is unavailable', 500) + maximum = OidcConfig.oidc_refresh_token_absolute_seconds + if isinstance(maximum, bool) or not isinstance(maximum, int) or maximum <= 0: + raise OAuthProtocolException('server_error', 'Token policy is unavailable', 500) + return min(value, maximum) + + @classmethod + def _request(cls, request: TokenRequest | Mapping[str, Any]) -> dict[str, Any]: + """ + 将 TokenRequest 或映射转换为请求字典 + + :param request: TokenRequest 实例或包含 Token Endpoint 字段的映射 + :return: 用于 Grant Type 分派的已验证 Token Endpoint 字段字典 + """ + + if isinstance(request, TokenRequest): + return request.model_dump() + if isinstance(request, Mapping): + try: + return TokenRequest(**dict(request)).model_dump() + except ValidationError: + cls._invalid_request() + cls._invalid_request() + + return {} + + @staticmethod + def _token_pepper(override: str | bytes | None) -> str | bytes: + """ + 选择 Token HMAC Pepper 配置 + + :param override: 覆盖配置 + :return: 用于令牌摘要计算的 Pepper 文本或字节串 + """ + + value = OidcConfig.oidc_token_hash_pepper if override is None else override + if not isinstance(value, (str, bytes)): + TokenService._server_error() + raw = value.encode('utf-8') if isinstance(value, str) else value + if len(raw) < TokenService._MIN_TOKEN_PEPPER_BYTES: + TokenService._server_error() + return value + + @classmethod + def _requested_scopes(cls, value: Any, fallback: Sequence[str]) -> list[str]: + """ + 解析请求中的 Scope 列表 + + :param value: Token Endpoint 的 scope 字符串,或已解析的 Scope code 序列 + :param fallback: 请求未提供 scope 时使用的已持久化 Scope code 序列 + :return: 去重校验后的 Scope code 字符串列表 + """ + + raw = fallback if value is None else value.split() if isinstance(value, str) else value + if not isinstance(raw, (list, tuple)) or not raw or any(not isinstance(item, str) or not item for item in raw): + cls._invalid_scope() + if len(set(raw)) != len(raw): + cls._invalid_scope() + return list(raw) + + @classmethod + def _requested_resources(cls, value: Any, fallback: Sequence[str]) -> list[str]: + """ + 解析请求中的 Resource 列表 + + :param value: Token Endpoint 的 resource audience 字符串或序列 + :param fallback: 请求未提供 resource 时使用的已持久化 audience 序列 + :return: 通过单 Resource 限制校验的 audience 字符串列表 + """ + + raw = fallback if value is None else [value] if isinstance(value, str) else value + if not isinstance(raw, (list, tuple)) or len(raw) > cls._MAX_RESOURCE_COUNT: + cls._invalid_scope() + if any(not isinstance(item, str) or not item for item in raw): + cls._invalid_scope() + return list(raw) + + @staticmethod + def _numeric_time(value: datetime | None) -> int: + """ + 将时间值转换为整数时间戳 + + :param value: 认证时间等需要编码进 JWT claim 的 datetime,或 None + :return: Unix 整数时间戳 + """ + + current = TimezoneUtil.to_optional_utc(value) + if current is None: + TokenService._server_error() + return int(current.timestamp()) + + @staticmethod + def _map_oauth_failure(error: Exception, fallback: str) -> None: + """ + 将底层异常映射为 OAuth 协议错误 + + :param error: Authorization Code 消费等底层流程抛出的异常 + :param fallback: 非 OAuthProtocolException 时使用的 OAuth error code + :return: None + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + if isinstance(error, OAuthProtocolException): + raise OAuthProtocolException(error.error, 'The authorization grant is invalid or expired', 400) from None + raise OAuthProtocolException(fallback, 'The authorization grant is invalid or expired', 400) from None + + @staticmethod + def _invalid_request() -> None: + """ + 抛出 invalid_request 协议错误 + + :return: None + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + raise OAuthProtocolException('invalid_request', 'Invalid token request', 400) + + @staticmethod + def _invalid_client() -> None: + """ + 抛出 invalid_client 协议错误 + + :return: None + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + raise OAuthProtocolException('invalid_client', 'Client authentication failed', 401) + + @staticmethod + def _invalid_grant() -> None: + """ + 抛出 invalid_grant 协议错误 + + :return: None + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + raise OAuthProtocolException('invalid_grant', 'The authorization grant is invalid or expired', 400) + + @staticmethod + def _invalid_scope() -> None: + """ + 抛出 invalid_scope 协议错误 + + :return: None + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + raise OAuthProtocolException('invalid_scope', 'Requested scope is not authorized', 400) + + @staticmethod + def _server_error() -> None: + """ + 抛出 server_error 协议错误 + + :return: None + :raises OAuthProtocolException: 请求不符合 OAuth 协议约束 + """ + + raise OAuthProtocolException('server_error', 'Token issuance is unavailable', 500) diff --git a/ruoyi-fastapi-backend/plugins/core/discovery/__init__.py b/ruoyi-fastapi-backend/plugins/core/discovery/__init__.py index bfdcbc116..e69de29bb 100644 --- a/ruoyi-fastapi-backend/plugins/core/discovery/__init__.py +++ b/ruoyi-fastapi-backend/plugins/core/discovery/__init__.py @@ -1,3 +0,0 @@ -""" -插件发现与注册表分层包。 -""" diff --git a/ruoyi-fastapi-backend/plugins/core/lifecycle/__init__.py b/ruoyi-fastapi-backend/plugins/core/lifecycle/__init__.py index bb1f05aa6..e69de29bb 100644 --- a/ruoyi-fastapi-backend/plugins/core/lifecycle/__init__.py +++ b/ruoyi-fastapi-backend/plugins/core/lifecycle/__init__.py @@ -1,3 +0,0 @@ -""" -插件生命周期能力分层包。 -""" diff --git a/ruoyi-fastapi-backend/plugins/core/management/service/startup_gateway.py b/ruoyi-fastapi-backend/plugins/core/management/service/startup_gateway.py index 5ceb3f41b..272469422 100644 --- a/ruoyi-fastapi-backend/plugins/core/management/service/startup_gateway.py +++ b/ruoyi-fastapi-backend/plugins/core/management/service/startup_gateway.py @@ -1,21 +1,16 @@ -from __future__ import annotations +from pathlib import Path +from typing import Any -from typing import TYPE_CHECKING, Any +from sqlalchemy.ext.asyncio import AsyncSession +from common.vo import CrudResponseModel +from plugins.core.discovery.scanner import DiscoveredPlugin from plugins.core.management.dao.dao import PluginDao +from plugins.core.management.entity.vo.schemas import PluginMigrationModel, PluginModel from plugins.core.management.service.gateway import PluginManagementRuntimeGateway from plugins.core.management.service.service import PluginService from plugins.core.state import PluginStateResolver -if TYPE_CHECKING: - from pathlib import Path - - from sqlalchemy.ext.asyncio import AsyncSession - - from common.vo import CrudResponseModel - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.management.entity.vo.schemas import PluginMigrationModel, PluginModel - class PluginManagementStartupGateway: """ diff --git a/ruoyi-fastapi-backend/plugins/core/manifest/__init__.py b/ruoyi-fastapi-backend/plugins/core/manifest/__init__.py index 14d3a76f4..0cbedb2a4 100644 --- a/ruoyi-fastapi-backend/plugins/core/manifest/__init__.py +++ b/ruoyi-fastapi-backend/plugins/core/manifest/__init__.py @@ -1,5 +1 @@ -""" -插件 manifest 分层包。 -""" - from plugins.core.manifest.schema import * # noqa: F403 diff --git a/ruoyi-fastapi-backend/plugins/core/manifest/menu_tree.py b/ruoyi-fastapi-backend/plugins/core/manifest/menu_tree.py index e1b844d5b..5f557e963 100644 --- a/ruoyi-fastapi-backend/plugins/core/manifest/menu_tree.py +++ b/ruoyi-fastapi-backend/plugins/core/manifest/menu_tree.py @@ -1,5 +1,3 @@ -from __future__ import annotations - from pathlib import Path from typing import TYPE_CHECKING @@ -18,7 +16,7 @@ class PluginMenuTree: """ @classmethod - def flatten(cls, menus: list[PluginMenuManifest]) -> list[PluginMenuManifest]: + def flatten(cls, menus: 'list[PluginMenuManifest]') -> 'list[PluginMenuManifest]': """ 展平插件菜单树。 @@ -33,7 +31,7 @@ def flatten(cls, menus: list[PluginMenuManifest]) -> list[PluginMenuManifest]: return flattened_menus @classmethod - def count(cls, menus: list[PluginMenuManifest]) -> int: + def count(cls, menus: 'list[PluginMenuManifest]') -> int: """ 统计插件菜单树节点数量。 @@ -43,7 +41,7 @@ def count(cls, menus: list[PluginMenuManifest]) -> int: return len(cls.flatten(menus)) @classmethod - def collect_permissions(cls, menus: list[PluginMenuManifest]) -> set[str]: + def collect_permissions(cls, menus: 'list[PluginMenuManifest]') -> set[str]: """ 收集插件菜单树中的权限标识。 @@ -55,7 +53,7 @@ def collect_permissions(cls, menus: list[PluginMenuManifest]) -> set[str]: @classmethod def collect_route_paths( cls, - menus: list[PluginMenuManifest], + menus: 'list[PluginMenuManifest]', parent_path: str = '', ) -> list[str]: """ @@ -85,7 +83,7 @@ def is_plugin_component(component: str) -> bool: return component.startswith('plugin/') @staticmethod - def resolve_plugin_view_path(manifest: PluginManifest, component: str) -> Path | None: + def resolve_plugin_view_path(manifest: 'PluginManifest', component: str) -> Path | None: """ 将插件组件路径解析为前端插件内视图路径。 diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/__init__.py b/ruoyi-fastapi-backend/plugins/core/runtime/__init__.py index b2a3666e1..d09521252 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/__init__.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/__init__.py @@ -1,7 +1,3 @@ -""" -插件运行时能力分层包。 -""" - from plugins.core.runtime.health import PluginHealthChecker, PluginHealthContext, PluginHealthResult from plugins.core.runtime.hooks import PluginHookContext, PluginHookResult, PluginHookRunner diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/result.py b/ruoyi-fastapi-backend/plugins/core/runtime/result.py index 92d8731cd..48035c232 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/result.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/result.py @@ -1,10 +1,7 @@ -from __future__ import annotations - +from collections.abc import Mapping from dataclasses import dataclass -from typing import TYPE_CHECKING -if TYPE_CHECKING: - from collections.abc import Mapping +from typing_extensions import Self @dataclass(frozen=True) @@ -23,7 +20,7 @@ def from_payload( payload: Mapping[str, object], *, default_message: str = '插件操作完成', - ) -> PluginOperationResult: + ) -> Self: """ 从插件运行时 payload 构建结果视图。 diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/batch.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/batch.py index f627a7b87..82db03d30 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/batch.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/batch.py @@ -1,26 +1,21 @@ -from __future__ import annotations - +from collections.abc import Mapping from pathlib import Path -from typing import TYPE_CHECKING, Protocol, cast +from typing import Protocol, cast +from plugins.core.capability import PluginRuntimeCapability +from plugins.core.discovery.scanner import DiscoveredPlugin from plugins.core.runtime.support import ( BatchOperationResultPayload, PluginBatchReportBuilder, PluginPayloadBuilder, PluginRuntimePayloadBuilder, ) +from plugins.core.types import PluginStateRecord from plugins.core.validation.plugin_deps import PluginBatchOperation, PluginDependencyPlanBuilder -if TYPE_CHECKING: - from collections.abc import Mapping - - from plugins.core.capability import PluginRuntimeCapability - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.types import PluginStateRecord - - from .context import PluginRuntimeContextService - from .dependency_container import PluginRuntimeDependencies - from .responses import PluginBatchResponse, PluginLifecycleResponse, PluginPlanResponse +from .context import PluginRuntimeContextService +from .dependency_container import PluginRuntimeDependencies +from .responses import PluginBatchResponse, PluginLifecycleResponse, PluginPlanResponse class PluginBatchRuntimeOperations(Protocol): diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/gateway.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/gateway.py index d9d683770..798449271 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/gateway.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/gateway.py @@ -1,32 +1,27 @@ -from __future__ import annotations - import subprocess import threading -from collections.abc import Callable -from typing import TYPE_CHECKING, Any, Literal, Protocol, TypeAlias, runtime_checkable - -if TYPE_CHECKING: - from collections.abc import Mapping - from contextlib import AbstractAsyncContextManager - from pathlib import Path - - from sqlalchemy.ext.asyncio import AsyncSession - - from common.vo import CrudResponseModel - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.lifecycle.migration import PluginMigrationResult - from plugins.core.lifecycle.purge import PluginPurgePlan - from plugins.core.management.entity.vo.schemas import ( - PluginConfigModel, - PluginConfigUpdateModel, - PluginConfigValueModel, - PluginMigrationModel, - PluginModel, - PluginOperationLogDetailModel, - PluginOperationLogExportQueryModel, - ) - from plugins.core.types import PluginConfigValue, PluginStateRecord - from plugins.core.validation.menus import PluginMenuConflictItem +from collections.abc import Callable, Mapping +from contextlib import AbstractAsyncContextManager +from pathlib import Path +from typing import Any, Literal, Protocol, TypeAlias, runtime_checkable + +from sqlalchemy.ext.asyncio import AsyncSession + +from common.vo import CrudResponseModel +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.lifecycle.migration import PluginMigrationResult +from plugins.core.lifecycle.purge import PluginPurgePlan +from plugins.core.management.entity.vo.schemas import ( + PluginConfigModel, + PluginConfigUpdateModel, + PluginConfigValueModel, + PluginMigrationModel, + PluginModel, + PluginOperationLogDetailModel, + PluginOperationLogExportQueryModel, +) +from plugins.core.types import PluginConfigValue, PluginStateRecord +from plugins.core.validation.menus import PluginMenuConflictItem PluginCommandOutputKind = Literal['status', 'stdout', 'stderr'] PluginCommandOutputCallback: TypeAlias = Callable[[PluginCommandOutputKind, str], None] diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/enable.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/enable.py index ceb380cb3..338dc5686 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/enable.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/enable.py @@ -1,9 +1,10 @@ -from __future__ import annotations - from dataclasses import dataclass from pathlib import Path -from typing import TYPE_CHECKING, cast +from typing import cast + +from sqlalchemy.ext.asyncio import AsyncSession +from plugins.core.discovery.scanner import DiscoveredPlugin from plugins.core.runtime.support import ( PluginEnablePayloadBuilder, PluginLifecyclePayloadBuilder, @@ -12,19 +13,13 @@ PluginRuntimePayloadBuilder, ) +from ..context import PluginRuntimeContextService +from ..dependency_container import PluginRuntimeDependencies +from ..responses import PluginLifecycleResponse from .common import PluginLifecycleUseCaseSupport +from .operations import PluginLifecycleRuntimeOperations from .runner import PluginLifecycleStep, PluginLifecycleStepFailed, PluginLifecycleStepRunner -if TYPE_CHECKING: - from sqlalchemy.ext.asyncio import AsyncSession - - from plugins.core.discovery.scanner import DiscoveredPlugin - - from ..context import PluginRuntimeContextService - from ..dependency_container import PluginRuntimeDependencies - from ..responses import PluginLifecycleResponse - from .operations import PluginLifecycleRuntimeOperations - @dataclass(slots=True) class PluginEnabledLifecycleContext: diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/install.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/install.py index 39d638942..3db74de1e 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/install.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/install.py @@ -1,12 +1,13 @@ -from __future__ import annotations - from dataclasses import dataclass from pathlib import Path -from typing import TYPE_CHECKING, cast +from typing import cast + +from sqlalchemy.ext.asyncio import AsyncSession -from plugins.core.lifecycle.migration import PluginMigrationError -from plugins.core.lifecycle.seed import PluginSeedRunner -from plugins.core.runtime.hooks import PluginHookRunner +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.lifecycle.migration import PluginMigrationError, PluginMigrationResult +from plugins.core.lifecycle.seed import PluginSeedResult, PluginSeedRunner +from plugins.core.runtime.hooks import PluginHookResult, PluginHookRunner from plugins.core.runtime.support import ( PluginLifecyclePayloadBuilder, PluginPayloadBuilder, @@ -14,22 +15,13 @@ PluginRuntimePayloadBuilder, ) +from ..context import PluginRuntimeContextService +from ..dependency_container import PluginRuntimeDependencies +from ..responses import PluginLifecycleResponse from .common import PluginLifecycleUseCaseSupport +from .operations import PluginLifecycleRuntimeOperations from .runner import PluginLifecycleStep, PluginLifecycleStepFailed, PluginLifecycleStepRunner -if TYPE_CHECKING: - from sqlalchemy.ext.asyncio import AsyncSession - - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.lifecycle.migration import PluginMigrationResult - from plugins.core.lifecycle.seed import PluginSeedResult - from plugins.core.runtime.hooks import PluginHookResult - - from ..context import PluginRuntimeContextService - from ..dependency_container import PluginRuntimeDependencies - from ..responses import PluginLifecycleResponse - from .operations import PluginLifecycleRuntimeOperations - @dataclass(slots=True) class PluginInstallLifecycleContext: diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/purge.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/purge.py index 439351dbe..d9e3c7a07 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/purge.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/purge.py @@ -1,10 +1,11 @@ -from __future__ import annotations - from dataclasses import dataclass from pathlib import Path -from typing import TYPE_CHECKING, cast +from typing import cast + +from sqlalchemy.ext.asyncio import AsyncSession -from plugins.core.runtime.hooks import PluginHookRunner +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.runtime.hooks import PluginHookResult, PluginHookRunner from plugins.core.runtime.support import ( PluginEnablePayloadBuilder, PluginLifecyclePayloadBuilder, @@ -14,20 +15,13 @@ PluginRuntimePayloadBuilder, ) +from ..context import PluginRuntimeContextService +from ..dependency_container import PluginRuntimeDependencies +from ..responses import PluginLifecycleResponse from .common import PluginLifecycleUseCaseSupport +from .operations import PluginLifecycleRuntimeOperations from .runner import PluginLifecycleStep, PluginLifecycleStepFailed, PluginLifecycleStepRunner -if TYPE_CHECKING: - from sqlalchemy.ext.asyncio import AsyncSession - - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.runtime.hooks import PluginHookResult - - from ..context import PluginRuntimeContextService - from ..dependency_container import PluginRuntimeDependencies - from ..responses import PluginLifecycleResponse - from .operations import PluginLifecycleRuntimeOperations - @dataclass(slots=True) class PluginPurgeLifecycleContext: diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/runner.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/runner.py index ca746bd03..a04d4689e 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/runner.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/runner.py @@ -1,10 +1,6 @@ -from __future__ import annotations - +from collections.abc import Awaitable, Callable from dataclasses import dataclass -from typing import TYPE_CHECKING, Generic, TypeVar - -if TYPE_CHECKING: - from collections.abc import Awaitable, Callable +from typing import Generic, TypeVar from ..responses import PluginLifecycleResponse diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/upgrade.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/upgrade.py index 2babbcaf1..2bb2c5775 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/upgrade.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle/upgrade.py @@ -1,36 +1,28 @@ -from __future__ import annotations - from dataclasses import dataclass from pathlib import Path -from typing import TYPE_CHECKING, cast +from typing import cast + +from sqlalchemy.ext.asyncio import AsyncSession -from plugins.core.lifecycle.migration import PluginMigrationError -from plugins.core.lifecycle.seed import PluginSeedRunner -from plugins.core.runtime.hooks import PluginHookRunner +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.lifecycle.migration import PluginMigrationError, PluginMigrationResult +from plugins.core.lifecycle.seed import PluginSeedResult, PluginSeedRunner +from plugins.core.runtime.hooks import PluginHookResult, PluginHookRunner from plugins.core.runtime.support import ( PluginLifecyclePayloadBuilder, PluginPayloadBuilder, PluginPrecheckContext, PluginRuntimePayloadBuilder, ) +from plugins.core.types import PluginStateRecord +from ..context import PluginRuntimeContextService +from ..dependency_container import PluginRuntimeDependencies +from ..responses import PluginLifecycleResponse from .common import PluginLifecycleUseCaseSupport +from .operations import PluginLifecycleRuntimeOperations from .runner import PluginLifecycleStep, PluginLifecycleStepFailed, PluginLifecycleStepRunner -if TYPE_CHECKING: - from sqlalchemy.ext.asyncio import AsyncSession - - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.lifecycle.migration import PluginMigrationResult - from plugins.core.lifecycle.seed import PluginSeedResult - from plugins.core.runtime.hooks import PluginHookResult - from plugins.core.types import PluginStateRecord - - from ..context import PluginRuntimeContextService - from ..dependency_container import PluginRuntimeDependencies - from ..responses import PluginLifecycleResponse - from .operations import PluginLifecycleRuntimeOperations - @dataclass(slots=True) class PluginUpgradeLifecycleContext: diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle_lock.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle_lock.py index 3fb56520b..9f28dd980 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle_lock.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/lifecycle_lock.py @@ -1,11 +1,11 @@ -from __future__ import annotations - import asyncio +from collections.abc import AsyncIterator from contextlib import asynccontextmanager, suppress from dataclasses import dataclass -from typing import TYPE_CHECKING, Protocol +from typing import Protocol from uuid import uuid4 +from redis import asyncio as aioredis from redis.exceptions import RedisError from common.constant import LockConstant @@ -13,11 +13,6 @@ from exceptions.exception import ServiceException from utils.log_util import logger -if TYPE_CHECKING: - from collections.abc import AsyncIterator - - from redis import asyncio as aioredis - @dataclass(frozen=True) class PluginLifecycleLockResult: diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/service/responses.py b/ruoyi-fastapi-backend/plugins/core/runtime/service/responses.py index 71df8d8c7..4ba88339e 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/service/responses.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/service/responses.py @@ -1,5 +1,3 @@ -from __future__ import annotations - from typing import TypeAlias from pydantic import Field diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/startup_coordination.py b/ruoyi-fastapi-backend/plugins/core/runtime/startup_coordination.py index b3c23efc7..bcb011980 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/startup_coordination.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/startup_coordination.py @@ -1,5 +1,3 @@ -from __future__ import annotations - from hashlib import sha256 from pathlib import Path diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/startup_gateway.py b/ruoyi-fastapi-backend/plugins/core/runtime/startup_gateway.py index f72c60bba..add46e110 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/startup_gateway.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/startup_gateway.py @@ -1,15 +1,11 @@ -from __future__ import annotations +from pathlib import Path +from typing import Any, NoReturn, Protocol, runtime_checkable -from typing import TYPE_CHECKING, Any, NoReturn, Protocol, runtime_checkable +from sqlalchemy.ext.asyncio import AsyncSession -if TYPE_CHECKING: - from pathlib import Path - - from sqlalchemy.ext.asyncio import AsyncSession - - from common.vo import CrudResponseModel - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.management.entity.vo.schemas import PluginMigrationModel, PluginModel +from common.vo import CrudResponseModel +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.management.entity.vo.schemas import PluginMigrationModel, PluginModel @runtime_checkable diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/batch_report.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/batch_report.py index 36cece3ed..907964529 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/batch_report.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/batch_report.py @@ -1,5 +1,3 @@ -from __future__ import annotations - from collections.abc import Awaitable, Callable, Mapping from dataclasses import dataclass from time import perf_counter diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/catalog.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/catalog.py index ce62efd48..b528cc553 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/catalog.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/catalog.py @@ -1,25 +1,21 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, TypeAlias +from typing import TypeAlias from pydantic import Field from plugins.core.discovery.registry import RegisteredPlugin +from plugins.core.discovery.scanner import DiscoveredPlugin from plugins.core.manifest.menu_tree import PluginMenuTree +from plugins.core.manifest.schema import ( + PluginConfigItemManifest, + PluginJobManifest, + PluginMenuManifest, + PluginPermissionManifest, +) +from plugins.core.types import PluginStateRecord, SupportsToPayload +from plugins.core.validation.dependencies import DependencyCheckItem from .base import PluginPayloadModel -if TYPE_CHECKING: - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.manifest.schema import ( - PluginConfigItemManifest, - PluginJobManifest, - PluginMenuManifest, - PluginPermissionManifest, - ) - from plugins.core.types import PluginStateRecord, SupportsToPayload - from plugins.core.validation.dependencies import DependencyCheckItem - class PluginCatalogSummaryPayload(PluginPayloadModel): """ diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/lifecycle.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/lifecycle.py index dd50d0f80..3aeb154e1 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/lifecycle.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/lifecycle.py @@ -1,23 +1,18 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Protocol, TypeAlias, cast +from collections.abc import Mapping +from typing import Protocol, TypeAlias, cast from pydantic import Field +from plugins.core.lifecycle.migration import PluginMigrationResult +from plugins.core.lifecycle.seed import PluginSeedResult +from plugins.core.runtime.hooks import PluginHookResult +from plugins.core.types import SupportsModelDump +from plugins.core.validation.menus import PluginMenuConflictItem + from . import PluginPayloadBuilder from .base import PluginPayloadModel - -if TYPE_CHECKING: - from collections.abc import Mapping - - from plugins.core.lifecycle.migration import PluginMigrationResult - from plugins.core.lifecycle.seed import PluginSeedResult - from plugins.core.runtime.hooks import PluginHookResult - from plugins.core.types import SupportsModelDump - from plugins.core.validation.menus import PluginMenuConflictItem - - from .plan import ActionPayload, VersionStatePayload - from .validation import MenuConflictItemPayload +from .plan import ActionPayload, VersionStatePayload +from .validation import MenuConflictItemPayload class SupportsOk(Protocol): diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/plan.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/plan.py index ae6996008..05fd29cf3 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/plan.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/plan.py @@ -1,26 +1,23 @@ -from __future__ import annotations - +from subprocess import CompletedProcess from typing import TYPE_CHECKING, Protocol, TypeAlias, cast from pydantic import Field +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.lifecycle.purge import PluginPurgePlan, PluginPurgePlanItem from plugins.core.manifest.menu_tree import PluginMenuTree +from plugins.core.validation.dependencies import DependencyCheckResult, DependencyInstallPlanItem +from plugins.core.validation.plugin_deps import ( + PluginDependencyCheckResult, + PluginDependencyPlan, + PluginDependencyPlanBlocker, + PluginDependencyPlanItem, +) from plugins.core.validation.versioning import PluginVersionComparator from .base import PluginPayloadModel if TYPE_CHECKING: - from subprocess import CompletedProcess - - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.lifecycle.purge import PluginPurgePlan, PluginPurgePlanItem - from plugins.core.validation.dependencies import DependencyCheckResult, DependencyInstallPlanItem - from plugins.core.validation.plugin_deps import ( - PluginDependencyCheckResult, - PluginDependencyPlan, - PluginDependencyPlanBlocker, - PluginDependencyPlanItem, - ) from plugins.core.validation.structure import PluginStructureCheckResult diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/purge.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/purge.py index aef5324d1..118625ef2 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/purge.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/purge.py @@ -1,15 +1,12 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, TypeAlias +from typing import TypeAlias from pydantic import Field +from plugins.core.lifecycle.purge import PluginPurgePlan + from . import PluginPayloadBuilder from .base import PluginPayloadModel -if TYPE_CHECKING: - from plugins.core.lifecycle.purge import PluginPurgePlan - class PluginPurgeStatePayload(PluginPayloadModel): """ diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/runtime.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/runtime.py index 629516048..f880db971 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/runtime.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/runtime.py @@ -1,24 +1,18 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Protocol, TypeAlias +from collections.abc import Mapping +from typing import Protocol, TypeAlias from pydantic import Field +from plugins.core.discovery.scanner import DiscoveredPlugin +from plugins.core.lifecycle.purge import PluginPurgePlan +from plugins.core.types import JSONObject +from plugins.core.validation.plugin_deps import PluginBatchOperation from utils.log_util import logger from . import PluginPayloadBuilder from .base import PluginPayloadModel - -if TYPE_CHECKING: - from collections.abc import Mapping - - from plugins.core.discovery.scanner import DiscoveredPlugin - from plugins.core.lifecycle.purge import PluginPurgePlan - from plugins.core.types import JSONObject - from plugins.core.validation.plugin_deps import PluginBatchOperation - - from .catalog import PluginMenuDiagnosticPlanPayload - from .plan import ActionPayload, VersionStatePayload +from .catalog import PluginMenuDiagnosticPlanPayload +from .plan import ActionPayload, VersionStatePayload class SupportsOk(Protocol): diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/validation.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/validation.py index e379106b1..ec86015fb 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/validation.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/payload/validation.py @@ -1,22 +1,16 @@ -from __future__ import annotations - -from typing import TYPE_CHECKING, Protocol, TypeAlias +from collections.abc import Mapping +from typing import Protocol, TypeAlias from pydantic import Field -from plugins.core.validation.result import PluginValidationLevelResolver, ValidationLevel +from plugins.core.validation.dependencies import DependencyCheckItem, DependencyCheckResult +from plugins.core.validation.menus import PluginMenuConflictItem +from plugins.core.validation.plugin_deps import PluginDependencyCheckItem +from plugins.core.validation.result import PluginValidationIssue, PluginValidationLevelResolver, ValidationLevel +from plugins.core.validation.structure import PluginStructureCheckItem from .base import PluginPayloadModel -if TYPE_CHECKING: - from collections.abc import Mapping - - from plugins.core.validation.dependencies import DependencyCheckItem, DependencyCheckResult - from plugins.core.validation.menus import PluginMenuConflictItem - from plugins.core.validation.plugin_deps import PluginDependencyCheckItem - from plugins.core.validation.result import PluginValidationIssue - from plugins.core.validation.structure import PluginStructureCheckItem - class DependencyItemPayload(PluginPayloadModel): """ diff --git a/ruoyi-fastapi-backend/plugins/core/runtime/support/precheck.py b/ruoyi-fastapi-backend/plugins/core/runtime/support/precheck.py index 6d79ab37d..607bba152 100644 --- a/ruoyi-fastapi-backend/plugins/core/runtime/support/precheck.py +++ b/ruoyi-fastapi-backend/plugins/core/runtime/support/precheck.py @@ -1,9 +1,8 @@ -from __future__ import annotations - from dataclasses import dataclass -from typing import TYPE_CHECKING, TypeAlias +from typing import TypeAlias from pydantic import Field +from typing_extensions import Self from plugins.core.runtime.support.payload.validation import ( DependencyItemPayload, @@ -12,18 +11,15 @@ StructureItemPayload, ValidationIssuePayload, ) +from plugins.core.validation.dependencies import DependencyCheckResult +from plugins.core.validation.manifest import PluginManifestCheckResult +from plugins.core.validation.menus import PluginMenuConflictResult +from plugins.core.validation.plugin_deps import PluginDependencyCheckResult +from plugins.core.validation.structure import PluginStructureCheckResult from .payload import PluginPayloadBuilder from .payload.base import PluginPayloadModel -if TYPE_CHECKING: - from plugins.core.validation.dependencies import DependencyCheckResult - from plugins.core.validation.manifest import PluginManifestCheckResult - from plugins.core.validation.menus import PluginMenuConflictResult - from plugins.core.validation.plugin_deps import PluginDependencyCheckResult - from plugins.core.validation.structure import PluginStructureCheckResult - - PrecheckOperationPayloadDict: TypeAlias = dict[ str, bool @@ -123,7 +119,7 @@ def build( plugin_dependency_result: PluginDependencyCheckResult, structure_result: PluginStructureCheckResult, menu_conflict_result: PluginMenuConflictResult, - ) -> PluginPrecheckContext: + ) -> Self: """ 从各类检查结果构建预检上下文。 diff --git a/ruoyi-fastapi-backend/plugins/core/types.py b/ruoyi-fastapi-backend/plugins/core/types.py index 353b13b20..1116c59ff 100644 --- a/ruoyi-fastapi-backend/plugins/core/types.py +++ b/ruoyi-fastapi-backend/plugins/core/types.py @@ -1,5 +1,3 @@ -from __future__ import annotations - from typing import Any, Protocol, TypeAlias from pydantic import JsonValue as PydanticJsonValue diff --git a/ruoyi-fastapi-backend/plugins/core/validation/__init__.py b/ruoyi-fastapi-backend/plugins/core/validation/__init__.py index 1de9edba1..e69de29bb 100644 --- a/ruoyi-fastapi-backend/plugins/core/validation/__init__.py +++ b/ruoyi-fastapi-backend/plugins/core/validation/__init__.py @@ -1,3 +0,0 @@ -""" -插件校验能力分层包。 -""" diff --git a/ruoyi-fastapi-backend/server.py b/ruoyi-fastapi-backend/server.py index 1555b9591..526eebeb8 100644 --- a/ruoyi-fastapi-backend/server.py +++ b/ruoyi-fastapi-backend/server.py @@ -14,6 +14,7 @@ from exceptions.handle import handle_exception from middlewares.handle import handle_middleware from module_admin.service.log_service import LogAggregatorService +from module_identity.service.runtime_service import OidcRuntimeService from plugins.core.runtime.application import get_plugin_application_runtime from sub_applications.handle import handle_sub_applications from utils.common_util import worship @@ -31,6 +32,7 @@ async def _start_background_tasks(app: FastAPI) -> None: """ await SchedulerManager.init_system_scheduler(app.state.redis) app.state.log_aggregator_task = asyncio.create_task(LogAggregatorService.consume_stream(app.state.redis)) + await OidcRuntimeService.start_background_tasks(app) async def _stop_background_tasks(app: FastAPI) -> None: @@ -48,6 +50,7 @@ async def _stop_background_tasks(app: FastAPI) -> None: await log_task except asyncio.CancelledError: pass + await OidcRuntimeService.stop_background_tasks(app) finally: try: redis = getattr(app.state, 'redis', None) @@ -78,9 +81,17 @@ async def _initialize_application_runtime(app: FastAPI, application_leader: bool stage='platform', log_success_enabled=application_leader, ) + try: + await OidcRuntimeService.refresh_cors_snapshot(app) + except Exception: + app.state.oidc_registered_cors_origins = () + logger.error('OIDC 已注册跨域来源快照加载失败') + await OidcRuntimeService.validate_runtime(app) async def create_plugin_entity_tables() -> None: - """在插件 writer 导入实体后同步插件表。""" + """ + 在插件 writer 导入实体后同步插件表。 + """ await init_create_table( stage='plugin_entities', log_success_enabled=True, @@ -98,6 +109,7 @@ async def create_plugin_entity_tables() -> None: ) await RedisUtil.init_sys_dict(app.state.redis) await RedisUtil.init_sys_config(app.state.redis) + app.state.application_leader = application_leader await _start_background_tasks(app) @@ -126,7 +138,14 @@ def _log_address_group( *, path: str = '', ) -> None: - """输出一组本地和网络访问地址。""" + """ + 输出一组本地和网络访问地址。 + + :param title: 地址分组标题。 + :param local_ip: 本机地址。 + :param network_ips: 网络地址列表。 + :param path: 地址路径。 + """ port = AppConfig.app_port links = [f'🏠 Local: http://{local_ip}:{port}{path}'] links.extend(f'📡 Network: http://{ip}:{port}{path}' for ip in network_ips) diff --git a/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql b/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql index e8bcdd1ce..71d0a9166 100644 --- a/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql +++ b/ruoyi-fastapi-backend/sql/ruoyi-fastapi-pg.sql @@ -262,6 +262,7 @@ insert into sys_menu values(107, '通知公告', 1, '8', 'notice', insert into sys_menu values(108, '日志管理', 1, '9', 'log', '', '', '', 1, 0, 'M', '0', '0', '', 'log', 'admin', current_timestamp, '', null, '日志管理菜单'); insert into sys_menu values(119, '文件管理', 1, '10', 'file', 'system/file/index', '', '', 1, 0, 'C', '0', '0', 'system:file:list', 'documentation', 'admin', current_timestamp, '', null, '文件管理菜单'); insert into sys_menu values(120, '插件管理', 1, '11', 'plugin', 'system/plugin/index', '', '', 1, 0, 'C', '0', '0', 'system:plugin:list', 'component', 'admin', current_timestamp, '', null, '插件管理菜单'); +insert into sys_menu values(121, '认证中心', 1, '12', 'oauth', '', '', '', 1, 0, 'M', '0', '0', '', 'oauth', 'admin', current_timestamp, '', null, '统一认证中心管理'); insert into sys_menu values(109, '在线用户', 2, '1', 'online', 'monitor/online/index', '', '', 1, 0, 'C', '0', '0', 'monitor:online:list', 'online', 'admin', current_timestamp, '', null, '在线用户菜单'); insert into sys_menu values(110, '定时任务', 2, '2', 'job', 'monitor/job/index', '', '', 1, 0, 'C', '0', '0', 'monitor:job:list', 'job', 'admin', current_timestamp, '', null, '定时任务菜单'); insert into sys_menu values(111, '数据监控', 2, '3', 'druid', 'monitor/druid/index', '', '', 1, 0, 'C', '0', '0', 'monitor:druid:list', 'druid', 'admin', current_timestamp, '', null, '数据监控菜单'); @@ -269,12 +270,19 @@ insert into sys_menu values(112, '服务监控', 2, '4', 'server', insert into sys_menu values(113, '缓存监控', 2, '5', 'cache', 'monitor/cache/index', '', '', 1, 0, 'C', '0', '0', 'monitor:cache:list', 'redis', 'admin', current_timestamp, '', null, '缓存监控菜单'); insert into sys_menu values(114, '缓存列表', 2, '6', 'cacheList', 'monitor/cache/list', '', '', 1, 0, 'C', '0', '0', 'monitor:cache:list', 'redis-list', 'admin', current_timestamp, '', null, '缓存列表菜单'); insert into sys_menu values(118, '传输加密', 2, '7', 'transportCrypto', 'monitor/transportCrypto/index', '', '', 1, 0, 'C', '0', '0', 'monitor:transportCrypto:list', 'chart', 'admin', current_timestamp, '', null, '传输加密监控菜单'); +insert into sys_menu values(122, 'OAuth审计', 2, '8', 'oauthAudit', 'monitor/oauthAudit/index', '', '', 1, 0, 'C', '0', '0', 'monitor:oauthAudit:list', 'form', 'admin', current_timestamp, '', null, 'OAuth 审计日志'); insert into sys_menu values(115, '表单构建', 3, '1', 'build', 'tool/build/index', '', '', 1, 0, 'C', '0', '0', 'tool:build:list', 'build', 'admin', current_timestamp, '', null, '表单构建菜单'); insert into sys_menu values(116, '代码生成', 3, '2', 'gen', 'tool/gen/index', '', '', 1, 0, 'C', '0', '0', 'tool:gen:list', 'code', 'admin', current_timestamp, '', null, '代码生成菜单'); insert into sys_menu values(117, '系统接口', 3, '3', 'swagger', 'tool/swagger/index', '', '', 1, 0, 'C', '0', '0', 'tool:swagger:list', 'swagger', 'admin', current_timestamp, '', null, '系统接口菜单'); -- 三级菜单 insert into sys_menu values(500, '操作日志', 108, '1', 'operlog', 'monitor/operlog/index', '', '', 1, 0, 'C', '0', '0', 'monitor:operlog:list', 'form', 'admin', current_timestamp, '', null, '操作日志菜单'); insert into sys_menu values(501, '登录日志', 108, '2', 'logininfor', 'monitor/logininfor/index', '', '', 1, 0, 'C', '0', '0', 'monitor:logininfor:list', 'logininfor', 'admin', current_timestamp, '', null, '登录日志菜单'); +insert into sys_menu values(502, '客户端管理', 121, '1', 'client', 'system/oauth/client/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthClient:list', 'client', 'admin', current_timestamp, '', null, 'OAuth Client 管理'); +insert into sys_menu values(503, '资源管理', 121, '2', 'resource', 'system/oauth/resource/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthResource:list', 'resource', 'admin', current_timestamp, '', null, 'OAuth Resource 管理'); +insert into sys_menu values(504, '范围管理', 121, '3', 'scope', 'system/oauth/scope/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthScope:list', 'scope', 'admin', current_timestamp, '', null, 'OAuth Scope 管理'); +insert into sys_menu values(505, '外部会话', 121, '4', 'session', 'system/oauth/session/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthSession:list', 'session', 'admin', current_timestamp, '', null, 'OIDC SSO Session 管理'); +insert into sys_menu values(506, '外部授权', 121, '5', 'grant', 'system/oauth/grant/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthGrant:list', 'grant', 'admin', current_timestamp, '', null, 'OAuth Grant 管理'); +insert into sys_menu values(507, '签名密钥', 121, '6', 'key', 'system/oauth/key/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthKey:list', 'key', 'admin', current_timestamp, '', null, 'OIDC Key 管理'); -- 用户管理按钮 insert into sys_menu values(1000, '用户查询', 100, '1', '', '', '', '', 1, 0, 'F', '0', '0', 'system:user:query', '#', 'admin', current_timestamp, '', null, ''); insert into sys_menu values(1001, '用户新增', 100, '2', '', '', '', '', 1, 0, 'F', '0', '0', 'system:user:add', '#', 'admin', current_timestamp, '', null, ''); @@ -345,6 +353,24 @@ insert into sys_menu values(1042, '登录查询', 501, '1', '#', '', '', '', 1, insert into sys_menu values(1043, '登录删除', 501, '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:logininfor:remove', '#', 'admin', current_timestamp, '', null, ''); insert into sys_menu values(1044, '日志导出', 501, '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:logininfor:export', '#', 'admin', current_timestamp, '', null, ''); insert into sys_menu values(1045, '账户解锁', 501, '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:logininfor:unlock', '#', 'admin', current_timestamp, '', null, ''); +-- 认证中心管理按钮 +insert into sys_menu values(1073, '客户端查询', 502, '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:query', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1074, '客户端新增', 502, '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:add', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1075, '客户端修改', 502, '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:edit', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1076, '客户端删除', 502, '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:remove', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1077, '密钥轮换', 502, '5', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:rotateSecret', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1079, '资源新增', 503, '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthResource:add', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1080, '资源修改', 503, '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthResource:edit', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1081, '资源删除', 503, '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthResource:remove', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1083, '范围新增', 504, '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthScope:add', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1084, '范围修改', 504, '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthScope:edit', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1085, '范围删除', 504, '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthScope:remove', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1086, '会话下线', 505, '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthSession:revoke', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1087, '授权撤销', 506, '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthGrant:revoke', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1088, '密钥轮换', 507, '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthKey:rotate', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1089, '密钥启用', 507, '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthKey:activate', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1090, '密钥退役', 507, '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthKey:retire', '#', 'admin', current_timestamp, '', null, ''); +insert into sys_menu values(1091, '审计导出', 122, '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:oauthAudit:export', '#', 'admin', current_timestamp, '', null, ''); -- 在线用户按钮 insert into sys_menu values(1046, '在线查询', 109, '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:online:query', '#', 'admin', current_timestamp, '', null, ''); insert into sys_menu values(1047, '批量强退', 109, '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:online:batchLogout', '#', 'admin', current_timestamp, '', null, ''); @@ -1620,6 +1646,662 @@ comment on column sys_plugin_operation_log.result is '完整执行结果JSON'; comment on column sys_plugin_operation_log.create_time is '创建时间'; comment on column sys_plugin_operation_log.remark is '备注'; +-- ---------------------------- +-- 统一认证中心相关表清理 +-- ---------------------------- +drop table if exists sys_oauth_audit_archive; +drop table if exists sys_oauth_audit_log; +drop table if exists sys_oidc_signing_key; +drop table if exists sys_oauth_refresh_token; +drop table if exists sys_sso_session_client; +drop table if exists sys_sso_session; +drop table if exists sys_oauth_access_policy; +drop table if exists sys_oauth_grant; +drop table if exists sys_oauth_client_resource; +drop table if exists sys_oauth_client_scope; +drop table if exists sys_oauth_scope; +drop table if exists sys_oauth_resource; +drop table if exists sys_oauth_client_uri; +drop table if exists sys_oauth_client_secret; +drop table if exists sys_oauth_client; +drop table if exists sys_identity_subject; + +-- ---------------------------- +-- 36、统一认证主体关联表 +-- ---------------------------- +create table sys_identity_subject ( + identity_id bigserial not null, + user_id bigint not null, + subject_id varchar(36) not null, + auth_version bigint not null default 1, + create_by varchar(64) default null, + create_time timestamp(3) with time zone not null, + update_by varchar(64) default null, + update_time timestamp(3) with time zone default null, + primary key (identity_id), + constraint uk_identity_subject_user unique (user_id), + constraint uk_identity_subject_subject unique (subject_id), + constraint fk_identity_subject_user foreign key (user_id) references sys_user (user_id) on delete restrict +); +create index idx_identity_subject_auth_version on sys_identity_subject (auth_version); +comment on table sys_identity_subject is '统一认证主体关联表'; +comment on column sys_identity_subject.identity_id is '内部主键'; +comment on column sys_identity_subject.user_id is '本地用户ID'; +comment on column sys_identity_subject.subject_id is 'OIDC Subject'; +comment on column sys_identity_subject.auth_version is '认证安全版本'; +comment on column sys_identity_subject.create_by is '创建者'; +comment on column sys_identity_subject.create_time is '创建时间'; +comment on column sys_identity_subject.update_by is '更新者'; +comment on column sys_identity_subject.update_time is '更新时间'; + +-- ---------------------------- +-- 初始化-统一认证主体关联表数据 +-- ---------------------------- +with seeded_users as materialized ( + select user_id, + overlay( + overlay(md5(user_id::text || ':' || clock_timestamp()::text || ':' || random()::text) + placing '4' from 13 for 1) + placing '8' from 17 for 1 + ) as uuid_seed + from sys_user +) +insert into sys_identity_subject (user_id, subject_id, auth_version, create_by, create_time) +select user_id, + substr(uuid_seed, 1, 8) || '-' || substr(uuid_seed, 9, 4) || '-' || + substr(uuid_seed, 13, 4) || '-' || substr(uuid_seed, 17, 4) || '-' || + substr(uuid_seed, 21, 12), + 1, + 'initial-sql', + current_timestamp +from seeded_users; + +-- ---------------------------- +-- 37、OAuth客户端表 +-- ---------------------------- +create table sys_oauth_client ( + client_pk bigserial not null, + client_id varchar(64) not null, + client_name varchar(100) not null, + client_type varchar(20) not null, + token_endpoint_auth_method varchar(32) not null, + grant_types jsonb not null, + response_types jsonb not null, + subject_type varchar(16) not null default 'public', + require_pkce smallint not null default 1, + require_consent smallint not null default 1, + trusted_client smallint not null default 0, + policy_version bigint not null default 1, + id_token_signed_response_alg varchar(16) not null default 'RS256', + access_token_ttl_seconds int4 default null, + refresh_token_idle_seconds int4 default null, + refresh_token_absolute_seconds int4 default null, + logo_uri varchar(500) default null, + policy_uri varchar(500) default null, + tos_uri varchar(500) default null, + backchannel_logout_session_required smallint not null default 1, + status char(1) not null default '0', + create_by varchar(64) not null default '', + create_time timestamp(3) with time zone not null, + update_by varchar(64) not null default '', + update_time timestamp(3) with time zone not null, + remark varchar(500) default null, + primary key (client_pk), + constraint uk_oauth_client_client_id unique (client_id) +); +create index idx_oauth_client_status on sys_oauth_client (status); +comment on table sys_oauth_client is 'OAuth客户端表'; +comment on column sys_oauth_client.client_pk is '内部主键'; +comment on column sys_oauth_client.client_id is 'Client ID'; +comment on column sys_oauth_client.client_name is '客户端名称'; +comment on column sys_oauth_client.client_type is 'Client 类型'; +comment on column sys_oauth_client.token_endpoint_auth_method is 'Token 端点认证方式'; +comment on column sys_oauth_client.grant_types is 'Grant Type 列表'; +comment on column sys_oauth_client.response_types is 'Response Type 列表'; +comment on column sys_oauth_client.subject_type is 'Subject 类型'; +comment on column sys_oauth_client.require_pkce is '是否要求 PKCE'; +comment on column sys_oauth_client.require_consent is '是否要求同意'; +comment on column sys_oauth_client.trusted_client is '是否受信任 Client'; +comment on column sys_oauth_client.policy_version is '安全策略版本'; +comment on column sys_oauth_client.id_token_signed_response_alg is 'ID Token 算法'; +comment on column sys_oauth_client.access_token_ttl_seconds is 'Access Token 有效期'; +comment on column sys_oauth_client.refresh_token_idle_seconds is 'Refresh Token 闲置有效期'; +comment on column sys_oauth_client.refresh_token_absolute_seconds is 'Refresh Token 绝对有效期'; +comment on column sys_oauth_client.logo_uri is 'Logo URI'; +comment on column sys_oauth_client.policy_uri is '隐私政策 URI'; +comment on column sys_oauth_client.tos_uri is '服务条款 URI'; +comment on column sys_oauth_client.backchannel_logout_session_required is '是否要求 Back-Channel Session'; +comment on column sys_oauth_client.status is '状态(0正常 1停用)'; +comment on column sys_oauth_client.create_by is '创建者'; +comment on column sys_oauth_client.create_time is '创建时间'; +comment on column sys_oauth_client.update_by is '更新者'; +comment on column sys_oauth_client.update_time is '更新时间'; +comment on column sys_oauth_client.remark is '备注'; + +-- ---------------------------- +-- 38、OAuth客户端密钥表 +-- ---------------------------- +create table sys_oauth_client_secret ( + secret_id varchar(36) not null, + client_pk bigint not null, + secret_hash varchar(100) not null, + secret_hint varchar(12) not null, + status varchar(16) not null default 'active', + not_before timestamp(3) with time zone not null, + expires_at timestamp(3) with time zone default null, + last_used_at timestamp(3) with time zone default null, + create_by varchar(64) not null, + create_time timestamp(3) with time zone not null, + revoked_by varchar(64) default null, + revoked_at timestamp(3) with time zone default null, + primary key (secret_id), + constraint fk_oauth_client_secret_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +); +create index idx_oauth_client_secret_client on sys_oauth_client_secret (client_pk, status); +comment on table sys_oauth_client_secret is 'OAuth客户端密钥表'; +comment on column sys_oauth_client_secret.secret_id is 'Secret ID'; +comment on column sys_oauth_client_secret.client_pk is 'Client 主键'; +comment on column sys_oauth_client_secret.secret_hash is 'Secret 强哈希'; +comment on column sys_oauth_client_secret.secret_hint is 'Secret 提示'; +comment on column sys_oauth_client_secret.status is 'Secret 状态'; +comment on column sys_oauth_client_secret.not_before is '生效时间'; +comment on column sys_oauth_client_secret.expires_at is '过期时间'; +comment on column sys_oauth_client_secret.last_used_at is '最近使用时间'; +comment on column sys_oauth_client_secret.create_by is '创建者'; +comment on column sys_oauth_client_secret.create_time is '创建时间'; +comment on column sys_oauth_client_secret.revoked_by is '撤销者'; +comment on column sys_oauth_client_secret.revoked_at is '撤销时间'; + +-- ---------------------------- +-- 39、OAuth客户端URI表 +-- ---------------------------- +create table sys_oauth_client_uri ( + uri_id bigserial not null, + client_pk bigint not null, + uri_type varchar(32) not null, + uri varchar(1000) not null, + uri_hash char(64) not null, + is_default smallint not null default 0, + status char(1) not null default '0', + create_time timestamp(3) with time zone not null, + primary key (uri_id), + constraint uk_oauth_client_uri_hash unique (client_pk, uri_type, uri_hash), + constraint fk_oauth_client_uri_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +); +create index idx_oauth_client_uri_type on sys_oauth_client_uri (client_pk, uri_type, status); +comment on table sys_oauth_client_uri is 'OAuth客户端URI表'; +comment on column sys_oauth_client_uri.uri_id is 'URI 主键'; +comment on column sys_oauth_client_uri.client_pk is 'Client 主键'; +comment on column sys_oauth_client_uri.uri_type is 'URI 类型'; +comment on column sys_oauth_client_uri.uri is '精确 URI'; +comment on column sys_oauth_client_uri.uri_hash is 'URI SHA-256 摘要'; +comment on column sys_oauth_client_uri.is_default is '是否默认 URI'; +comment on column sys_oauth_client_uri.status is '状态(0正常 1停用)'; +comment on column sys_oauth_client_uri.create_time is '创建时间'; + +-- ---------------------------- +-- 40、OAuth资源服务器表 +-- ---------------------------- +create table sys_oauth_resource ( + resource_pk bigserial not null, + resource_id varchar(64) not null, + resource_name varchar(100) not null, + audience varchar(500) not null, + token_format varchar(16) not null default 'jwt', + signing_alg varchar(16) not null default 'RS256', + access_token_ttl_seconds int4 default null, + introspection_client_pk bigint default null, + allowed_claims jsonb not null, + status char(1) not null default '0', + create_by varchar(64) not null, + create_time timestamp(3) with time zone not null, + update_by varchar(64) not null, + update_time timestamp(3) with time zone not null, + remark varchar(500) default null, + primary key (resource_pk), + constraint uk_oauth_resource_resource_id unique (resource_id), + constraint uk_oauth_resource_audience unique (audience), + constraint fk_oauth_resource_introspection_client foreign key (introspection_client_pk) references sys_oauth_client (client_pk) on delete restrict +); +create index idx_oauth_resource_status on sys_oauth_resource (status); +comment on table sys_oauth_resource is 'OAuth资源服务器表'; +comment on column sys_oauth_resource.resource_pk is '内部主键'; +comment on column sys_oauth_resource.resource_id is 'Resource ID'; +comment on column sys_oauth_resource.resource_name is 'Resource 名称'; +comment on column sys_oauth_resource.audience is 'Access Token audience'; +comment on column sys_oauth_resource.token_format is 'Token 格式'; +comment on column sys_oauth_resource.signing_alg is '签名算法'; +comment on column sys_oauth_resource.access_token_ttl_seconds is 'Access Token 有效期'; +comment on column sys_oauth_resource.introspection_client_pk is 'Introspection Client 主键'; +comment on column sys_oauth_resource.allowed_claims is '允许的 Claims'; +comment on column sys_oauth_resource.status is '状态(0正常 1停用)'; +comment on column sys_oauth_resource.create_by is '创建者'; +comment on column sys_oauth_resource.create_time is '创建时间'; +comment on column sys_oauth_resource.update_by is '更新者'; +comment on column sys_oauth_resource.update_time is '更新时间'; +comment on column sys_oauth_resource.remark is '备注'; + +-- ---------------------------- +-- 41、OAuth权限范围表 +-- ---------------------------- +create table sys_oauth_scope ( + scope_pk bigserial not null, + scope_code varchar(100) not null, + scope_name varchar(100) not null, + scope_type varchar(16) not null, + resource_pk bigint default null, + claims jsonb not null, + consent_required smallint not null default 1, + sensitive smallint not null default 0, + status char(1) not null default '0', + create_by varchar(64) not null, + create_time timestamp(3) with time zone not null, + update_by varchar(64) not null, + update_time timestamp(3) with time zone not null, + remark varchar(500) default null, + primary key (scope_pk), + constraint uk_oauth_scope_code unique (scope_code), + constraint fk_oauth_scope_resource foreign key (resource_pk) references sys_oauth_resource (resource_pk) on delete restrict +); +create index idx_oauth_scope_status on sys_oauth_scope (status); +create index idx_oauth_scope_resource on sys_oauth_scope (resource_pk); +alter sequence sys_oauth_scope_scope_pk_seq restart 8; +comment on table sys_oauth_scope is 'OAuth权限范围表'; +comment on column sys_oauth_scope.scope_pk is '内部主键'; +comment on column sys_oauth_scope.scope_code is 'Scope 编码'; +comment on column sys_oauth_scope.scope_name is 'Scope 名称'; +comment on column sys_oauth_scope.scope_type is 'Scope 类型'; +comment on column sys_oauth_scope.resource_pk is 'Resource 主键'; +comment on column sys_oauth_scope.claims is 'Claims 列表'; +comment on column sys_oauth_scope.consent_required is '是否需要同意'; +comment on column sys_oauth_scope.sensitive is '是否敏感'; +comment on column sys_oauth_scope.status is '状态(0正常 1停用)'; +comment on column sys_oauth_scope.create_by is '创建者'; +comment on column sys_oauth_scope.create_time is '创建时间'; +comment on column sys_oauth_scope.update_by is '更新者'; +comment on column sys_oauth_scope.update_time is '更新时间'; +comment on column sys_oauth_scope.remark is '备注'; + +-- ---------------------------- +-- 初始化-OAuth权限范围表数据 +-- ---------------------------- +insert into sys_oauth_scope values(1, 'openid', 'OpenID', 'identity', null, '["sub"]'::jsonb, 1, 0, '0', 'system', current_timestamp, 'system', current_timestamp, 'OIDC 必需身份范围'); +insert into sys_oauth_scope values(2, 'profile', '基础资料', 'identity', null, '["name", "preferred_username", "picture", "updated_at"]'::jsonb, 1, 0, '0', 'system', current_timestamp, 'system', current_timestamp, 'OIDC Profile'); +insert into sys_oauth_scope values(3, 'email', '邮箱', 'identity', null, '["email", "email_verified"]'::jsonb, 1, 1, '0', 'system', current_timestamp, 'system', current_timestamp, 'OIDC Email'); +insert into sys_oauth_scope values(4, 'phone', '手机号', 'identity', null, '["phone_number", "phone_number_verified"]'::jsonb, 1, 1, '0', 'system', current_timestamp, 'system', current_timestamp, 'OIDC Phone'); +insert into sys_oauth_scope values(5, 'roles', '角色', 'identity', null, '["roles"]'::jsonb, 1, 1, '0', 'system', current_timestamp, 'system', current_timestamp, '外部角色 Claim'); +insert into sys_oauth_scope values(6, 'dept', '部门', 'identity', null, '["dept_id", "dept_name"]'::jsonb, 1, 1, '0', 'system', current_timestamp, 'system', current_timestamp, '外部部门 Claim'); +insert into sys_oauth_scope values(7, 'offline_access', '离线访问', 'identity', null, '[]'::jsonb, 1, 1, '0', 'system', current_timestamp, 'system', current_timestamp, '允许签发 Refresh Token'); + +-- ---------------------------- +-- 42、OAuth客户端和权限范围关联表 +-- ---------------------------- +create table sys_oauth_client_scope ( + client_pk bigint not null, + scope_pk bigint not null, + is_default smallint not null default 0, + pre_authorized smallint not null default 0, + claim_filter jsonb default null, + create_time timestamp(3) with time zone not null, + primary key (client_pk, scope_pk), + constraint fk_oauth_client_scope_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict, + constraint fk_oauth_client_scope_scope foreign key (scope_pk) references sys_oauth_scope (scope_pk) on delete restrict +); +create index idx_oauth_client_scope_scope on sys_oauth_client_scope (scope_pk); +comment on table sys_oauth_client_scope is 'OAuth客户端和权限范围关联表'; +comment on column sys_oauth_client_scope.client_pk is 'Client 主键'; +comment on column sys_oauth_client_scope.scope_pk is 'Scope 主键'; +comment on column sys_oauth_client_scope.is_default is '是否默认 Scope'; +comment on column sys_oauth_client_scope.pre_authorized is '是否预授权'; +comment on column sys_oauth_client_scope.claim_filter is 'Client Claim 过滤策略'; +comment on column sys_oauth_client_scope.create_time is '创建时间'; + +-- ---------------------------- +-- 43、OAuth客户端和资源服务器关联表 +-- ---------------------------- +create table sys_oauth_client_resource ( + client_pk bigint not null, + resource_pk bigint not null, + is_default smallint not null default 0, + create_time timestamp(3) with time zone not null, + primary key (client_pk, resource_pk), + constraint fk_oauth_client_resource_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict, + constraint fk_oauth_client_resource_resource foreign key (resource_pk) references sys_oauth_resource (resource_pk) on delete restrict +); +create index idx_oauth_client_resource_resource on sys_oauth_client_resource (resource_pk); +comment on table sys_oauth_client_resource is 'OAuth客户端和资源服务器关联表'; +comment on column sys_oauth_client_resource.client_pk is 'Client 主键'; +comment on column sys_oauth_client_resource.resource_pk is 'Resource 主键'; +comment on column sys_oauth_client_resource.is_default is '是否默认 Resource'; +comment on column sys_oauth_client_resource.create_time is '创建时间'; + +-- ---------------------------- +-- 44、用户应用访问控制表 +-- ---------------------------- +create table sys_oauth_access_policy ( + user_id bigint not null, + client_pk bigint not null, + access_status varchar(16) not null default 'allowed', + reason varchar(200), + update_by varchar(64) not null, + update_time timestamp(3) with time zone not null, + primary key (user_id, client_pk), + constraint fk_oauth_access_user foreign key (user_id) references sys_user (user_id) on delete restrict, + constraint fk_oauth_access_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +); +comment on table sys_oauth_access_policy is 'OAuth用户应用访问控制表'; +comment on column sys_oauth_access_policy.user_id is '用户ID'; +comment on column sys_oauth_access_policy.client_pk is 'Client 主键'; +comment on column sys_oauth_access_policy.access_status is 'allowed允许 blocked禁止'; +comment on column sys_oauth_access_policy.reason is '访问控制原因'; +comment on column sys_oauth_access_policy.update_by is '操作人'; +comment on column sys_oauth_access_policy.update_time is '操作时间'; + +-- ---------------------------- +-- 45、OAuth授权记录表 +-- ---------------------------- +create table sys_oauth_grant ( + grant_id varchar(36) not null, + user_id bigint not null, + subject_id varchar(36) not null, + client_pk bigint not null, + granted_scopes jsonb not null, + granted_resources jsonb not null, + remembered_scopes jsonb default null, + remembered_resources jsonb default null, + client_policy_version bigint not null, + status varchar(16) not null default 'active', + consented_at timestamp(3) with time zone not null, + expires_at timestamp(3) with time zone default null, + revoked_at timestamp(3) with time zone default null, + revoke_reason varchar(200) default null, + last_used_at timestamp(3) with time zone default null, + primary key (grant_id), + constraint fk_oauth_grant_user foreign key (user_id) references sys_user (user_id) on delete restrict, + constraint fk_oauth_grant_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +); +create index idx_oauth_grant_user on sys_oauth_grant (user_id, status); +create index idx_oauth_grant_client on sys_oauth_grant (client_pk, status); +create index idx_oauth_grant_user_client on sys_oauth_grant (user_id, client_pk, status); +comment on table sys_oauth_grant is 'OAuth授权记录表'; +comment on column sys_oauth_grant.grant_id is 'Grant ID'; +comment on column sys_oauth_grant.user_id is '用户ID'; +comment on column sys_oauth_grant.subject_id is 'Subject 快照'; +comment on column sys_oauth_grant.client_pk is 'Client 主键'; +comment on column sys_oauth_grant.granted_scopes is '已同意 Scope'; +comment on column sys_oauth_grant.granted_resources is '已同意 Resource audience'; +comment on column sys_oauth_grant.client_policy_version is 'Client Policy Version'; +comment on column sys_oauth_grant.status is 'Grant 状态'; +comment on column sys_oauth_grant.consented_at is '同意时间'; +comment on column sys_oauth_grant.expires_at is '过期时间'; +comment on column sys_oauth_grant.revoked_at is '撤销时间'; +comment on column sys_oauth_grant.revoke_reason is '撤销原因'; +comment on column sys_oauth_grant.last_used_at is '最近使用时间'; + +-- ---------------------------- +-- 46、OIDC单点登录会话表 +-- ---------------------------- +create table sys_sso_session ( + sid varchar(36) not null, + session_secret_hash char(64) not null, + user_id bigint not null, + subject_id varchar(36) not null, + auth_version bigint not null, + auth_time timestamp(3) with time zone not null, + last_seen_at timestamp(3) with time zone not null, + idle_expires_at timestamp(3) with time zone not null, + absolute_expires_at timestamp(3) with time zone not null, + acr varchar(100) not null, + amr jsonb not null, + remember_me smallint not null default 0, + ip_address varchar(128) default null, + user_agent_hash char(64) default null, + status varchar(16) not null default 'active', + revoked_at timestamp(3) with time zone default null, + revoke_reason varchar(200) default null, + create_time timestamp(3) with time zone not null, + primary key (sid), + constraint fk_sso_session_user foreign key (user_id) references sys_user (user_id) on delete restrict +); +create index idx_sso_session_user on sys_sso_session (user_id, status); +create index idx_sso_session_idle on sys_sso_session (status, idle_expires_at); +create index idx_sso_session_absolute on sys_sso_session (status, absolute_expires_at); +comment on table sys_sso_session is 'OIDC单点登录会话表'; +comment on column sys_sso_session.sid is 'OIDC Session ID'; +comment on column sys_sso_session.session_secret_hash is 'SSO Cookie 摘要'; +comment on column sys_sso_session.user_id is '用户ID'; +comment on column sys_sso_session.subject_id is 'Subject 快照'; +comment on column sys_sso_session.auth_version is '认证安全版本'; +comment on column sys_sso_session.auth_time is '认证时间'; +comment on column sys_sso_session.last_seen_at is '最近活动时间'; +comment on column sys_sso_session.idle_expires_at is '闲置过期时间'; +comment on column sys_sso_session.absolute_expires_at is '绝对过期时间'; +comment on column sys_sso_session.acr is '认证上下文'; +comment on column sys_sso_session.amr is '认证方式'; +comment on column sys_sso_session.remember_me is '是否长期会话'; +comment on column sys_sso_session.ip_address is '登录 IP'; +comment on column sys_sso_session.user_agent_hash is 'User-Agent 摘要'; +comment on column sys_sso_session.status is 'Session 状态'; +comment on column sys_sso_session.revoked_at is '撤销时间'; +comment on column sys_sso_session.revoke_reason is '撤销原因'; +comment on column sys_sso_session.create_time is '创建时间'; + +-- ---------------------------- +-- 47、SSO会话与参与应用关联表 +-- ---------------------------- +create table sys_sso_session_client ( + sid varchar(36) not null, + client_pk bigint not null, + create_time timestamp(3) with time zone not null, + last_used_at timestamp(3) with time zone not null, + primary key (sid, client_pk), + constraint fk_sso_session_client_sid foreign key (sid) references sys_sso_session (sid) on delete restrict, + constraint fk_sso_session_client_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +); +create index idx_sso_session_client_client on sys_sso_session_client (client_pk); +comment on table sys_sso_session_client is 'SSO会话参与应用'; +comment on column sys_sso_session_client.sid is 'SSO Session ID'; +comment on column sys_sso_session_client.client_pk is 'Client 主键'; +comment on column sys_sso_session_client.create_time is '首次授权时间'; +comment on column sys_sso_session_client.last_used_at is '最近授权时间'; + +-- ---------------------------- +-- 48、OAuth刷新令牌表 +-- ---------------------------- +create table sys_oauth_refresh_token ( + token_id varchar(36) not null, + token_hash char(64) not null, + family_id varchar(36) not null, + parent_token_id varchar(36) default null, + replaced_by_token_id varchar(36) default null, + grant_id varchar(36) not null, + user_id bigint not null, + subject_id varchar(36) not null, + auth_version bigint not null, + client_pk bigint not null, + sid varchar(36) not null, + scopes jsonb not null, + resources jsonb not null, + status varchar(24) not null default 'active', + issued_at timestamp(3) with time zone not null, + last_used_at timestamp(3) with time zone default null, + idle_expires_at timestamp(3) with time zone not null, + absolute_expires_at timestamp(3) with time zone not null, + revoked_at timestamp(3) with time zone default null, + revoke_reason varchar(200) default null, + reuse_detected_at timestamp(3) with time zone default null, + primary key (token_id), + constraint uk_oauth_refresh_token_hash unique (token_hash), + constraint fk_oauth_refresh_parent foreign key (parent_token_id) references sys_oauth_refresh_token (token_id) on delete restrict, + constraint fk_oauth_refresh_replaced_by foreign key (replaced_by_token_id) references sys_oauth_refresh_token (token_id) on delete restrict, + constraint fk_oauth_refresh_grant foreign key (grant_id) references sys_oauth_grant (grant_id) on delete restrict, + constraint fk_oauth_refresh_user foreign key (user_id) references sys_user (user_id) on delete restrict, + constraint fk_oauth_refresh_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict, + constraint fk_oauth_refresh_sid foreign key (sid) references sys_sso_session (sid) on delete restrict +); +create index idx_oauth_refresh_family on sys_oauth_refresh_token (family_id, status); +create index idx_oauth_refresh_user on sys_oauth_refresh_token (user_id, status); +create index idx_oauth_refresh_client on sys_oauth_refresh_token (client_pk, status); +create index idx_oauth_refresh_sid on sys_oauth_refresh_token (sid, status); +create index idx_oauth_refresh_expire on sys_oauth_refresh_token (status, absolute_expires_at); +comment on table sys_oauth_refresh_token is 'OAuth刷新令牌表'; +comment on column sys_oauth_refresh_token.token_id is 'Token ID'; +comment on column sys_oauth_refresh_token.token_hash is 'Token HMAC 摘要'; +comment on column sys_oauth_refresh_token.family_id is 'Token Family ID'; +comment on column sys_oauth_refresh_token.parent_token_id is '父 Token ID'; +comment on column sys_oauth_refresh_token.replaced_by_token_id is '替代 Token ID'; +comment on column sys_oauth_refresh_token.grant_id is 'Grant ID'; +comment on column sys_oauth_refresh_token.user_id is '用户ID'; +comment on column sys_oauth_refresh_token.subject_id is 'Subject 快照'; +comment on column sys_oauth_refresh_token.auth_version is '认证安全版本'; +comment on column sys_oauth_refresh_token.client_pk is 'Client 主键'; +comment on column sys_oauth_refresh_token.sid is 'SSO Session ID'; +comment on column sys_oauth_refresh_token.scopes is '绑定 Scope'; +comment on column sys_oauth_refresh_token.resources is '绑定 Resource audience'; +comment on column sys_oauth_refresh_token.status is 'Token 状态'; +comment on column sys_oauth_refresh_token.issued_at is '签发时间'; +comment on column sys_oauth_refresh_token.last_used_at is '最近使用时间'; +comment on column sys_oauth_refresh_token.idle_expires_at is '闲置过期时间'; +comment on column sys_oauth_refresh_token.absolute_expires_at is '绝对过期时间'; +comment on column sys_oauth_refresh_token.revoked_at is '撤销时间'; +comment on column sys_oauth_refresh_token.revoke_reason is '撤销原因'; +comment on column sys_oauth_refresh_token.reuse_detected_at is '重放检测时间'; + +-- ---------------------------- +-- 49、OIDC签名密钥表 +-- ---------------------------- +create table sys_oidc_signing_key ( + key_pk bigserial not null, + kid varchar(100) not null, + key_use varchar(16) not null default 'sig', + alg varchar(16) not null default 'RS256', + public_jwk jsonb not null, + private_key_ref varchar(1000) default null, + private_key_ciphertext text default null, + status varchar(16) not null, + publish_at timestamp(3) with time zone not null, + signing_start_at timestamp(3) with time zone default null, + signing_stop_at timestamp(3) with time zone default null, + remove_from_jwks_at timestamp(3) with time zone default null, + create_by varchar(64) not null, + create_time timestamp(3) with time zone not null, + remark varchar(500) default null, + primary key (key_pk), + constraint uk_oidc_signing_key_kid unique (kid), + constraint ck_oidc_signing_key_private_material check (num_nonnulls(private_key_ref, private_key_ciphertext) = 1) +); +create index idx_oidc_signing_key_status_publish on sys_oidc_signing_key (status, publish_at); +create index idx_oidc_signing_key_jwks_remove on sys_oidc_signing_key (status, remove_from_jwks_at); +comment on table sys_oidc_signing_key is 'OIDC签名密钥表'; +comment on column sys_oidc_signing_key.key_pk is '内部主键'; +comment on column sys_oidc_signing_key.kid is 'JWKS Key ID'; +comment on column sys_oidc_signing_key.key_use is 'JWK 用途'; +comment on column sys_oidc_signing_key.alg is '签名算法'; +comment on column sys_oidc_signing_key.public_jwk is '公开 JWK'; +comment on column sys_oidc_signing_key.private_key_ref is 'KMS/HSM/文件引用'; +comment on column sys_oidc_signing_key.private_key_ciphertext is '加密私钥材料'; +comment on column sys_oidc_signing_key.status is '密钥状态'; +comment on column sys_oidc_signing_key.publish_at is '发布时间'; +comment on column sys_oidc_signing_key.signing_start_at is '开始签名时间'; +comment on column sys_oidc_signing_key.signing_stop_at is '停止签名时间'; +comment on column sys_oidc_signing_key.remove_from_jwks_at is '移出 JWKS 时间'; +comment on column sys_oidc_signing_key.create_by is '创建者'; +comment on column sys_oidc_signing_key.create_time is '创建时间'; +comment on column sys_oidc_signing_key.remark is '备注'; + +-- ---------------------------- +-- 50、OAuth审计日志表 +-- ---------------------------- +create table sys_oauth_audit_log ( + event_id bigserial not null, + trace_id varchar(64) default null, + event_type varchar(64) not null, + result varchar(16) not null, + risk_level varchar(16) not null default 'normal', + client_id varchar(64) default null, + resource_id varchar(64) default null, + user_id bigint default null, + subject_id varchar(36) default null, + sid varchar(36) default null, + grant_id varchar(36) default null, + token_id varchar(36) default null, + ip_address varchar(128) default null, + user_agent varchar(500) default null, + failure_code varchar(64) default null, + detail jsonb default null, + create_time timestamp(3) with time zone not null, + primary key (event_id) +); +create index idx_oauth_audit_time on sys_oauth_audit_log (create_time); +create index idx_oauth_audit_client on sys_oauth_audit_log (client_id, create_time); +create index idx_oauth_audit_user on sys_oauth_audit_log (user_id, create_time); +create index idx_oauth_audit_event on sys_oauth_audit_log (event_type, result, create_time); +create index idx_oauth_audit_risk on sys_oauth_audit_log (risk_level, create_time); +comment on table sys_oauth_audit_log is 'OAuth审计日志表'; +comment on column sys_oauth_audit_log.event_id is '事件ID'; +comment on column sys_oauth_audit_log.trace_id is '链路追踪ID'; +comment on column sys_oauth_audit_log.event_type is '事件类型'; +comment on column sys_oauth_audit_log.result is '结果'; +comment on column sys_oauth_audit_log.risk_level is '风险等级'; +comment on column sys_oauth_audit_log.client_id is 'Client ID 快照'; +comment on column sys_oauth_audit_log.resource_id is 'Resource ID 快照'; +comment on column sys_oauth_audit_log.user_id is '用户ID快照'; +comment on column sys_oauth_audit_log.subject_id is 'Subject 快照'; +comment on column sys_oauth_audit_log.sid is 'SSO Session ID'; +comment on column sys_oauth_audit_log.grant_id is 'Grant ID'; +comment on column sys_oauth_audit_log.token_id is 'Token ID'; +comment on column sys_oauth_audit_log.ip_address is '客户端 IP'; +comment on column sys_oauth_audit_log.user_agent is '脱敏 User-Agent'; +comment on column sys_oauth_audit_log.failure_code is '失败码'; +comment on column sys_oauth_audit_log.detail is '脱敏扩展详情'; +comment on column sys_oauth_audit_log.create_time is '事件时间'; + +-- ---------------------------- +-- 51、OAuth审计归档表 +-- ---------------------------- +create table sys_oauth_audit_archive ( + event_id bigint not null, + trace_id varchar(64) default null, + event_type varchar(64) not null, + result varchar(16) not null, + risk_level varchar(16) not null, + client_id varchar(64) default null, + resource_id varchar(64) default null, + user_id bigint default null, + subject_id varchar(36) default null, + sid varchar(36) default null, + grant_id varchar(36) default null, + token_id varchar(36) default null, + ip_address varchar(128) default null, + user_agent varchar(500) default null, + failure_code varchar(64) default null, + detail jsonb default null, + create_time timestamp(3) with time zone not null, + archived_at timestamp(3) with time zone not null, + primary key (event_id) +); +create index idx_oauth_audit_archive_time on sys_oauth_audit_archive (create_time); +create index idx_oauth_audit_archive_event on sys_oauth_audit_archive (event_type, result, create_time); +comment on table sys_oauth_audit_archive is 'OAuth审计归档表'; +comment on column sys_oauth_audit_archive.event_id is '原事件ID'; +comment on column sys_oauth_audit_archive.trace_id is '链路追踪ID'; +comment on column sys_oauth_audit_archive.event_type is '事件类型'; +comment on column sys_oauth_audit_archive.result is '结果'; +comment on column sys_oauth_audit_archive.risk_level is '风险等级'; +comment on column sys_oauth_audit_archive.client_id is 'Client ID 快照'; +comment on column sys_oauth_audit_archive.resource_id is 'Resource ID 快照'; +comment on column sys_oauth_audit_archive.user_id is '用户ID快照'; +comment on column sys_oauth_audit_archive.subject_id is 'Subject 快照'; +comment on column sys_oauth_audit_archive.sid is 'SSO Session ID'; +comment on column sys_oauth_audit_archive.grant_id is 'Grant ID'; +comment on column sys_oauth_audit_archive.token_id is 'Token ID'; +comment on column sys_oauth_audit_archive.ip_address is '客户端 IP'; +comment on column sys_oauth_audit_archive.user_agent is '脱敏 User-Agent'; +comment on column sys_oauth_audit_archive.failure_code is '失败码'; +comment on column sys_oauth_audit_archive.detail is '脱敏扩展详情'; +comment on column sys_oauth_audit_archive.create_time is '事件时间'; +comment on column sys_oauth_audit_archive.archived_at is '归档时间'; + CREATE OR REPLACE FUNCTION "find_in_set"(int8, varchar) RETURNS "pg_catalog"."bool" AS $BODY$ DECLARE diff --git a/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql b/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql index 78d9e2258..186062307 100644 --- a/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql +++ b/ruoyi-fastapi-backend/sql/ruoyi-fastapi.sql @@ -179,6 +179,7 @@ insert into sys_menu values('107', '通知公告', '1', '8', 'notice', insert into sys_menu values('108', '日志管理', '1', '9', 'log', '', '', '', 1, 0, 'M', '0', '0', '', 'log', 'admin', sysdate(), '', null, '日志管理菜单'); insert into sys_menu values('119', '文件管理', '1', '10', 'file', 'system/file/index', '', '', 1, 0, 'C', '0', '0', 'system:file:list', 'documentation', 'admin', sysdate(), '', null, '文件管理菜单'); insert into sys_menu values('120', '插件管理', '1', '11', 'plugin', 'system/plugin/index', '', '', 1, 0, 'C', '0', '0', 'system:plugin:list', 'component', 'admin', sysdate(), '', null, '插件管理菜单'); +insert into sys_menu values('121', '认证中心', '1', '12', 'oauth', '', '', '', 1, 0, 'M', '0', '0', '', 'oauth', 'admin', sysdate(), '', null, '统一认证中心管理'); insert into sys_menu values('109', '在线用户', '2', '1', 'online', 'monitor/online/index', '', '', 1, 0, 'C', '0', '0', 'monitor:online:list', 'online', 'admin', sysdate(), '', null, '在线用户菜单'); insert into sys_menu values('110', '定时任务', '2', '2', 'job', 'monitor/job/index', '', '', 1, 0, 'C', '0', '0', 'monitor:job:list', 'job', 'admin', sysdate(), '', null, '定时任务菜单'); insert into sys_menu values('111', '数据监控', '2', '3', 'druid', 'monitor/druid/index', '', '', 1, 0, 'C', '0', '0', 'monitor:druid:list', 'druid', 'admin', sysdate(), '', null, '数据监控菜单'); @@ -186,12 +187,19 @@ insert into sys_menu values('112', '服务监控', '2', '4', 'server', insert into sys_menu values('113', '缓存监控', '2', '5', 'cache', 'monitor/cache/index', '', '', 1, 0, 'C', '0', '0', 'monitor:cache:list', 'redis', 'admin', sysdate(), '', null, '缓存监控菜单'); insert into sys_menu values('114', '缓存列表', '2', '6', 'cacheList', 'monitor/cache/list', '', '', 1, 0, 'C', '0', '0', 'monitor:cache:list', 'redis-list', 'admin', sysdate(), '', null, '缓存列表菜单'); insert into sys_menu values('118', '传输加密', '2', '7', 'transportCrypto', 'monitor/transportCrypto/index', '', '', 1, 0, 'C', '0', '0', 'monitor:transportCrypto:list', 'chart', 'admin', sysdate(), '', null, '传输加密监控菜单'); +insert into sys_menu values('122', 'OAuth审计', '2', '8', 'oauthAudit', 'monitor/oauthAudit/index', '', '', 1, 0, 'C', '0', '0', 'monitor:oauthAudit:list', 'form', 'admin', sysdate(), '', null, 'OAuth 审计日志'); insert into sys_menu values('115', '表单构建', '3', '1', 'build', 'tool/build/index', '', '', 1, 0, 'C', '0', '0', 'tool:build:list', 'build', 'admin', sysdate(), '', null, '表单构建菜单'); insert into sys_menu values('116', '代码生成', '3', '2', 'gen', 'tool/gen/index', '', '', 1, 0, 'C', '0', '0', 'tool:gen:list', 'code', 'admin', sysdate(), '', null, '代码生成菜单'); insert into sys_menu values('117', '系统接口', '3', '3', 'swagger', 'tool/swagger/index', '', '', 1, 0, 'C', '0', '0', 'tool:swagger:list', 'swagger', 'admin', sysdate(), '', null, '系统接口菜单'); -- 三级菜单 insert into sys_menu values('500', '操作日志', '108', '1', 'operlog', 'monitor/operlog/index', '', '', 1, 0, 'C', '0', '0', 'monitor:operlog:list', 'form', 'admin', sysdate(), '', null, '操作日志菜单'); insert into sys_menu values('501', '登录日志', '108', '2', 'logininfor', 'monitor/logininfor/index', '', '', 1, 0, 'C', '0', '0', 'monitor:logininfor:list', 'logininfor', 'admin', sysdate(), '', null, '登录日志菜单'); +insert into sys_menu values('502', '客户端管理', '121', '1', 'client', 'system/oauth/client/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthClient:list', 'client', 'admin', sysdate(), '', null, 'OAuth Client 管理'); +insert into sys_menu values('503', '资源管理', '121', '2', 'resource', 'system/oauth/resource/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthResource:list', 'resource', 'admin', sysdate(), '', null, 'OAuth Resource 管理'); +insert into sys_menu values('504', '范围管理', '121', '3', 'scope', 'system/oauth/scope/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthScope:list', 'scope', 'admin', sysdate(), '', null, 'OAuth Scope 管理'); +insert into sys_menu values('505', '外部会话', '121', '4', 'session', 'system/oauth/session/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthSession:list', 'session', 'admin', sysdate(), '', null, 'OIDC SSO Session 管理'); +insert into sys_menu values('506', '外部授权', '121', '5', 'grant', 'system/oauth/grant/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthGrant:list', 'grant', 'admin', sysdate(), '', null, 'OAuth Grant 管理'); +insert into sys_menu values('507', '签名密钥', '121', '6', 'key', 'system/oauth/key/index', '', '', 1, 0, 'C', '0', '0', 'system:oauthKey:list', 'key', 'admin', sysdate(), '', null, 'OIDC Key 管理'); -- 用户管理按钮 insert into sys_menu values('1000', '用户查询', '100', '1', '', '', '', '', 1, 0, 'F', '0', '0', 'system:user:query', '#', 'admin', sysdate(), '', null, ''); insert into sys_menu values('1001', '用户新增', '100', '2', '', '', '', '', 1, 0, 'F', '0', '0', 'system:user:add', '#', 'admin', sysdate(), '', null, ''); @@ -262,6 +270,24 @@ insert into sys_menu values('1042', '登录查询', '501', '1', '#', '', '', '', insert into sys_menu values('1043', '登录删除', '501', '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:logininfor:remove', '#', 'admin', sysdate(), '', null, ''); insert into sys_menu values('1044', '日志导出', '501', '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:logininfor:export', '#', 'admin', sysdate(), '', null, ''); insert into sys_menu values('1045', '账户解锁', '501', '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:logininfor:unlock', '#', 'admin', sysdate(), '', null, ''); +-- 认证中心管理按钮 +insert into sys_menu values('1073', '客户端查询', '502', '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:query', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1074', '客户端新增', '502', '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:add', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1075', '客户端修改', '502', '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:edit', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1076', '客户端删除', '502', '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:remove', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1077', '密钥轮换', '502', '5', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthClient:rotateSecret', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1079', '资源新增', '503', '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthResource:add', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1080', '资源修改', '503', '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthResource:edit', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1081', '资源删除', '503', '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthResource:remove', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1083', '范围新增', '504', '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthScope:add', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1084', '范围修改', '504', '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthScope:edit', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1085', '范围删除', '504', '4', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthScope:remove', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1086', '会话下线', '505', '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthSession:revoke', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1087', '授权撤销', '506', '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthGrant:revoke', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1088', '密钥轮换', '507', '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthKey:rotate', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1089', '密钥启用', '507', '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthKey:activate', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1090', '密钥退役', '507', '3', '#', '', '', '', 1, 0, 'F', '0', '0', 'system:oauthKey:retire', '#', 'admin', sysdate(), '', null, ''); +insert into sys_menu values('1091', '审计导出', '122', '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:oauthAudit:export', '#', 'admin', sysdate(), '', null, ''); -- 在线用户按钮 insert into sys_menu values('1046', '在线查询', '109', '1', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:online:query', '#', 'admin', sysdate(), '', null, ''); insert into sys_menu values('1047', '批量强退', '109', '2', '#', '', '', '', 1, 0, 'F', '0', '0', 'monitor:online:batchLogout', '#', 'admin', sysdate(), '', null, ''); @@ -1152,3 +1178,421 @@ create table sys_plugin_operation_log ( remark varchar(500) default null comment '备注', primary key (operation_id) ) engine=innodb comment = '插件批量操作审计日志表'; + +-- ---------------------------- +-- 统一认证中心相关表清理 +-- ---------------------------- +drop table if exists sys_oauth_audit_archive; +drop table if exists sys_oauth_audit_log; +drop table if exists sys_oidc_signing_key; +drop table if exists sys_oauth_refresh_token; +drop table if exists sys_sso_session_client; +drop table if exists sys_sso_session; +drop table if exists sys_oauth_access_policy; +drop table if exists sys_oauth_grant; +drop table if exists sys_oauth_client_resource; +drop table if exists sys_oauth_client_scope; +drop table if exists sys_oauth_scope; +drop table if exists sys_oauth_resource; +drop table if exists sys_oauth_client_uri; +drop table if exists sys_oauth_client_secret; +drop table if exists sys_oauth_client; +drop table if exists sys_identity_subject; + +-- ---------------------------- +-- 36、统一认证主体关联表 +-- ---------------------------- +create table sys_identity_subject ( + identity_id bigint not null auto_increment comment '内部主键', + user_id bigint not null comment '本地用户ID', + subject_id varchar(36) not null comment 'OIDC Subject', + auth_version bigint not null default 1 comment '认证安全版本', + create_by varchar(64) default null comment '创建者', + create_time datetime(3) not null comment '创建时间', + update_by varchar(64) default null comment '更新者', + update_time datetime(3) default null comment '更新时间', + primary key (identity_id), + unique key uk_identity_subject_user (user_id), + unique key uk_identity_subject_subject (subject_id), + key idx_identity_subject_auth_version (auth_version), + constraint fk_identity_subject_user foreign key (user_id) references sys_user (user_id) on delete restrict +) engine=innodb comment = '统一认证主体关联表'; + +-- ---------------------------- +-- 初始化-统一认证主体关联表数据 +-- ---------------------------- +insert into sys_identity_subject (user_id, subject_id, auth_version, create_by, create_time) +select user_id, uuid(), 1, 'initial-sql', UTC_TIMESTAMP(3) from sys_user; + +-- ---------------------------- +-- 37、OAuth客户端表 +-- ---------------------------- +create table sys_oauth_client ( + client_pk bigint not null auto_increment comment '内部主键', + client_id varchar(64) not null comment 'Client ID', + client_name varchar(100) not null comment '客户端名称', + client_type varchar(20) not null comment 'Client 类型', + token_endpoint_auth_method varchar(32) not null comment 'Token 端点认证方式', + grant_types json not null comment 'Grant Type 列表', + response_types json not null comment 'Response Type 列表', + subject_type varchar(16) not null default 'public' comment 'Subject 类型', + require_pkce smallint not null default 1 comment '是否要求 PKCE', + require_consent smallint not null default 1 comment '是否要求同意', + trusted_client smallint not null default 0 comment '是否受信任 Client', + policy_version bigint not null default 1 comment '安全策略版本', + id_token_signed_response_alg varchar(16) not null default 'RS256' comment 'ID Token 算法', + access_token_ttl_seconds int default null comment 'Access Token 有效期', + refresh_token_idle_seconds int default null comment 'Refresh Token 闲置有效期', + refresh_token_absolute_seconds int default null comment 'Refresh Token 绝对有效期', + logo_uri varchar(500) default null comment 'Logo URI', + policy_uri varchar(500) default null comment '隐私政策 URI', + tos_uri varchar(500) default null comment '服务条款 URI', + backchannel_logout_session_required smallint not null default 1 comment '是否要求 Back-Channel Session', + status char(1) not null default '0' comment '状态(0正常 1停用)', + create_by varchar(64) not null default '' comment '创建者', + create_time datetime(3) not null comment '创建时间', + update_by varchar(64) not null default '' comment '更新者', + update_time datetime(3) not null comment '更新时间', + remark varchar(500) default null comment '备注', + primary key (client_pk), + unique key uk_oauth_client_client_id (client_id), + key idx_oauth_client_status (status) +) engine=innodb comment = 'OAuth客户端表'; + +-- ---------------------------- +-- 38、OAuth客户端密钥表 +-- ---------------------------- +create table sys_oauth_client_secret ( + secret_id varchar(36) not null comment 'Secret ID', + client_pk bigint not null comment 'Client 主键', + secret_hash varchar(100) not null comment 'Secret 强哈希', + secret_hint varchar(12) not null comment 'Secret 提示', + status varchar(16) not null default 'active' comment 'Secret 状态', + not_before datetime(3) not null comment '生效时间', + expires_at datetime(3) default null comment '过期时间', + last_used_at datetime(3) default null comment '最近使用时间', + create_by varchar(64) not null comment '创建者', + create_time datetime(3) not null comment '创建时间', + revoked_by varchar(64) default null comment '撤销者', + revoked_at datetime(3) default null comment '撤销时间', + primary key (secret_id), + key idx_oauth_client_secret_client (client_pk, status), + constraint fk_oauth_client_secret_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +) engine=innodb comment = 'OAuth客户端密钥表'; + +-- ---------------------------- +-- 39、OAuth客户端URI表 +-- ---------------------------- +create table sys_oauth_client_uri ( + uri_id bigint not null auto_increment comment 'URI 主键', + client_pk bigint not null comment 'Client 主键', + uri_type varchar(32) not null comment 'URI 类型', + uri varchar(1000) not null comment '精确 URI', + uri_hash char(64) not null comment 'URI SHA-256 摘要', + is_default smallint not null default 0 comment '是否默认 URI', + status char(1) not null default '0' comment '状态(0正常 1停用)', + create_time datetime(3) not null comment '创建时间', + primary key (uri_id), + unique key uk_oauth_client_uri_hash (client_pk, uri_type, uri_hash), + key idx_oauth_client_uri_type (client_pk, uri_type, status), + constraint fk_oauth_client_uri_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +) engine=innodb comment = 'OAuth客户端URI表'; + +-- ---------------------------- +-- 40、OAuth资源服务器表 +-- ---------------------------- +create table sys_oauth_resource ( + resource_pk bigint not null auto_increment comment '内部主键', + resource_id varchar(64) not null comment 'Resource ID', + resource_name varchar(100) not null comment 'Resource 名称', + audience varchar(500) not null comment 'Access Token audience', + token_format varchar(16) not null default 'jwt' comment 'Token 格式', + signing_alg varchar(16) not null default 'RS256' comment '签名算法', + access_token_ttl_seconds int default null comment 'Access Token 有效期', + introspection_client_pk bigint default null comment 'Introspection Client 主键', + allowed_claims json not null comment '允许的 Claims', + status char(1) not null default '0' comment '状态(0正常 1停用)', + create_by varchar(64) not null comment '创建者', + create_time datetime(3) not null comment '创建时间', + update_by varchar(64) not null comment '更新者', + update_time datetime(3) not null comment '更新时间', + remark varchar(500) default null comment '备注', + primary key (resource_pk), + unique key uk_oauth_resource_resource_id (resource_id), + unique key uk_oauth_resource_audience (audience), + key idx_oauth_resource_status (status), + constraint fk_oauth_resource_introspection_client foreign key (introspection_client_pk) references sys_oauth_client (client_pk) on delete restrict +) engine=innodb comment = 'OAuth资源服务器表'; + +-- ---------------------------- +-- 41、OAuth权限范围表 +-- ---------------------------- +create table sys_oauth_scope ( + scope_pk bigint not null auto_increment comment '内部主键', + scope_code varchar(100) not null comment 'Scope 编码', + scope_name varchar(100) not null comment 'Scope 名称', + scope_type varchar(16) not null comment 'Scope 类型', + resource_pk bigint default null comment 'Resource 主键', + claims json not null comment 'Claims 列表', + consent_required smallint not null default 1 comment '是否需要同意', + `sensitive` smallint not null default 0 comment '是否敏感', + status char(1) not null default '0' comment '状态(0正常 1停用)', + create_by varchar(64) not null comment '创建者', + create_time datetime(3) not null comment '创建时间', + update_by varchar(64) not null comment '更新者', + update_time datetime(3) not null comment '更新时间', + remark varchar(500) default null comment '备注', + primary key (scope_pk), + unique key uk_oauth_scope_code (scope_code), + key idx_oauth_scope_status (status), + key idx_oauth_scope_resource (resource_pk), + constraint fk_oauth_scope_resource foreign key (resource_pk) references sys_oauth_resource (resource_pk) on delete restrict +) engine=innodb comment = 'OAuth权限范围表'; + +-- ---------------------------- +-- 初始化-OAuth权限范围表数据 +-- ---------------------------- +insert into sys_oauth_scope values(1, 'openid', 'OpenID', 'identity', null, json_array('sub'), 1, 0, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), 'OIDC 必需身份范围'); +insert into sys_oauth_scope values(2, 'profile', '基础资料', 'identity', null, json_array('name', 'preferred_username', 'picture', 'updated_at'), 1, 0, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), 'OIDC Profile'); +insert into sys_oauth_scope values(3, 'email', '邮箱', 'identity', null, json_array('email', 'email_verified'), 1, 1, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), 'OIDC Email'); +insert into sys_oauth_scope values(4, 'phone', '手机号', 'identity', null, json_array('phone_number', 'phone_number_verified'), 1, 1, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), 'OIDC Phone'); +insert into sys_oauth_scope values(5, 'roles', '角色', 'identity', null, json_array('roles'), 1, 1, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), '外部角色 Claim'); +insert into sys_oauth_scope values(6, 'dept', '部门', 'identity', null, json_array('dept_id', 'dept_name'), 1, 1, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), '外部部门 Claim'); +insert into sys_oauth_scope values(7, 'offline_access', '离线访问', 'identity', null, json_array(), 1, 1, '0', 'system', UTC_TIMESTAMP(3), 'system', UTC_TIMESTAMP(3), '允许签发 Refresh Token'); + +-- ---------------------------- +-- 42、OAuth客户端和权限范围关联表 +-- ---------------------------- +create table sys_oauth_client_scope ( + client_pk bigint not null comment 'Client 主键', + scope_pk bigint not null comment 'Scope 主键', + is_default smallint not null default 0 comment '是否默认 Scope', + pre_authorized smallint not null default 0 comment '是否预授权', + claim_filter json default null comment 'Client Claim 过滤策略', + create_time datetime(3) not null comment '创建时间', + primary key (client_pk, scope_pk), + key idx_oauth_client_scope_scope (scope_pk), + constraint fk_oauth_client_scope_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict, + constraint fk_oauth_client_scope_scope foreign key (scope_pk) references sys_oauth_scope (scope_pk) on delete restrict +) engine=innodb comment = 'OAuth客户端和权限范围关联表'; + +-- ---------------------------- +-- 43、OAuth客户端和资源服务器关联表 +-- ---------------------------- +create table sys_oauth_client_resource ( + client_pk bigint not null comment 'Client 主键', + resource_pk bigint not null comment 'Resource 主键', + is_default smallint not null default 0 comment '是否默认 Resource', + create_time datetime(3) not null comment '创建时间', + primary key (client_pk, resource_pk), + key idx_oauth_client_resource_resource (resource_pk), + constraint fk_oauth_client_resource_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict, + constraint fk_oauth_client_resource_resource foreign key (resource_pk) references sys_oauth_resource (resource_pk) on delete restrict +) engine=innodb comment = 'OAuth客户端和资源服务器关联表'; + +-- ---------------------------- +-- 44、用户应用访问控制表 +-- ---------------------------- +create table sys_oauth_access_policy ( + user_id bigint not null comment '用户ID', + client_pk bigint not null comment 'Client 主键', + access_status varchar(16) not null default 'allowed' comment 'allowed允许 blocked禁止', + reason varchar(200) default null comment '访问控制原因', + update_by varchar(64) not null comment '操作人', + update_time datetime(3) not null comment '操作时间', + primary key (user_id, client_pk), + constraint fk_oauth_access_user foreign key (user_id) references sys_user (user_id) on delete restrict, + constraint fk_oauth_access_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +) engine=innodb comment = 'OAuth用户应用访问控制表'; + +-- ---------------------------- +-- 45、OAuth授权记录表 +-- ---------------------------- +create table sys_oauth_grant ( + grant_id varchar(36) not null comment 'Grant ID', + user_id bigint not null comment '用户ID', + subject_id varchar(36) not null comment 'Subject 快照', + client_pk bigint not null comment 'Client 主键', + granted_scopes json not null comment '已同意 Scope', + granted_resources json not null comment '已同意 Resource audience', + remembered_scopes json default null comment '后续可免确认的 Scope', + remembered_resources json default null comment '后续可免确认的 Resource audience', + client_policy_version bigint not null comment 'Client Policy Version', + status varchar(16) not null default 'active' comment 'Grant 状态', + consented_at datetime(3) not null comment '同意时间', + expires_at datetime(3) default null comment '过期时间', + revoked_at datetime(3) default null comment '撤销时间', + revoke_reason varchar(200) default null comment '撤销原因', + last_used_at datetime(3) default null comment '最近使用时间', + primary key (grant_id), + key idx_oauth_grant_user (user_id, status), + key idx_oauth_grant_client (client_pk, status), + key idx_oauth_grant_user_client (user_id, client_pk, status), + constraint fk_oauth_grant_user foreign key (user_id) references sys_user (user_id) on delete restrict, + constraint fk_oauth_grant_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +) engine=innodb comment = 'OAuth授权记录表'; + +-- ---------------------------- +-- 46、OIDC单点登录会话表 +-- ---------------------------- +create table sys_sso_session ( + sid varchar(36) not null comment 'OIDC Session ID', + session_secret_hash char(64) not null comment 'SSO Cookie 摘要', + user_id bigint not null comment '用户ID', + subject_id varchar(36) not null comment 'Subject 快照', + auth_version bigint not null comment '认证安全版本', + auth_time datetime(3) not null comment '认证时间', + last_seen_at datetime(3) not null comment '最近活动时间', + idle_expires_at datetime(3) not null comment '闲置过期时间', + absolute_expires_at datetime(3) not null comment '绝对过期时间', + acr varchar(100) not null comment '认证上下文', + amr json not null comment '认证方式', + remember_me smallint not null default 0 comment '是否长期会话', + ip_address varchar(128) default null comment '登录 IP', + user_agent_hash char(64) default null comment 'User-Agent 摘要', + status varchar(16) not null default 'active' comment 'Session 状态', + revoked_at datetime(3) default null comment '撤销时间', + revoke_reason varchar(200) default null comment '撤销原因', + create_time datetime(3) not null comment '创建时间', + primary key (sid), + key idx_sso_session_user (user_id, status), + key idx_sso_session_idle (status, idle_expires_at), + key idx_sso_session_absolute (status, absolute_expires_at), + constraint fk_sso_session_user foreign key (user_id) references sys_user (user_id) on delete restrict +) engine=innodb comment = 'OIDC单点登录会话表'; + +-- ---------------------------- +-- 47、SSO会话与参与应用关联表 +-- ---------------------------- +create table sys_sso_session_client ( + sid varchar(36) not null comment 'SSO Session ID', + client_pk bigint not null comment 'Client 主键', + create_time datetime(3) not null comment '首次授权时间', + last_used_at datetime(3) not null comment '最近授权时间', + primary key (sid, client_pk), + key idx_sso_session_client_client (client_pk), + constraint fk_sso_session_client_sid foreign key (sid) references sys_sso_session (sid) on delete restrict, + constraint fk_sso_session_client_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict +) engine=innodb comment = 'SSO会话参与应用'; + +-- ---------------------------- +-- 48、OAuth刷新令牌表 +-- ---------------------------- +create table sys_oauth_refresh_token ( + token_id varchar(36) not null comment 'Token ID', + token_hash char(64) not null comment 'Token HMAC 摘要', + family_id varchar(36) not null comment 'Token Family ID', + parent_token_id varchar(36) default null comment '父 Token ID', + replaced_by_token_id varchar(36) default null comment '替代 Token ID', + grant_id varchar(36) not null comment 'Grant ID', + user_id bigint not null comment '用户ID', + subject_id varchar(36) not null comment 'Subject 快照', + auth_version bigint not null comment '认证安全版本', + client_pk bigint not null comment 'Client 主键', + sid varchar(36) not null comment 'SSO Session ID', + scopes json not null comment '绑定 Scope', + resources json not null comment '绑定 Resource audience', + status varchar(24) not null default 'active' comment 'Token 状态', + issued_at datetime(3) not null comment '签发时间', + last_used_at datetime(3) default null comment '最近使用时间', + idle_expires_at datetime(3) not null comment '闲置过期时间', + absolute_expires_at datetime(3) not null comment '绝对过期时间', + revoked_at datetime(3) default null comment '撤销时间', + revoke_reason varchar(200) default null comment '撤销原因', + reuse_detected_at datetime(3) default null comment '重放检测时间', + primary key (token_id), + unique key uk_oauth_refresh_token_hash (token_hash), + key idx_oauth_refresh_family (family_id, status), + key idx_oauth_refresh_user (user_id, status), + key idx_oauth_refresh_client (client_pk, status), + key idx_oauth_refresh_sid (sid, status), + key idx_oauth_refresh_expire (status, absolute_expires_at), + constraint fk_oauth_refresh_parent foreign key (parent_token_id) references sys_oauth_refresh_token (token_id) on delete restrict, + constraint fk_oauth_refresh_replaced_by foreign key (replaced_by_token_id) references sys_oauth_refresh_token (token_id) on delete restrict, + constraint fk_oauth_refresh_grant foreign key (grant_id) references sys_oauth_grant (grant_id) on delete restrict, + constraint fk_oauth_refresh_user foreign key (user_id) references sys_user (user_id) on delete restrict, + constraint fk_oauth_refresh_client foreign key (client_pk) references sys_oauth_client (client_pk) on delete restrict, + constraint fk_oauth_refresh_sid foreign key (sid) references sys_sso_session (sid) on delete restrict +) engine=innodb comment = 'OAuth刷新令牌表'; + +-- ---------------------------- +-- 49、OIDC签名密钥表 +-- ---------------------------- +create table sys_oidc_signing_key ( + key_pk bigint not null auto_increment comment '内部主键', + kid varchar(100) not null comment 'JWKS Key ID', + key_use varchar(16) not null default 'sig' comment 'JWK 用途', + alg varchar(16) not null default 'RS256' comment '签名算法', + public_jwk json not null comment '公开 JWK', + private_key_ref varchar(1000) default null comment 'KMS/HSM/文件引用', + private_key_ciphertext text default null comment '加密私钥材料', + status varchar(16) not null comment '密钥状态', + publish_at datetime(3) not null comment '发布时间', + signing_start_at datetime(3) default null comment '开始签名时间', + signing_stop_at datetime(3) default null comment '停止签名时间', + remove_from_jwks_at datetime(3) default null comment '移出 JWKS 时间', + create_by varchar(64) not null comment '创建者', + create_time datetime(3) not null comment '创建时间', + remark varchar(500) default null comment '备注', + primary key (key_pk), + unique key uk_oidc_signing_key_kid (kid), + key idx_oidc_signing_key_status_publish (status, publish_at), + key idx_oidc_signing_key_jwks_remove (status, remove_from_jwks_at), + constraint ck_oidc_signing_key_private_material check (((case when private_key_ref is null then 0 else 1 end) + (case when private_key_ciphertext is null then 0 else 1 end)) = 1) +) engine=innodb comment = 'OIDC签名密钥表'; + +-- ---------------------------- +-- 50、OAuth审计日志表 +-- ---------------------------- +create table sys_oauth_audit_log ( + event_id bigint not null auto_increment comment '事件ID', + trace_id varchar(64) default null comment '链路追踪ID', + event_type varchar(64) not null comment '事件类型', + result varchar(16) not null comment '结果', + risk_level varchar(16) not null default 'normal' comment '风险等级', + client_id varchar(64) default null comment 'Client ID 快照', + resource_id varchar(64) default null comment 'Resource ID 快照', + user_id bigint default null comment '用户ID快照', + subject_id varchar(36) default null comment 'Subject 快照', + sid varchar(36) default null comment 'SSO Session ID', + grant_id varchar(36) default null comment 'Grant ID', + token_id varchar(36) default null comment 'Token ID', + ip_address varchar(128) default null comment '客户端 IP', + user_agent varchar(500) default null comment '脱敏 User-Agent', + failure_code varchar(64) default null comment '失败码', + detail json default null comment '脱敏扩展详情', + create_time datetime(3) not null comment '事件时间', + primary key (event_id), + key idx_oauth_audit_time (create_time), + key idx_oauth_audit_client (client_id, create_time), + key idx_oauth_audit_user (user_id, create_time), + key idx_oauth_audit_event (event_type, result, create_time), + key idx_oauth_audit_risk (risk_level, create_time) +) engine=innodb comment = 'OAuth审计日志表'; + +-- ---------------------------- +-- 51、OAuth审计归档表 +-- ---------------------------- +create table sys_oauth_audit_archive ( + event_id bigint not null comment '原事件ID', + trace_id varchar(64) default null comment '链路追踪ID', + event_type varchar(64) not null comment '事件类型', + result varchar(16) not null comment '结果', + risk_level varchar(16) not null comment '风险等级', + client_id varchar(64) default null comment 'Client ID 快照', + resource_id varchar(64) default null comment 'Resource ID 快照', + user_id bigint default null comment '用户ID快照', + subject_id varchar(36) default null comment 'Subject 快照', + sid varchar(36) default null comment 'SSO Session ID', + grant_id varchar(36) default null comment 'Grant ID', + token_id varchar(36) default null comment 'Token ID', + ip_address varchar(128) default null comment '客户端 IP', + user_agent varchar(500) default null comment '脱敏 User-Agent', + failure_code varchar(64) default null comment '失败码', + detail json default null comment '脱敏扩展详情', + create_time datetime(3) not null comment '事件时间', + archived_at datetime(3) not null comment '归档时间', + primary key (event_id), + key idx_oauth_audit_archive_time (create_time), + key idx_oauth_audit_archive_event (event_type, result, create_time) +) engine=innodb comment = 'OAuth审计归档表'; diff --git a/ruoyi-fastapi-backend/tests/cli/root/test_bootstrap.py b/ruoyi-fastapi-backend/tests/cli/root/test_bootstrap.py index efb95481b..c13565b80 100644 --- a/ruoyi-fastapi-backend/tests/cli/root/test_bootstrap.py +++ b/ruoyi-fastapi-backend/tests/cli/root/test_bootstrap.py @@ -23,8 +23,7 @@ def test_config_env_loads_app_env_from_selected_run_env() -> None: 'processAppEnv': os.environ.get('APP_ENV'), }, ensure_ascii=False)) """ - process_env = dict(os.environ) - process_env.pop('APP_ENV', None) + process_env = {key: value for key, value in os.environ.items() if key != 'APP_ENV' and not key.startswith('OIDC_')} completed = subprocess.run( [sys.executable, '-c', script, '--env', 'dockermy'], cwd=BACKEND_DIR, diff --git a/ruoyi-fastapi-backend/tests/cli/root/test_contract_app_ops.py b/ruoyi-fastapi-backend/tests/cli/root/test_contract_app_ops.py index e83797bc4..cd923225f 100644 --- a/ruoyi-fastapi-backend/tests/cli/root/test_contract_app_ops.py +++ b/ruoyi-fastapi-backend/tests/cli/root/test_contract_app_ops.py @@ -74,6 +74,7 @@ def test_app_doctor_text_output_has_stable_check_structure( assert 'database:' in completed.stdout assert 'redis:' in completed.stdout assert 'crypto:' in completed.stdout + assert 'oidc:' in completed.stdout def test_app_doctor_json_output_has_stable_contract( @@ -88,10 +89,11 @@ def test_app_doctor_json_output_has_stable_contract( assert completed.stderr == '' assert payload['env'] == 'dev' assert isinstance(payload['ok'], bool) - assert set(payload) == {'env', 'database', 'redis', 'crypto', 'ok'} + assert set(payload) == {'env', 'database', 'redis', 'crypto', 'oidc', 'ok'} assert_check_payload_contract(payload['database'], True) assert_check_payload_contract(payload['redis'], True) assert_check_payload_contract(payload['crypto'], False) + assert_check_payload_contract(payload['oidc'], False) def test_app_env_json_output_has_stable_contract( diff --git a/ruoyi-fastapi-backend/tests/cli/root/test_contract_cli.py b/ruoyi-fastapi-backend/tests/cli/root/test_contract_cli.py index b43c4dfe4..6ac66d727 100644 --- a/ruoyi-fastapi-backend/tests/cli/root/test_contract_cli.py +++ b/ruoyi-fastapi-backend/tests/cli/root/test_contract_cli.py @@ -14,6 +14,7 @@ def test_root_help_shows_commands_without_completion_options( assert 'Usage: ruoyi [OPTIONS] COMMAND [ARGS]...' in completed.stdout assert 'app' in completed.stdout assert 'db' in completed.stdout + assert 'oidc' in completed.stdout assert 'completion' in completed.stdout assert 'wizard' in completed.stdout assert 'tui' in completed.stdout @@ -22,6 +23,35 @@ def test_root_help_shows_commands_without_completion_options( assert '--show-completion' not in completed.stdout +def test_oidc_key_bootstrap_dry_run_has_stable_contract( + run_cli_command: Callable[..., subprocess.CompletedProcess[str]], + parse_json_stdout: Callable[[subprocess.CompletedProcess[str]], dict], +) -> None: + """OIDC 初始化命令支持不连接基础设施的部署预演。""" + completed = run_cli_command( + 'oidc', + 'key', + 'bootstrap', + '--env=dev', + '--output=json', + '--dry-run', + '--yes', + '--kid=release-primary', + ) + payload = parse_json_stdout(completed) + + assert completed.returncode == SUCCESS + assert payload == { + 'ok': True, + 'message': 'OIDC 签名密钥初始化预演完成,未写入数据库', + 'kid': 'release-primary', + 'created': False, + 'active': False, + 'dryRun': True, + 'env': 'dev', + } + + def test_completion_show_bash_outputs_completion_script( run_cli_command: Callable[..., subprocess.CompletedProcess[str]], ) -> None: diff --git a/ruoyi-fastapi-backend/tests/cli/root/test_guards.py b/ruoyi-fastapi-backend/tests/cli/root/test_guards.py index 5b973294c..51ad432be 100644 --- a/ruoyi-fastapi-backend/tests/cli/root/test_guards.py +++ b/ruoyi-fastapi-backend/tests/cli/root/test_guards.py @@ -30,6 +30,7 @@ def test_dangerous_command_rules_cover_expected_commands() -> None: 'config set', 'config sync-cache', 'crypto rotate', + 'oidc key bootstrap', 'job run-once', 'job pause', 'job resume', diff --git a/ruoyi-fastapi-backend/tests/config/test_database_registry.py b/ruoyi-fastapi-backend/tests/config/test_database_registry.py index 61d0f2f24..91c25b408 100644 --- a/ruoyi-fastapi-backend/tests/config/test_database_registry.py +++ b/ruoyi-fastapi-backend/tests/config/test_database_registry.py @@ -1,21 +1,17 @@ -from __future__ import annotations - +from collections.abc import AsyncGenerator from contextlib import asynccontextmanager from types import SimpleNamespace -from typing import TYPE_CHECKING from unittest.mock import AsyncMock, MagicMock, call import pytest from sqlalchemy import URL from sqlalchemy.exc import OperationalError +from typing_extensions import Self from common.aspect.db_session import DBSessionDependency, get_db_session_provider from config import database from exceptions.exception import DataSourceInitializationException, DataSourceUnavailableException -if TYPE_CHECKING: - from collections.abc import AsyncGenerator - EXPECTED_CONNECT_TIMEOUT = 7 @@ -108,7 +104,7 @@ def __init__(self, should_fail: bool = False, session_timezone: str = '+00:00') self.should_fail = should_fail self.session_timezone = session_timezone - async def __aenter__(self) -> _Begin: + async def __aenter__(self) -> Self: if self.should_fail: raise RuntimeError('password=secret') return self @@ -483,7 +479,6 @@ async def test_dispose_all_attempts_every_engine_when_disposal_fails(monkeypatch def test_dependency_provider_is_cached_per_source() -> None: - get_db_session_provider.cache_clear() assert get_db_session_provider('reporting') is get_db_session_provider('reporting') assert get_db_session_provider('reporting') is not get_db_session_provider('archive') assert DBSessionDependency('reporting').dependency is get_db_session_provider('reporting') diff --git a/ruoyi-fastapi-backend/tests/config/test_oidc_settings.py b/ruoyi-fastapi-backend/tests/config/test_oidc_settings.py new file mode 100644 index 000000000..d756d4a45 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/config/test_oidc_settings.py @@ -0,0 +1,72 @@ +import pytest +from pydantic import ValidationError + +from config.env import OidcSettings + + +def _enabled_values() -> dict[str, object]: + return { + '_env_file': None, + 'oidc_enabled': True, + 'oidc_issuer': 'https://auth.example.com', + 'oidc_public_base_url': 'https://auth.example.com', + 'oidc_token_hash_pepper': 'p' * 32, + 'oidc_active_kid': '2026-primary', + 'oidc_signing_private_key_path': '/secure/oidc-private.pem', + 'oidc_interaction_login_url': 'https://auth.example.com/auth-center/login', + 'oidc_interaction_consent_url': 'https://auth.example.com/auth-center/consent', + 'oidc_interaction_error_url': 'https://auth.example.com/auth-center/error', + } + + +def test_oidc_disabled_does_not_require_runtime_secrets() -> None: + settings = OidcSettings(_env_file=None, oidc_enabled=False, oidc_issuer='', oidc_public_base_url='') + + assert settings.oidc_enabled is False + + +def test_oidc_enabled_accepts_secure_defaults() -> None: + settings = OidcSettings(**_enabled_values()) + + assert settings.oidc_require_pkce is True + assert settings.oidc_legacy_auth_isolation_enabled is True + assert settings.oidc_sso_cookie_name.startswith('__Host-') + + +@pytest.mark.parametrize( + 'overrides', + [ + {'oidc_require_pkce': False}, + {'oidc_legacy_auth_isolation_enabled': False}, + {'oidc_token_hash_pepper': 'short'}, + {'oidc_issuer': 'https://auth.example.com/realm'}, + {'oidc_sso_cookie_name': 'ruoyi-sso'}, + {'oidc_interaction_login_url': 'https://ui.example.com/login'}, + ], +) +def test_oidc_enabled_rejects_unsafe_configuration(overrides: dict[str, object]) -> None: + with pytest.raises(ValidationError): + OidcSettings(**(_enabled_values() | overrides)) + + +def test_oidc_cors_origin_rejects_userinfo_and_non_local_http(monkeypatch: pytest.MonkeyPatch) -> None: + with pytest.raises(ValidationError): + OidcSettings(**(_enabled_values() | {'oidc_cors_allowed_origins': 'https://user:pass@client.example'})) + + monkeypatch.setenv('APP_ENV', 'prod') + with pytest.raises(ValidationError): + OidcSettings(**(_enabled_values() | {'oidc_cors_allowed_origins': 'http://client.example'})) + + monkeypatch.setenv('APP_ENV', 'dev') + with pytest.raises(ValidationError): + OidcSettings(**(_enabled_values() | {'oidc_cors_allowed_origins': 'http://client.example'})) + settings = OidcSettings(**(_enabled_values() | {'oidc_cors_allowed_origins': 'http://localhost:5173'})) + assert settings.cors_origin_list == ('http://localhost:5173',) + + +def test_oidc_pepper_cannot_reuse_signing_key_encryption_key() -> None: + values = _enabled_values() + values['oidc_signing_key_encryption_key'] = values['oidc_token_hash_pepper'] + + with pytest.raises(ValidationError): + OidcSettings(**values) diff --git a/ruoyi-fastapi-backend/tests/middlewares/test_oidc_cors_middleware.py b/ruoyi-fastapi-backend/tests/middlewares/test_oidc_cors_middleware.py new file mode 100644 index 000000000..762c8008f --- /dev/null +++ b/ruoyi-fastapi-backend/tests/middlewares/test_oidc_cors_middleware.py @@ -0,0 +1,233 @@ +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI, Response, status +from fastapi.middleware.cors import CORSMiddleware +from fastapi.testclient import TestClient + +from config.env import AppConfig, OidcConfig +from middlewares.handle import handle_middleware +from middlewares.oidc_cors_middleware import OidcCorsMiddleware + + +@pytest.fixture(autouse=True) +def _mock_snapshot_refresh(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr('middlewares.oidc_cors_middleware.OidcRuntimeService.ensure_cors_snapshot', AsyncMock()) + + +def _client(monkeypatch: pytest.MonkeyPatch) -> TestClient: + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_cors_allowed_origins', 'https://allowed.example') + app = FastAPI() + + @app.get('/.well-known/openid-configuration') + async def discovery() -> dict[str, str]: + return {'issuer': 'https://auth.example.com'} + + @app.get('/oauth2/authorize') + async def authorize() -> dict[str, bool]: + return {'ok': True} + + @app.post('/oauth2/token') + async def token() -> Response: + return Response( + content='{"access_token":"token"}', + media_type='application/json', + headers={'Vary': 'Accept-Encoding, Origin'}, + ) + + @app.get('/business') + async def business() -> dict[str, bool]: + return {'ok': True} + + app.add_middleware( + CORSMiddleware, + allow_origins=['*'], + allow_credentials=True, + allow_methods=['*'], + allow_headers=['*'], + ) + app.add_middleware(OidcCorsMiddleware) + return TestClient(app) + + +def test_authorize_disallows_cors_but_discovery_remains_public(monkeypatch: pytest.MonkeyPatch) -> None: + client = _client(monkeypatch) + origin = {'Origin': 'https://external.example'} + + discovery = client.get('/.well-known/openid-configuration', headers=origin) + authorize = client.get('/oauth2/authorize', headers=origin) + + assert discovery.status_code == status.HTTP_200_OK + assert discovery.headers['access-control-allow-origin'] == origin['Origin'] + assert authorize.status_code == status.HTTP_403_FORBIDDEN + assert 'access-control-allow-origin' not in authorize.headers + + +def test_token_preflight_requires_explicit_origin(monkeypatch: pytest.MonkeyPatch) -> None: + client = _client(monkeypatch) + allowed = client.options( + '/oauth2/token', + headers={ + 'Origin': 'https://allowed.example', + 'Access-Control-Request-Method': 'POST', + 'Access-Control-Request-Headers': 'Authorization, Content-Type', + }, + ) + denied = client.options( + '/oauth2/token', + headers={ + 'Origin': 'https://denied.example', + 'Access-Control-Request-Method': 'POST', + }, + ) + + assert allowed.status_code == status.HTTP_200_OK + assert allowed.headers['access-control-allow-origin'] == 'https://allowed.example' + assert denied.status_code == status.HTTP_403_FORBIDDEN + assert 'access-control-allow-origin' not in denied.headers + + +def test_token_form_post_from_disallowed_origin_is_rejected_before_downstream( + monkeypatch: pytest.MonkeyPatch, +) -> None: + response = _client(monkeypatch).post( + '/oauth2/token', + data={'grant_type': 'authorization_code', 'code': 'opaque-code'}, + headers={'Origin': 'https://denied.example', 'Content-Type': 'application/x-www-form-urlencoded'}, + ) + + assert response.status_code == status.HTTP_403_FORBIDDEN + assert 'access-control-allow-origin' not in response.headers + + +def test_allowed_response_preserves_non_origin_vary_tokens(monkeypatch: pytest.MonkeyPatch) -> None: + client = _client(monkeypatch) + + response = client.post('/oauth2/token', headers={'Origin': 'https://allowed.example'}) + + assert response.status_code == status.HTTP_200_OK + assert response.headers['access-control-allow-origin'] == 'https://allowed.example' + assert 'accept-encoding' in response.headers.get('vary', '').lower() + assert 'origin' in response.headers.get('vary', '').lower() + + +def test_existing_business_cors_remains_unchanged(monkeypatch: pytest.MonkeyPatch) -> None: + response = _client(monkeypatch).get('/business', headers={'Origin': 'https://current-web.example'}) + + assert response.status_code == status.HTTP_200_OK + baseline_client = _client(monkeypatch) + monkeypatch.setattr(OidcConfig, 'oidc_enabled', False) + baseline = baseline_client.get('/business', headers={'Origin': 'https://current-web.example'}) + assert response.headers['access-control-allow-origin'] == baseline.headers['access-control-allow-origin'] + assert response.headers['access-control-allow-credentials'] == baseline.headers['access-control-allow-credentials'] + + +def test_runtime_registered_origin_is_exact_and_does_not_allow_similar_hosts( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """协议层注册缓存可扩展白名单,但相近 Host、路径和尾点仍拒绝。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_cors_allowed_origins', '') + app = FastAPI() + app.state.oidc_registered_cors_origins = frozenset({'https://registered.example'}) + + @app.post('/oauth2/token') + async def token() -> Response: + return Response('{}', media_type='application/json') + + app.add_middleware(OidcCorsMiddleware) + client = TestClient(app) + allowed = client.post('/oauth2/token', headers={'Origin': 'https://registered.example'}) + similar = client.post('/oauth2/token', headers={'Origin': 'https://registered.example.evil'}) + trailing = client.post('/oauth2/token', headers={'Origin': 'https://registered.example.'}) + path = client.post('/oauth2/token', headers={'Origin': 'https://registered.example/path'}) + assert allowed.status_code == status.HTTP_200_OK + assert similar.status_code == status.HTTP_403_FORBIDDEN + assert trailing.status_code == status.HTTP_403_FORBIDDEN + assert path.status_code == status.HTTP_403_FORBIDDEN + + +def test_interaction_api_requires_exact_issuer_origin(monkeypatch: pytest.MonkeyPatch) -> None: + """交互登录、改密和完成接口不继承全局宽松 CORS。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com/oidc') + app = FastAPI() + + @app.post('/auth/interaction/i-1/login') + async def login() -> Response: + return Response('{}', media_type='application/json') + + @app.get('/auth/interaction/i-1/complete') + async def complete() -> Response: + return Response('{}', media_type='application/json') + + app.add_middleware( + CORSMiddleware, + allow_origins=['*'], + allow_credentials=True, + allow_methods=['*'], + allow_headers=['*'], + ) + app.add_middleware(OidcCorsMiddleware) + client = TestClient(app) + + allowed = client.post('/auth/interaction/i-1/login', headers={'Origin': 'https://auth.example.com'}) + denied = client.post('/auth/interaction/i-1/login', headers={'Origin': 'https://evil.example'}) + port_spoof = client.post('/auth/interaction/i-1/login', headers={'Origin': 'https://auth.example.com:443'}) + scheme_spoof = client.get('/auth/interaction/i-1/complete', headers={'Origin': 'http://auth.example.com'}) + + assert allowed.status_code == status.HTTP_200_OK + assert allowed.headers['access-control-allow-origin'] == 'https://auth.example.com' + assert denied.status_code == status.HTTP_403_FORBIDDEN + assert port_spoof.status_code == status.HTTP_403_FORBIDDEN + assert scheme_spoof.status_code == status.HTTP_403_FORBIDDEN + assert '*' not in allowed.headers.get('access-control-allow-origin', '') + + +def test_interaction_preflight_rejects_cross_origin(monkeypatch: pytest.MonkeyPatch) -> None: + """交互接口的预检不能被内层宽松 CORS 伪造成功。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com') + app = FastAPI() + + @app.post('/auth/interaction/i-1/change-password') + async def change_password() -> Response: + return Response('{}') + + app.add_middleware(CORSMiddleware, allow_origins=['*'], allow_credentials=True, allow_methods=['*']) + app.add_middleware(OidcCorsMiddleware) + response = TestClient(app).options( + '/auth/interaction/i-1/change-password', + headers={'Origin': 'https://evil.example', 'Access-Control-Request-Method': 'POST'}, + ) + assert response.status_code == status.HTTP_403_FORBIDDEN + assert 'access-control-allow-origin' not in response.headers + + +def test_oidc_cors_is_bypassed_when_provider_is_disabled(monkeypatch: pytest.MonkeyPatch) -> None: + """关闭认证中心时 OIDC 路径不得被专用 CORS 中间件拦截。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', False) + app = FastAPI() + + @app.post('/oauth2/token') + async def token() -> Response: + return Response('{}', media_type='application/json') + + app.add_middleware(OidcCorsMiddleware) + response = TestClient(app).post('/oauth2/token', headers={'Origin': 'https://denied.example'}) + + assert response.status_code == status.HTTP_200_OK + + +@pytest.mark.parametrize('enabled', [False, True]) +def test_global_registration_follows_oidc_switch(monkeypatch: pytest.MonkeyPatch, enabled: bool) -> None: + """全局中间件注册应与 OIDC_ENABLED 开关保持一致。""" + monkeypatch.setattr(AppConfig, 'app_demo_mode', False) + monkeypatch.setattr(OidcConfig, 'oidc_enabled', enabled) + + app = FastAPI() + handle_middleware(app) + + registered = any(item.cls is OidcCorsMiddleware for item in app.user_middleware) + assert registered is enabled diff --git a/ruoyi-fastapi-backend/tests/middlewares/test_transport_crypto_oidc_exclusion.py b/ruoyi-fastapi-backend/tests/middlewares/test_transport_crypto_oidc_exclusion.py new file mode 100644 index 000000000..e5227f1a7 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/middlewares/test_transport_crypto_oidc_exclusion.py @@ -0,0 +1,25 @@ +import pytest + +from config.env import AppConfig, TransportCryptoConfig +from middlewares.transport_crypto_middleware import TransportCryptoMiddleware + + +def test_standard_oidc_paths_are_always_excluded_from_transport_envelope(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(TransportCryptoConfig, 'transport_crypto_exclude_paths', '') + + assert all( + TransportCryptoMiddleware._is_excluded_path(path) for path in TransportCryptoMiddleware._STANDARD_OIDC_PATHS + ) + + +def test_interaction_api_is_not_implicitly_excluded() -> None: + assert not TransportCryptoMiddleware._is_excluded_path('/auth/interaction/interaction-id/login') + + +def test_app_root_path_is_removed_before_oidc_exclusion(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(AppConfig, 'app_root_path', '/dev-api') + + normalized_path = TransportCryptoMiddleware._normalize_path('/dev-api/oauth2/token') + + assert normalized_path == '/oauth2/token' + assert TransportCryptoMiddleware._is_excluded_path(normalized_path) diff --git a/ruoyi-fastapi-backend/tests/module_admin/service/test_identity_security_integration.py b/ruoyi-fastapi-backend/tests/module_admin/service/test_identity_security_integration.py new file mode 100644 index 000000000..1cf342ae0 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_admin/service/test_identity_security_integration.py @@ -0,0 +1,140 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from module_admin.dao.role_dao import RoleDao +from module_admin.dao.user_dao import UserDao +from module_admin.entity.vo.role_vo import AddRoleModel +from module_admin.entity.vo.user_vo import AddUserModel, ResetUserModel +from module_admin.service.role_service import RoleService +from module_admin.service.user_service import UserService +from module_identity.service.identity_service import IdentitySecurityEventService, IdentitySubjectService + + +@pytest.mark.asyncio +async def test_add_user_creates_subject_before_commit(monkeypatch: pytest.MonkeyPatch) -> None: + """新增用户必须在同一会话提交前建立稳定 Subject。""" + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + monkeypatch.setattr(UserService, 'check_user_name_unique_services', AsyncMock(return_value=True)) + monkeypatch.setattr(UserDao, 'add_user_dao', AsyncMock(return_value=SimpleNamespace(user_id=7))) + create_subject = AsyncMock() + monkeypatch.setattr(IdentitySubjectService, 'create_for_new_user', create_subject) + payload = AddUserModel(userName='new-user', nickName='New User', createBy='admin') + + await UserService.add_user_services(db, payload) + + create_subject.assert_awaited_once_with(db, user_id=7, create_by='admin') + db.commit.assert_awaited_once() + db.rollback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_password_security_failure_rolls_back_user_change(monkeypatch: pytest.MonkeyPatch) -> None: + """密码安全失效失败时不得提交已写入的新密码。""" + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + monkeypatch.setattr(UserDao, 'edit_user_dao', AsyncMock()) + monkeypatch.setattr('module_admin.service.user_service.PwdUtil.get_password_hash', lambda value: f'hash:{value}') + security_event = AsyncMock(side_effect=RuntimeError('security state unavailable')) + monkeypatch.setattr(IdentitySecurityEventService, 'handle_user_event', security_event) + payload = ResetUserModel(userId=8, password='New123!', updateBy='admin') + + with pytest.raises(RuntimeError, match='security state unavailable'): + await UserService.reset_user_services(db, payload) + + security_event.assert_awaited_once_with(db, 8, 'password_changed', actor='admin') + db.commit.assert_not_awaited() + db.rollback.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_disabling_role_invalidates_members_before_commit(monkeypatch: pytest.MonkeyPatch) -> None: + """角色停用必须在角色更新事务中处理受影响用户。""" + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + monkeypatch.setattr( + RoleService, + 'role_detail_services', + AsyncMock(return_value=AddRoleModel(roleId=5, roleName='普通角色', roleKey='common', roleSort=1, status='0')), + ) + monkeypatch.setattr(RoleDao, 'edit_role_dao', AsyncMock()) + security_event = AsyncMock() + monkeypatch.setattr(IdentitySecurityEventService, 'handle_role_event', security_event) + payload = AddRoleModel( + roleId=5, + roleName='普通角色', + roleKey='common', + roleSort=1, + status='1', + type='status', + updateBy='admin', + ) + + await RoleService.edit_role_services(db, payload) + + security_event.assert_awaited_once_with(db, 5, 'role_disabled', actor='admin') + db.commit.assert_awaited_once() + db.rollback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_changing_role_key_invalidates_external_role_claims(monkeypatch: pytest.MonkeyPatch) -> None: + """角色标识变化必须让成员的旧外部角色声明失效。""" + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + monkeypatch.setattr( + RoleService, + 'role_detail_services', + AsyncMock(return_value=AddRoleModel(roleId=5, roleName='普通角色', roleKey='reader', roleSort=1, status='0')), + ) + monkeypatch.setattr(RoleDao, 'edit_role_dao', AsyncMock()) + monkeypatch.setattr(RoleDao, 'delete_role_menu_dao', AsyncMock()) + monkeypatch.setattr(RoleDao, 'add_role_menu_dao', AsyncMock()) + monkeypatch.setattr(RoleDao, 'list_role_menu_ids', AsyncMock(return_value=[])) + monkeypatch.setattr(RoleService, 'check_role_name_unique_services', AsyncMock(return_value=True)) + monkeypatch.setattr(RoleService, 'check_role_key_unique_services', AsyncMock(return_value=True)) + security_event = AsyncMock() + monkeypatch.setattr(IdentitySecurityEventService, 'handle_role_event', security_event) + payload = AddRoleModel( + roleId=5, + roleName='普通角色', + roleKey='auditor', + roleSort=1, + status='0', + menuIds=[], + updateBy='admin', + ) + + await RoleService.edit_role_services(db, payload) + + security_event.assert_awaited_once_with(db, 5, 'role_claim_changed', actor='admin') + db.commit.assert_awaited_once() + db.rollback.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_changing_role_menu_invalidates_external_claims(monkeypatch: pytest.MonkeyPatch) -> None: + """角色菜单授权变化必须让成员的旧外部声明失效。""" + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + role = AddRoleModel(roleId=5, roleName='普通角色', roleKey='reader', roleSort=1, status='0') + monkeypatch.setattr(RoleService, 'role_detail_services', AsyncMock(return_value=role)) + monkeypatch.setattr(RoleDao, 'list_role_menu_ids', AsyncMock(return_value=[10])) + monkeypatch.setattr(RoleDao, 'edit_role_dao', AsyncMock()) + monkeypatch.setattr(RoleDao, 'delete_role_menu_dao', AsyncMock()) + monkeypatch.setattr(RoleDao, 'add_role_menu_dao', AsyncMock()) + monkeypatch.setattr(RoleService, 'check_role_name_unique_services', AsyncMock(return_value=True)) + monkeypatch.setattr(RoleService, 'check_role_key_unique_services', AsyncMock(return_value=True)) + security_event = AsyncMock() + monkeypatch.setattr(IdentitySecurityEventService, 'handle_role_event', security_event) + payload = AddRoleModel( + roleId=5, + roleName='普通角色', + roleKey='reader', + roleSort=1, + status='0', + menuIds=[11], + updateBy='admin', + ) + + await RoleService.edit_role_services(db, payload) + + security_event.assert_awaited_once_with(db, 5, 'role_claim_changed', actor='admin') + db.commit.assert_awaited_once() diff --git a/ruoyi-fastapi-backend/tests/module_identity/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/architecture/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/architecture/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/architecture/test_architecture_boundaries.py b/ruoyi-fastapi-backend/tests/module_identity/architecture/test_architecture_boundaries.py new file mode 100644 index 000000000..066254220 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/architecture/test_architecture_boundaries.py @@ -0,0 +1,382 @@ +import ast +from pathlib import Path + +_BACKEND_ROOT = Path(__file__).resolve().parents[3] +_IDENTITY_ROOT = _BACKEND_ROOT / 'module_identity' +_SQL_CALL_NAMES = {'delete', 'insert', 'select', 'update'} +_DB_EXECUTION_METHODS = {'add', 'delete', 'execute', 'flush', 'scalar', 'scalars'} +_HTTP_ROUTE_METHODS = {'api_route', 'delete', 'get', 'head', 'options', 'patch', 'post', 'put'} + + +def _python_files(directory: str) -> list[Path]: + """返回指定认证模块目录下的生产 Python 文件。""" + return sorted((_IDENTITY_ROOT / directory).glob('*.py')) + + +def _parse(path: Path) -> ast.Module: + """将生产文件解析为抽象语法树。""" + return ast.parse(path.read_text(encoding='utf-8'), filename=str(path)) + + +def _imported_modules(tree: ast.Module) -> set[str]: + """收集语法树中的完整导入模块路径。""" + modules: set[str] = set() + for node in ast.walk(tree): + if isinstance(node, ast.Import): + modules.update(alias.name for alias in node.names) + elif isinstance(node, ast.ImportFrom) and node.module: + modules.add(node.module) + return modules + + +def test_controllers_do_not_depend_on_dao_or_database_models() -> None: + """Controller 只能依赖 DTO 和 Service,不得越层访问 DAO 或 DO。""" + violations: list[str] = [] + for path in _python_files('controller'): + modules = _imported_modules(_parse(path)) + forbidden = sorted( + module for module in modules if module.startswith(('module_identity.dao', 'module_identity.entity.do')) + ) + if forbidden: + violations.append(f'{path.name}: {", ".join(forbidden)}') + assert not violations, '\n'.join(violations) + + +def test_services_do_not_construct_or_execute_database_statements() -> None: + """Service 负责编排事务,SQL 构造和数据库执行必须位于 DAO。""" + violations: list[str] = [] + for path in _python_files('service'): + tree = _parse(path) + for node in ast.walk(tree): + if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id in _SQL_CALL_NAMES: + violations.append(f'{path.name}:{node.lineno}: {node.func.id}()') + if not ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr in _DB_EXECUTION_METHODS + and isinstance(node.func.value, ast.Name) + and (node.func.value.id == 'db' or node.func.value.id.endswith('_db')) + ): + continue + violations.append(f'{path.name}:{node.lineno}: {node.func.value.id}.{node.func.attr}()') + assert not violations, '\n'.join(violations) + + +def test_services_do_not_contain_controller_dependencies() -> None: + """Service 不得包含路由注册器或数据库依赖注入声明。""" + violations: list[str] = [] + for path in _python_files('service'): + tree = _parse(path) + imported_names = { + alias.name for node in ast.walk(tree) if isinstance(node, ast.ImportFrom) for alias in node.names + } + forbidden = sorted(imported_names & {'APIRouterPro', 'DBSessionDependency'}) + if forbidden: + violations.append(f'{path.name}: {", ".join(forbidden)}') + assert not violations, '\n'.join(violations) + + +def test_controller_dependencies_use_annotated_style() -> None: + """Controller 依赖必须使用项目统一的 ``Annotated`` 声明。""" + violations: list[str] = [] + for path in _python_files('controller'): + for node in ast.walk(_parse(path)): + if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + continue + defaults = [*node.args.defaults, *node.args.kw_defaults] + violations.extend( + f'{path.name}:{default.lineno}: {default.func.id}()' + for default in defaults + if isinstance(default, ast.Call) + and isinstance(default.func, ast.Name) + and default.func.id.endswith('Dependency') + ) + assert not violations, '\n'.join(violations) + + +def _route_decorators(node: ast.FunctionDef | ast.AsyncFunctionDef) -> list[ast.Call]: + """返回函数上的 HTTP 路由装饰器。""" + return [ + decorator + for decorator in node.decorator_list + if isinstance(decorator, ast.Call) + and isinstance(decorator.func, ast.Attribute) + and decorator.func.attr in _HTTP_ROUTE_METHODS + ] + + +def _doc_param_names(docstring: str) -> set[str]: + """提取 Sphinx 文档字符串中的参数名称。""" + return { + line.removeprefix(':param ').split(':', 1)[0].strip() + for line in docstring.splitlines() + if line.startswith(':param ') + } + + +def _block_docstring(path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef) -> bool: + """检查函数文档字符串是否采用项目统一的独占行格式。""" + if not node.body or not isinstance(node.body[0], ast.Expr): + return False + expression = node.body[0].value + if not isinstance(expression, ast.Constant) or not isinstance(expression.value, str): + return False + lines = path.read_text(encoding='utf-8').splitlines() + opening_index = expression.lineno - 1 + if lines[opening_index].strip() not in {'"""', "'''"}: + opening_index -= 1 + opening = lines[opening_index].strip() + closing = lines[expression.end_lineno - 1].strip() + doc_lines = expression.value.splitlines() + return opening in {'"""', "'''"} and closing == opening and len(doc_lines) > 1 and bool(doc_lines[1].strip()) + + +def _internal_doc_violations( + path: Path, + node: ast.FunctionDef | ast.AsyncFunctionDef, + first_route_lineno: int | None, +) -> list[str]: + """返回内部函数的顺序和文档违规信息。""" + violations: list[str] = [] + docstring = ast.get_docstring(node) + if node.col_offset != 0: + violations.append(f'{path.name}:{node.lineno}: 非路由函数必须位于模块级作用域') + if first_route_lineno is not None and node.lineno >= first_route_lineno: + violations.append(f'{path.name}:{node.lineno}: 非路由函数必须位于首个路由函数之前') + if docstring is None or not _block_docstring(path, node): + violations.append(f'{path.name}:{node.lineno}: 非路由函数必须使用独占行多行 docstring') + return violations + arguments = [*node.args.posonlyargs, *node.args.args, *node.args.kwonlyargs] + if node.args.vararg is not None: + arguments.append(node.args.vararg) + if node.args.kwarg is not None: + arguments.append(node.args.kwarg) + documented_params = _doc_param_names(docstring) + violations.extend( + f'{path.name}:{node.lineno}: 缺少 :param {argument.arg}:' + for argument in arguments + if argument.arg not in {'self', 'cls'} and argument.arg not in documented_params + ) + if node.returns is not None and ':return:' not in docstring: + violations.append(f'{path.name}:{node.lineno}: 缺少 :return:') + return violations + + +def _route_doc_violations( + path: Path, + node: ast.FunctionDef | ast.AsyncFunctionDef, + route_decorators: list[ast.Call], +) -> list[str]: + """返回路由函数文档和 OpenAPI 文案违规信息。""" + violations: list[str] = [] + if ast.get_docstring(node) is not None: + violations.append(f'{path.name}:{node.lineno}: 路由函数不得使用 docstring') + for decorator in route_decorators: + keywords = {keyword.arg: keyword.value for keyword in decorator.keywords if keyword.arg is not None} + summary = keywords.get('summary') + description = keywords.get('description') + if ( + not isinstance(summary, ast.Constant) + or not isinstance(summary.value, str) + or not summary.value.endswith('接口') + ): + violations.append(f'{path.name}:{node.lineno}: 路由 summary 必须是以“接口”结尾的字符串') + if ( + not isinstance(description, ast.Constant) + or not isinstance(description.value, str) + or not description.value.startswith('用于') + ): + violations.append(f'{path.name}:{node.lineno}: 路由 description 必须是以“用于”开头的字符串') + return violations + + +def test_identity_packages_do_not_reexport_implementation_details() -> None: + """包入口保持空白,调用方必须从职责明确的模块显式导入。""" + violations = [ + str(path.relative_to(_BACKEND_ROOT)) + for path in sorted(_IDENTITY_ROOT.rglob('__init__.py')) + if path.read_text(encoding='utf-8').strip() + ] + assert not violations, '\n'.join(violations) + + +def test_removed_facades_and_merged_modules_do_not_return() -> None: + """禁止重新引入旧 Facade 或已经合并的薄模块。""" + removed_paths = [ + _IDENTITY_ROOT / 'exceptions.py', + _IDENTITY_ROOT / 'controller' / 'logout_controller.py', + _IDENTITY_ROOT / 'controller' / 'userinfo_controller.py', + _IDENTITY_ROOT / 'service' / 'oauth_management_application_service.py', + _IDENTITY_ROOT / 'service' / 'authorization_flow_service.py', + _IDENTITY_ROOT / 'service' / 'token_endpoint_service.py', + _IDENTITY_ROOT / 'service' / 'oauth_audit_management_service.py', + _IDENTITY_ROOT / 'service' / 'interaction_audit_service.py', + _IDENTITY_ROOT / 'service' / 'authorization_code_service.py', + _IDENTITY_ROOT / 'service' / 'claim_service.py', + _IDENTITY_ROOT / 'service' / 'credential_authentication_service.py', + _IDENTITY_ROOT / 'service' / 'identity_security_event_service.py', + _IDENTITY_ROOT / 'service' / 'identity_subject_service.py', + _IDENTITY_ROOT / 'service' / 'interaction_completion_service.py', + _IDENTITY_ROOT / 'service' / 'interaction_consent_service.py', + _IDENTITY_ROOT / 'service' / 'interaction_flow_service.py', + _IDENTITY_ROOT / 'service' / 'interaction_login_service.py', + _IDENTITY_ROOT / 'service' / 'introspection_service.py', + _IDENTITY_ROOT / 'service' / 'logout_service.py', + _IDENTITY_ROOT / 'service' / 'oauth_client_management_service.py', + _IDENTITY_ROOT / 'service' / 'oauth_management_base.py', + _IDENTITY_ROOT / 'service' / 'oauth_resource_management_service.py', + _IDENTITY_ROOT / 'service' / 'oidc_key_management_service.py', + _IDENTITY_ROOT / 'service' / 'rate_limit_service.py', + _IDENTITY_ROOT / 'service' / 'revocation_service.py', + _IDENTITY_ROOT / 'service' / 'sso_session_service.py', + _IDENTITY_ROOT / 'service' / 'transaction_coordinator.py', + _IDENTITY_ROOT / 'service' / 'userinfo_service.py', + _BACKEND_ROOT / 'exceptions' / 'oidc_messages.py', + ] + assert not [str(path.relative_to(_BACKEND_ROOT)) for path in removed_paths if path.exists()] + + authorization_tree = _parse(_IDENTITY_ROOT / 'service' / 'authorization_service.py') + assert not [ + node.name + for node in authorization_tree.body + if isinstance(node, ast.ClassDef) and node.name == 'AuthorizationRequestService' + ] + + +def test_oidc_utility_does_not_depend_on_application_or_protocol_services() -> None: + """公共工具只处理传入的数据,不能反向加载配置、持久化、HTTP 或业务异常。""" + path = _BACKEND_ROOT / 'utils' / 'oidc_util.py' + modules = _imported_modules(_parse(path)) + forbidden = { + 'config', + 'exceptions', + 'module_identity', + 'module_admin', + 'redis', + 'sqlalchemy', + 'fastapi', + 'starlette', + } + assert not sorted(module for module in modules if module.split('.')[0] in forbidden) + + +def test_protocol_and_interaction_services_do_not_import_fastapi_http_types() -> None: + """协议与交互核心 Service 只接收领域数据,不得依赖 FastAPI HTTP 类型。""" + names = { + 'authorization_service.py', + 'consent_service.py', + 'identity_service.py', + 'interaction_service.py', + 'session_service.py', + 'token_protocol_service.py', + 'token_service.py', + } + forbidden = {'Request', 'Response', 'JSONResponse', 'RedirectResponse', 'HTTPException'} + violations: list[str] = [] + for path in _python_files('service'): + if path.name not in names: + continue + tree = _parse(path) + imported_modules = _imported_modules(tree) + http_modules = sorted(module for module in imported_modules if module.startswith(('fastapi', 'starlette'))) + imported_names = { + alias.name + for node in ast.walk(tree) + if isinstance(node, (ast.Import, ast.ImportFrom)) + for alias in node.names + } + if imported_names & forbidden or http_modules: + violations.append(f'{path.name}: names={sorted(imported_names & forbidden)}, modules={http_modules}') + assert not violations, '\n'.join(violations) + + +def test_dependencies_delegate_database_access_to_daos() -> None: + """协议依赖不得自行拼 SQL 或依赖 ORM DO,只能调用明确的 DAO。""" + path = _IDENTITY_ROOT / 'dependencies.py' + tree = _parse(path) + imported_modules = _imported_modules(tree) + forbidden_modules = { + module + for module in imported_modules + if module == 'sqlalchemy' + or module.startswith(('sqlalchemy.sql', 'sqlalchemy.orm', 'module_identity.entity.do')) + } + assert not forbidden_modules, f'{path.name}: {sorted(forbidden_modules)}' + + forbidden_calls = [ + f'{path.name}:{node.lineno}: {node.func.id}()' + for node in ast.walk(tree) + if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id in _SQL_CALL_NAMES + ] + forbidden_calls.extend( + f'{path.name}:{node.lineno}: {node.func.value.id}.{node.func.attr}()' + for node in ast.walk(tree) + if isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and isinstance(node.func.value, ast.Name) + and node.func.value.id in {'db', 'query_db'} + and node.func.attr in _DB_EXECUTION_METHODS + ) + assert not forbidden_calls, '\n'.join(forbidden_calls) + + dao_calls = { + (node.func.value.id, node.func.attr) + for node in ast.walk(tree) + if isinstance(node, ast.Call) and isinstance(node.func, ast.Attribute) and isinstance(node.func.value, ast.Name) + } + assert ('OidcKeyDao', 'get_verifying') in dao_calls + + +def test_logout_service_has_no_direct_http_client_or_socket_dependency() -> None: + """Logout Service 的协议编排不得携带 HTTP 客户端和 Socket 实现依赖。""" + path = _IDENTITY_ROOT / 'service' / 'session_service.py' + imported_modules = _imported_modules(_parse(path)) + forbidden = sorted( + module + for module in imported_modules + if module in {'socket', 'httpx', 'httpcore'} or module.startswith(('httpx.', 'httpcore.')) + ) + assert not forbidden, f'{path.name}: {forbidden}' + + +def test_admin_user_dao_and_identity_services_keep_identity_queries_in_identity_layer() -> None: + """Admin UserDao 不得出现身份专用查询名,身份 Service 也不得反向依赖它。""" + user_dao = _BACKEND_ROOT / 'module_admin' / 'dao' / 'user_dao.py' + source = user_dao.read_text(encoding='utf-8').lower() + forbidden_terms = ('identity', 'subject', 'auth_version', 'sso_session') + assert not [term for term in forbidden_terms if term in source], user_dao.name + + violations: list[str] = [] + for path in _python_files('service'): + tree = _parse(path) + modules = _imported_modules(tree) + imported_names = { + alias.name + for node in ast.walk(tree) + if isinstance(node, (ast.Import, ast.ImportFrom)) + for alias in node.names + } + if 'module_admin.dao.user_dao' in modules or 'UserDao' in imported_names: + violations.append(path.name) + assert not violations, ', '.join(violations) + + +def test_protocol_controllers_do_not_manage_database_transactions() -> None: + """协议与交互 Controller 只适配 HTTP,事务必须由核心 Service 完成。""" + violations: list[str] = [] + for path in ( + _IDENTITY_ROOT / 'controller' / 'authorization_controller.py', + _IDENTITY_ROOT / 'controller' / 'interaction_controller.py', + _IDENTITY_ROOT / 'controller' / 'token_controller.py', + ): + tree = _parse(path) + for node in ast.walk(tree): + if ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr in {'commit', 'rollback'} + ): + violations.append(f'{path.name}:{node.lineno}: {node.func.attr}()') + elif isinstance(node, ast.Name) and node.id == 'AfterCommitCoordinator': + violations.append(f'{path.name}:{node.lineno}: AfterCommitCoordinator') + assert not violations, '\n'.join(violations) diff --git a/ruoyi-fastapi-backend/tests/module_identity/architecture/test_management_architecture.py b/ruoyi-fastapi-backend/tests/module_identity/architecture/test_management_architecture.py new file mode 100644 index 000000000..86c61e9ab --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/architecture/test_management_architecture.py @@ -0,0 +1,17 @@ +import re +from pathlib import Path + +BACKEND_ROOT = Path(__file__).parents[3] +MANAGEMENT_SERVICES = ('oauth_management_service.py',) + + +def test_all_management_services_are_framework_and_app_state_free() -> None: + """所有管理 Service 不得反向依赖 FastAPI 请求或应用状态。""" + service_dir = BACKEND_ROOT / 'module_identity' / 'service' + forbidden_request = re.compile(r'\bRequest\b') + forbidden_fastapi_import = re.compile(r'^\s*(?:from|import)\s+fastapi(?:\.|\s|$)', re.MULTILINE) + for path in sorted(service_dir.glob('*management_service.py')): + source = path.read_text(encoding='utf-8') + assert forbidden_fastapi_import.search(source) is None, path.name + assert forbidden_request.search(source) is None, path.name + assert 'app.state' not in source, path.name diff --git a/ruoyi-fastapi-backend/tests/module_identity/conftest.py b/ruoyi-fastapi-backend/tests/module_identity/conftest.py new file mode 100644 index 000000000..365de8b36 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/conftest.py @@ -0,0 +1,100 @@ +import json +import os +from unittest.mock import AsyncMock + +os.environ.setdefault( + 'DB_SOURCES', + json.dumps( + { + 'primary': { + 'db_type': 'mysql', + 'db_host': 'localhost', + 'db_port': 3306, + 'db_username': 'test', + 'db_password': 'test', + 'db_database': 'test', + } + } + ), +) +os.environ.setdefault('DB_DEFAULT_SOURCE', 'primary') + +import pytest +import pytest_asyncio +from sqlalchemy.dialects.mysql import TINYINT +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlalchemy.ext.compiler import compiles + +from config.database import Base +from module_admin.entity.do.user_do import SysUser as _sys_user_model # noqa: N813, F401 +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.entity.do import ( + identity_subject_do as _identity_subject_models, # noqa: F401 +) +from module_identity.entity.do import oauth_audit_do as _oauth_audit_models # noqa: F401 +from module_identity.entity.do import oauth_client_do as _oauth_client_models # noqa: F401 +from module_identity.entity.do import oauth_grant_do as _oauth_grant_models # noqa: F401 +from module_identity.entity.do import oauth_resource_do as _oauth_resource_models # noqa: F401 +from module_identity.entity.do import oidc_key_do as _oidc_key_models # noqa: F401 +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_resource_do import SysOAuthScope + + +@compiles(TINYINT, 'sqlite') +def _compile_mysql_tinyint_for_sqlite(_type: TINYINT, _compiler: object, **_kwargs: object) -> str: + """让现有 MySQL 布尔列可在认证模块的 SQLite 测试库中建表。""" + return 'SMALLINT' + + +@pytest_asyncio.fixture +async def data_session() -> AsyncSession: + """创建只包含认证模块表的独立内存 SQLite 异步会话。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + identity_tables = [ + table + for table_name, table in Base.metadata.tables.items() + if table_name == 'sys_user' or table_name.startswith(('sys_identity_', 'sys_oauth_', 'sys_oidc_', 'sys_sso_')) + ] + async with engine.begin() as connection: + await connection.run_sync( + lambda sync_connection: Base.metadata.create_all(sync_connection, tables=identity_tables) + ) + session_factory = async_sessionmaker(engine, expire_on_commit=False) + async with session_factory() as session: + yield session + await engine.dispose() + + +@pytest.fixture +def interaction_page_metadata( + monkeypatch: pytest.MonkeyPatch, +) -> tuple[SysOAuthClient, list[SysOAuthScope]]: + """提供当前已登记的页面元数据,管理备注必须保持私有。""" + client = SysOAuthClient( + client_pk=1001, + client_id='portal-client', + client_name='示例门户', + policy_uri='https://portal.example/privacy', + remark='private-client-note', + ) + scopes = [ + SysOAuthScope( + scope_pk=1, + scope_code='openid', + scope_name='确认身份', + consent_required=0, + sensitive=0, + remark='private-scope-note', + ), + SysOAuthScope( + scope_pk=2, + scope_code='profile', + scope_name='基本资料', + consent_required=1, + sensitive=1, + remark='private-scope-note', + ), + ] + monkeypatch.setattr(OAuthClientDao, 'get_by_pk', AsyncMock(return_value=client)) + monkeypatch.setattr(OAuthClientDao, 'list_scopes', AsyncMock(return_value=scopes)) + return client, scopes diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_auth_center_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_auth_center_controller.py new file mode 100644 index 000000000..bf72f349f --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_auth_center_controller.py @@ -0,0 +1,26 @@ +import pytest +from fastapi import FastAPI +from httpx import ASGITransport, AsyncClient + +from config.env import OidcConfig +from module_identity.controller.auth_center_controller import auth_center_controller + +_HTTP_OK = 200 + + +@pytest.mark.asyncio +@pytest.mark.parametrize('enabled', [False, True]) +async def test_public_status_tracks_config_without_auth_or_runtime_dependencies( + monkeypatch: pytest.MonkeyPatch, enabled: bool +) -> None: + """未初始化 Redis、数据库及密钥的匿名请求仍可读取最新开关。""" + app = FastAPI() + app.include_router(auth_center_controller) + async with AsyncClient(transport=ASGITransport(app=app), base_url='http://test') as client: + for current in (enabled, not enabled): + monkeypatch.setattr(OidcConfig, 'oidc_enabled', current) + response = await client.get('/auth/status') + assert response.status_code == _HTTP_OK + assert response.json()['data'] == {'enabled': current} + assert response.headers['cache-control'] == 'no-store' + assert response.headers['pragma'] == 'no-cache' diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_authorization_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_authorization_controller.py new file mode 100644 index 000000000..697369835 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_authorization_controller.py @@ -0,0 +1,278 @@ +import asyncio +import time +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from typing import Any + +import pytest +from sqlalchemy import select + +from common.constant import OidcAuditEvent +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.controller import authorization_controller as controller +from module_identity.controller.authorization_controller import authorize +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.redis_keys import OidcRedisKey +from module_identity.service import infrastructure_service +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import AuthorizationService +from module_identity.service.interaction_service import InteractionService +from tests.module_identity.services.test_authorization_service import _seed_authorization_data +from tests.module_identity.support.redis_fakes import FakeRedis + +_REDIRECT_URI = 'https://portal.example/callback?tenant=one' +_CHALLENGE = 'A' * 43 +_HTTP_OK = 200 +_HTTP_NOT_FOUND = 404 +_HTTP_SEE_OTHER = 303 +_HTTP_TOO_MANY_REQUESTS = 429 + + +class _QueryParams: + """提供 Starlette QueryParams 所需的多值读取接口。""" + + def __init__(self, values: list[tuple[str, str]]) -> None: + self.values = values + + def multi_items(self) -> list[tuple[str, str]]: + return self.values + + +class _Request: + """构造不携带 Legacy 会话的认证请求。""" + + def __init__(self, values: list[tuple[str, str]], redis: FakeRedis | None = None) -> None: + self.method = 'GET' + self.query_params = _QueryParams(values) + self.cookies: dict[str, str] = {} + self.headers: dict[str, str] = {} + self.client = SimpleNamespace(host='198.51.100.10') + self.app = SimpleNamespace(state=SimpleNamespace(redis=redis or FakeRedis())) + + +class _RateRedis(FakeRedis): + """仅实现授权端点固定窗口 Lua 等效语义的 FakeRedis。""" + + async def eval(self, script: str, numkeys: int, *args: Any) -> Any: + if script != infrastructure_service.OidcRateLimiter._SCRIPT: + return await super().eval(script, numkeys, *args) + async with self.lock: + key = str(args[0]) + current = int(self.values.get(key, ('0', None))[0]) + 1 + ttl = int(args[1]) + self.values[key] = (str(current), time.monotonic() + ttl) + return [current, ttl] + + +def _values(**overrides: str) -> list[tuple[str, str]]: + """构造有效授权 Query 参数。""" + values = { + 'response_type': 'code', + 'client_id': 'authorization-client', + 'redirect_uri': _REDIRECT_URI, + 'scope': 'openid profile', + 'nonce': 'nonce-value', + 'state': 'opaque-state', + 'code_challenge': _CHALLENGE, + 'code_challenge_method': 'S256', + } + values.update(overrides) + return list(values.items()) + + +@pytest.fixture +def oidc_enabled(monkeypatch: pytest.MonkeyPatch) -> None: + """配置控制器测试使用的安全 OIDC 参数。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com') + monkeypatch.setattr(OidcConfig, 'oidc_pkce_methods', 'S256') + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', 'authorization-controller-pepper-' + 'x' * 32) + monkeypatch.setattr(OidcConfig, 'oidc_interaction_login_url', 'https://auth.example.com/auth-center/login') + monkeypatch.setattr(OidcConfig, 'oidc_interaction_consent_url', 'https://auth.example.com/auth-center/consent') + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorize_creates_real_interaction_with_fragment_csrf(data_session: Any) -> None: + """有效请求真实执行校验和 Interaction 创建,CSRF 只进入 Fragment。""" + await _seed_authorization_data(data_session) + response = await authorize(_Request(_values()), data_session) + assert response.status_code == _HTTP_SEE_OTHER + location = response.headers['location'] + assert 'interaction=' in location + assert '#csrf=' in location + assert 'csrf=' not in location.split('?', 1)[1].split('#', 1)[0] + assert response.headers['cache-control'] == 'no-store' + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorize_invalid_redirect_is_local_and_plain_pkce_is_redirectable(data_session: Any) -> None: + """未注册 Redirect 不得跳转;精确注册后 plain PKCE 才允许协议错误回跳。""" + await _seed_authorization_data(data_session) + with pytest.raises(OAuthProtocolException) as local: + await authorize(_Request(_values(redirect_uri='https://evil.example/callback')), data_session) + assert local.value.can_redirect is False + + with pytest.raises(OAuthProtocolException) as redirectable: + await authorize(_Request(_values(code_challenge_method='plain')), data_session) + assert redirectable.value.error == 'invalid_request' + assert redirectable.value.can_redirect is True + with pytest.raises(OAuthProtocolException) as unsupported: + await authorize(_Request(_values(response_type='token')), data_session) + assert unsupported.value.error == 'unsupported_response_type' + assert unsupported.value.can_redirect is True + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorize_max_age_discards_stale_sso(monkeypatch: pytest.MonkeyPatch, data_session: Any) -> None: + """超过 max_age 的有效 SSO 只能重新进入登录交互。""" + await _seed_authorization_data(data_session) + stale = SimpleNamespace( + sid='sid-stale', + user_id=2, + subject_id='11111111-1111-4111-8111-111111111111', + auth_version=1, + auth_time=datetime.now(timezone.utc) - timedelta(seconds=30), + ) + + async def load_stale(*args: object, **kwargs: object) -> Any: + return stale + + monkeypatch.setattr(AuthorizationService, '_load_sso_session', load_stale) + response = await authorize(_Request(_values(max_age='1')), data_session) + assert response.status_code == _HTTP_SEE_OTHER + assert '/auth-center/login' in response.headers['location'] + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorize_ip_limit_is_atomic_and_hashes_ip(monkeypatch: pytest.MonkeyPatch) -> None: + """授权固定窗口并发计数只使用 HMAC-IP Redis Key。""" + redis = _RateRedis() + request = _Request(_values(), redis) + request.client = SimpleNamespace(host='198.51.100.10') + await asyncio.gather(*(controller._enforce_authorization_rate_limit(request, redis) for _ in range(30))) + with pytest.raises(OAuthProtocolException) as limited: + await controller._enforce_authorization_rate_limit(request, redis) + assert limited.value.status_code == _HTTP_TOO_MANY_REQUESTS + assert limited.value.redirect_uri is None + assert all('198.51.100.10' not in key for key in redis.values) + assert any(key.startswith('oidc:rate_limit:authorize:ip:') for key in redis.values) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorize_audit_events_are_structured_and_secret_free(data_session: Any) -> None: + """授权成功事件落库且不携带协议秘密。""" + await _seed_authorization_data(data_session) + await authorize(_Request(_values()), data_session) + rows = ( + ( + await data_session.execute( + select(SysOAuthAuditLog).where(SysOAuthAuditLog.client_id == 'authorization-client') + ) + ) + .scalars() + .all() + ) + event_types = {row.event_type for row in rows} + assert OidcAuditEvent.AUTHORIZE_REQUESTED in event_types + assert OidcAuditEvent.AUTHORIZE_SUCCEEDED not in event_types + for row in rows: + assert 'nonce-value' not in str(row.detail) + assert _CHALLENGE not in str(row.detail) + assert 'opaque-state' not in str(row.detail) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorize_prompt_none_without_sso_returns_login_required_without_interaction(data_session: Any) -> None: + """prompt=none 没有新 SSO 时直接返回 login_required,不创建页面状态。""" + await _seed_authorization_data(data_session) + redis = FakeRedis() + with pytest.raises(OAuthProtocolException) as raised: + await authorize(_Request(_values(prompt='none'), redis), data_session) + assert raised.value.error == 'login_required' + assert raised.value.can_redirect is True + assert not any(key.startswith('oidc:interaction:') for key in redis.values) + + +@pytest.mark.asyncio +async def test_authorize_disabled_is_local_404(data_session: Any, monkeypatch: pytest.MonkeyPatch) -> None: + """OIDC 关闭时授权入口返回本地 404。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', False) + response = await authorize(_Request(_values()), data_session) + assert response.status_code == _HTTP_NOT_FOUND + assert response.headers['cache-control'] == 'no-store' + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_complete_code_audit_failure_cleans_code_and_marker(monkeypatch: pytest.MonkeyPatch) -> None: + """Authorization 完成阶段审计失败时授权码与 marker 均被补偿。""" + redis = FakeRedis() + interaction_id = 'authorization-completion-id' + marker = OidcRedisKey.interaction(f'{interaction_id}-completion') + await redis.set(OidcRedisKey.interaction(interaction_id), 'present', ex=60) + record = { + 'interactionId': interaction_id, + 'status': 'completed', + 'clientPk': 1, + 'clientId': 'authorization-client', + 'redirectUri': _REDIRECT_URI, + 'scopes': ['openid'], + 'resources': [], + 'state': 'state-value', + 'nonce': 'nonce-value', + 'codeChallenge': _CHALLENGE, + 'codeChallengeMethod': 'S256', + 'authenticatedSid': 'sid-1', + 'grantId': None, + } + session = SimpleNamespace( + sid='sid-1', + user_id=2, + subject_id='11111111-1111-4111-8111-111111111111', + auth_version=1, + auth_time=None, + ) + monkeypatch.setattr(InteractionService, 'get_record', lambda *_args, **_kwargs: _async(record)) + monkeypatch.setattr( + AuthorizationService, + 'verified_redirect_for_client', + lambda *_args, **_kwargs: _async(_REDIRECT_URI), + ) + monkeypatch.setattr(AuthorizationService, 'active_session', lambda *_args, **_kwargs: _async(session)) + monkeypatch.setattr(AuditService, 'record', lambda *_args, **_kwargs: _raise_async(RuntimeError('audit'))) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationService._complete_authorization(_Db(), redis, interaction_id) + assert raised.value.error == 'server_error' + assert await redis.get(marker) is None + assert not any(key.startswith('oidc:authorization_code:') for key in redis.values) + + +class _Db: + """记录授权完成测试的事务动作。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +async def _async(value: object) -> Any: + """构造异步测试结果。""" + return value + + +async def _raise_async(error: Exception) -> Any: + """构造异步异常结果。""" + raise error diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_client_ip_resolution.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_client_ip_resolution.py new file mode 100644 index 000000000..8c6a400b4 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_client_ip_resolution.py @@ -0,0 +1,153 @@ +import hashlib +import hmac +import json +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import Request, status + +from config.env import AppConfig, OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.controller.authorization_controller import _enforce_authorization_rate_limit, logout +from module_identity.controller.interaction_controller import captcha, login_endpoint +from module_identity.entity.vo.interaction_vo import CaptchaResponseModel +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.infrastructure_service import OidcRateLimiter +from module_identity.service.interaction_service import ( + CaptchaOutcome, + InteractionFlowService, + InteractionLoginOutcome, + InteractionLoginService, +) +from module_identity.service.logout_confirmation_service import LogoutConfirmationService + +_PEPPER = 'client-ip-regression-pepper-' + 'x' * 32 +_PROXY_IP = '10.0.0.2' +_USER_IP = '198.51.100.23' + + +def _request(peer: str | None, headers: dict[str, str]) -> Request: + body = json.dumps({'userName': 'alice', 'password': 'test-password'}).encode() + + async def receive() -> dict[str, object]: + return {'type': 'http.request', 'body': body, 'more_body': False} + + return Request( + { + 'type': 'http', + 'http_version': '1.1', + 'method': 'GET', + 'scheme': 'https', + 'path': '/', + 'query_string': b'', + 'headers': [ + (name.lower().encode(), value.encode()) + for name, value in {'content-type': 'application/json', **headers}.items() + ], + 'client': (peer, 12345) if peer is not None else None, + 'server': ('auth.example', 443), + 'app': SimpleNamespace(state=SimpleNamespace(redis=object())), + }, + receive, + ) + + +@pytest.fixture +def oidc_enabled(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example') + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', _PEPPER) + monkeypatch.setattr(AppConfig, 'app_trusted_proxy_ips', _PROXY_IP) + + +@pytest.fixture( + params=[ + (_PROXY_IP, 1, {'X-Forwarded-For': _USER_IP, 'X-Real-IP': '203.0.113.9'}, _USER_IP), + (_PROXY_IP, 1, {'X-Real-IP': _USER_IP}, _USER_IP), + (_USER_IP, 1, {'X-Forwarded-For': '203.0.113.9', 'X-Real-IP': '203.0.113.9'}, _USER_IP), + (_PROXY_IP, 0, {'X-Forwarded-For': _USER_IP}, _PROXY_IP), + (_PROXY_IP, 1, {}, _PROXY_IP), + ], + ids=['trusted-forwarded-for', 'trusted-real-ip', 'untrusted-forged-headers', 'proxy-disabled', 'no-proxy-header'], +) +def client_request( + request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch, oidc_enabled: None +) -> tuple[Request, str]: + peer, hops, headers, expected = request.param + monkeypatch.setattr(AppConfig, 'app_trusted_proxy_hops', hops) + return _request(peer, headers), expected + + +@pytest.mark.asyncio +async def test_login_passes_resolved_ip_to_credentials_and_session_flow( + client_request: tuple[Request, str], monkeypatch: pytest.MonkeyPatch +) -> None: + request, expected = client_request + request.scope['method'] = 'POST' + login = AsyncMock(return_value=InteractionLoginOutcome(failure_message='认证信息无效')) + monkeypatch.setattr(InteractionLoginService, 'login', login) + + await login_endpoint(request, 'interaction-1', object(), 'csrf') + + assert login.await_args.args[5] == expected + + +@pytest.mark.asyncio +async def test_captcha_uses_the_resolved_client_ip( + client_request: tuple[Request, str], monkeypatch: pytest.MonkeyPatch +) -> None: + request, expected = client_request + generate = AsyncMock(return_value=CaptchaOutcome(result=CaptchaResponseModel(captcha_enabled=False))) + monkeypatch.setattr(InteractionFlowService, 'captcha', generate) + + await captcha(request, 'interaction-1') + + generate.assert_awaited_once_with(request.app.state.redis, 'interaction-1', expected) + + +@pytest.mark.asyncio +async def test_authorize_rate_limit_hashes_the_resolved_client_ip( + client_request: tuple[Request, str], monkeypatch: pytest.MonkeyPatch +) -> None: + request, expected = client_request + enforce = AsyncMock() + monkeypatch.setattr(OidcRateLimiter, 'enforce', enforce) + + await _enforce_authorization_rate_limit(request, request.app.state.redis) + + digest = hmac.new(_PEPPER.encode(), expected.encode(), hashlib.sha256).hexdigest() + assert enforce.await_args.args[1] == OidcRedisKey.authorize_ip_rate_limit(digest) + + +@pytest.mark.asyncio +async def test_logout_rate_limit_hashes_the_resolved_client_ip( + client_request: tuple[Request, str], monkeypatch: pytest.MonkeyPatch +) -> None: + request, expected = client_request + enforce = AsyncMock() + monkeypatch.setattr(OidcRateLimiter, 'enforce', enforce) + monkeypatch.setattr(LogoutConfirmationService, 'issue', AsyncMock(return_value=('confirmation', 'nonce'))) + monkeypatch.setattr(LogoutConfirmationService, 'form_redirect_origin', AsyncMock(return_value=None)) + + response = await logout(request, object()) + + assert response.status_code == status.HTTP_200_OK + digest = hmac.new(_PEPPER.encode(), expected.encode(), hashlib.sha256).hexdigest() + assert enforce.await_args.args[1] == OidcRedisKey.logout_rate_limit(digest) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_authorization_still_rejects_unknown_peer_ip(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(AppConfig, 'app_trusted_proxy_hops', 1) + enforce = AsyncMock() + monkeypatch.setattr(OidcRateLimiter, 'enforce', enforce) + request = _request(None, {'X-Forwarded-For': _USER_IP}) + + with pytest.raises(OAuthProtocolException) as error: + await _enforce_authorization_rate_limit(request, request.app.state.redis) + + assert error.value.error == 'temporarily_unavailable' + assert error.value.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + enforce.assert_not_awaited() diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_discovery_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_discovery_controller.py new file mode 100644 index 000000000..b7decbb58 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_discovery_controller.py @@ -0,0 +1,168 @@ +from types import SimpleNamespace + +import pytest +from starlette.requests import Request + +from config.env import OidcConfig +from module_identity.controller import discovery_controller as controller +from module_identity.service.key_service import KeyServiceError + +_HTTP_OK = 200 +_HTTP_NOT_MODIFIED = 304 +_HTTP_NOT_FOUND = 404 +_HTTP_UNAVAILABLE = 503 + + +def _request(headers: dict[str, str] | None = None, host: str = 'evil.example') -> Request: + """创建带可控 Host 和缓存请求头的 Starlette 请求。""" + raw_headers = [(b'host', host.encode())] + raw_headers.extend((key.lower().encode(), value.encode()) for key, value in (headers or {}).items()) + return Request({'type': 'http', 'method': 'GET', 'path': '/', 'headers': raw_headers}) + + +def _enabled(monkeypatch: pytest.MonkeyPatch) -> OidcConfig: + """将全局配置设为启用且固定 issuer。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com') + monkeypatch.setattr(controller, 'OidcConfig', OidcConfig) + return OidcConfig + + +class _DiscoveryDb: + """提供 Discovery 动态 Scope 查询的最小显式会话替身。""" + + class _Scalars: + def all(self) -> list[str]: + return [] + + class _Result: + def scalars(self) -> '_DiscoveryDb._Scalars': + return _DiscoveryDb._Scalars() + + async def execute(self, statement: object) -> '_DiscoveryDb._Result': + return _DiscoveryDb._Result() + + +@pytest.mark.asyncio +async def test_discovery_is_raw_static_and_supports_304(monkeypatch: pytest.MonkeyPatch) -> None: + """发现响应使用静态 issuer、裸 JSON 和 ETag 304。""" + _enabled(monkeypatch) + response = await controller.openid_configuration(_request(host='attacker.example'), _DiscoveryDb()) + assert response.status_code == _HTTP_OK + assert response.media_type == 'application/json' + assert response.headers['cache-control'] == 'public, max-age=300' + assert b'attacker.example' not in response.body + not_modified = await controller.openid_configuration( + _request({'If-None-Match': response.headers['etag']}), _DiscoveryDb() + ) + assert not_modified.status_code == _HTTP_NOT_MODIFIED + assert not not_modified.body + listed = await controller.openid_configuration( + _request({'If-None-Match': '"other", ' + response.headers['etag']}), _DiscoveryDb() + ) + assert listed.status_code == _HTTP_NOT_MODIFIED + wildcard = await controller.openid_configuration(_request({'If-None-Match': '*'}), _DiscoveryDb()) + assert wildcard.status_code == _HTTP_NOT_MODIFIED + + +@pytest.mark.asyncio +async def test_oauth_metadata_does_not_claim_oidc_only_fields(monkeypatch: pytest.MonkeyPatch) -> None: + """RFC 8414 子集不要求 userinfo/end_session 等 OIDC 专属字段。""" + _enabled(monkeypatch) + response = await controller.oauth_authorization_server_metadata(_request(), _DiscoveryDb()) + assert response.status_code == _HTTP_OK + assert b'userinfo_endpoint' not in response.body + assert b'end_session_endpoint' not in response.body + assert b'client_credentials' in response.body + + +@pytest.mark.asyncio +async def test_oidc_metadata_declares_implemented_machine_and_backchannel_flows( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """OIDC Discovery 声明当前已实现的机器授权和 Back-Channel 能力。""" + _enabled(monkeypatch) + response = await controller.openid_configuration(_request(), _DiscoveryDb()) + assert response.status_code == _HTTP_OK + assert b'client_credentials' in response.body + assert b'backchannel_logout_supported' in response.body + assert b'backchannel_logout_session_supported' in response.body + + +@pytest.mark.asyncio +async def test_discovery_publishes_only_active_database_scopes(monkeypatch: pytest.MonkeyPatch) -> None: + """Discovery 应发布启用的 Resource Scope,并在停用后通过 ETag 变化。""" + _enabled(monkeypatch) + + class _Scalars: + def all(self) -> list[str]: + return ['openid', 'orders.read', 'zzz.read'] + + class _Result: + def scalars(self) -> _Scalars: + return _Scalars() + + class _Db: + async def execute(self, statement: object) -> _Result: + return _Result() + + response = await controller.openid_configuration(_request(), _Db()) + assert response.status_code == _HTTP_OK + assert b'orders.read' in response.body + assert response.headers['etag'] + + +@pytest.mark.asyncio +async def test_discovery_scope_query_failure_is_no_store_503(monkeypatch: pytest.MonkeyPatch) -> None: + """动态 Scope 查询失败时不得发布静态或不完整元数据。""" + _enabled(monkeypatch) + + class _FailingDb: + async def execute(self, statement: object) -> object: + raise RuntimeError('database unavailable') + + response = await controller.openid_configuration(_request(), _FailingDb()) + assert response.status_code == _HTTP_UNAVAILABLE + assert response.headers['cache-control'] == 'no-store' + assert response.body == b'{"error":"temporarily_unavailable"}' + + +@pytest.mark.asyncio +async def test_discovery_fails_closed_when_disabled(monkeypatch: pytest.MonkeyPatch) -> None: + """OIDC 关闭时所有 Discovery 请求均拒绝。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', False) + monkeypatch.setattr(controller, 'OidcConfig', OidcConfig) + response = await controller.openid_configuration(_request(), SimpleNamespace()) + assert response.status_code == _HTTP_NOT_FOUND + assert response.headers['cache-control'] == 'no-store' + assert response.body == b'{"error":"not_found"}' + + +@pytest.mark.asyncio +async def test_jwks_is_raw_and_uses_key_service(monkeypatch: pytest.MonkeyPatch) -> None: + """JWKS 端点返回标准裸 keys JSON 和缓存头。""" + _enabled(monkeypatch) + + async def build(*args: object, **kwargs: object) -> dict[str, list[dict[str, str]]]: + return {'keys': [{'kty': 'RSA', 'use': 'sig', 'kid': 'k1', 'alg': 'RS256', 'n': 'n', 'e': 'AQAB'}]} + + monkeypatch.setattr(controller.KeyService, 'build_jwks', build) + response = await controller.jwks(_request(), SimpleNamespace()) + assert response.status_code == _HTTP_OK + assert response.headers['cache-control'] == 'public, max-age=300' + assert response.body.startswith(b'{"keys"') + + +@pytest.mark.asyncio +async def test_jwks_fails_closed_on_invalid_public_key(monkeypatch: pytest.MonkeyPatch) -> None: + """JWKS 公钥记录非法时返回 503,不发布污染数据。""" + _enabled(monkeypatch) + + async def build(*args: object, **kwargs: object) -> dict[str, list[dict[str, str]]]: + raise KeyServiceError('public JWK RSA numbers are invalid') + + monkeypatch.setattr(controller.KeyService, 'build_jwks', build) + response = await controller.jwks(_request(), SimpleNamespace()) + assert response.status_code == _HTTP_UNAVAILABLE + assert response.headers['cache-control'] == 'no-store' + assert response.body == b'{"error":"temporarily_unavailable"}' diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_interaction_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_interaction_controller.py new file mode 100644 index 000000000..05adb4824 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_interaction_controller.py @@ -0,0 +1,936 @@ +import json +from dataclasses import dataclass +from datetime import datetime, timezone +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI +from fastapi.testclient import TestClient + +from common.aspect.db_session import get_db_session_provider +from common.constant import OidcAuditEvent +from common.enums import RedisInitKeyConfig +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException, OidcInteractionException +from exceptions.handle import handle_exception +from module_admin.service.user_service import UserService +from module_identity.controller.interaction_controller import ( + _failure_response, + _safe_json_body, + cancel, + captcha, + get_interaction, + interaction_controller, +) +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao, OAuthGrantSnapshot +from module_identity.dependencies import require_oidc_protocol_ready +from module_identity.entity.vo.interaction_vo import ( + CaptchaResponseModel, + ChangePasswordModel, + InteractionConsentModel, + InteractionLoginModel, + InteractionResultModel, + InteractionViewModel, +) +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import ( + AuthorizationCodeService, + AuthorizationContext, + ClientSnapshot, + InteractionCompletionService, + ScopeSnapshot, +) +from module_identity.service.consent_service import ConsentResult, ConsentService, InteractionConsentService +from module_identity.service.identity_service import ( + CredentialAuthenticationResult, + CredentialAuthenticationService, + IdentitySecurityEventService, + IdentitySubjectService, +) +from module_identity.service.infrastructure_service import AfterCommitCoordinator, OidcRateLimiter, RateLimitUnavailable +from module_identity.service.interaction_service import ( + InteractionFlowService, + InteractionLoginOutcome, + InteractionLoginService, + InteractionService, +) +from module_identity.service.session_service import SsoSessionDao, SsoSessionService +from tests.module_identity.support.redis_fakes import FakeRedis +from utils.oidc_util import OidcUtil +from utils.pwd_util import PwdUtil + +_PEPPER = 'interaction-controller-pepper-' + 'x' * 32 +_CHALLENGE = 'A' * 43 +_HTTP_OK = 200 +_HTTP_NOT_FOUND = 404 +_HTTP_SEE_OTHER = 303 +_HTTP_BAD_REQUEST = 400 +_HTTP_TOO_MANY_REQUESTS = 429 +_HTTP_SERVICE_UNAVAILABLE = 503 +_MIN_EXPECTED_COMMITS = 2 + + +class _Db: + """记录交互控制器事务边界。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + self.events: list[object] = [] + + def add(self, value: object) -> None: + self.events.append(value) + + async def flush(self) -> None: + return None + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +@dataclass +class _User: + """登录成功所需的最小用户标量。""" + + user_id: int = 2 + + +def _payload(status: str = 'awaiting_login', **overrides: object) -> dict[str, object]: + """构造 Interaction 内部白名单载荷。""" + payload: dict[str, object] = { + 'interactionId': 'interaction-controller-id', + 'clientPk': 1001, + 'clientId': 'portal-client', + 'redirectUri': 'https://portal.example/callback', + 'responseType': 'code', + 'scopes': ['openid', 'profile'], + 'resources': [], + 'state': 'opaque-state', + 'nonce': 'opaque-nonce', + 'codeChallenge': _CHALLENGE, + 'codeChallengeMethod': 'S256', + 'prompt': [], + 'maxAge': None, + 'consentRequired': status == 'awaiting_consent', + 'authenticatedSid': 'sid-1' if status in {'awaiting_consent', 'completed'} else None, + 'userId': 2 if status in {'awaiting_consent', 'completed'} else None, + 'subjectId': '11111111-1111-4111-8111-111111111111' if status in {'awaiting_consent', 'completed'} else None, + 'authVersion': 1 if status in {'awaiting_consent', 'completed'} else None, + } + payload.update(overrides) + return payload + + +def _request(redis: FakeRedis, csrf: str | None = None) -> SimpleNamespace: + """构造交互页面请求替身。""" + return SimpleNamespace( + app=SimpleNamespace(state=SimpleNamespace(redis=redis)), + headers={'x-csrf-token': csrf} if csrf else {}, + client=SimpleNamespace(host='127.0.0.1'), + cookies={}, + ) + + +def _context() -> AuthorizationContext: + """构造同意校验使用的不可变上下文。""" + return AuthorizationContext( + client=ClientSnapshot(1001, 'portal-client', 1, ('authorization_code',), ('code',), True, True, False), + redirect_uri='https://portal.example/callback', + scopes=('openid', 'profile'), + scope_models=( + ScopeSnapshot(1, 'openid', 'identity', None, True), + ScopeSnapshot(2, 'profile', 'identity', None, True), + ), + pre_authorized_scopes=frozenset({'openid'}), + resource=None, + state='opaque-state', + nonce='opaque-nonce', + code_challenge=_CHALLENGE, + code_challenge_method='S256', + prompt=None, + max_age=None, + ) + + +@pytest.fixture +def oidc_enabled(monkeypatch: pytest.MonkeyPatch) -> None: + """配置安全 OIDC 参数。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', _PEPPER) + monkeypatch.setattr(OidcConfig, 'oidc_authorization_code_ttl_seconds', 90) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_login_and_complete_real_interaction_lifecycle(monkeypatch: pytest.MonkeyPatch) -> None: + """真实 Interaction Redis 状态机贯通 CSRF、登录、完成和授权码一次消费标记。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + result = CredentialAuthenticationResult( + user=_User(), + dept=None, + acr='urn:ruoyi:acr:pwd', + amr=('pwd',), + remember_me=False, + password_change_required=False, + password_change_reason=None, + ) + subject = SimpleNamespace(subject_id='11111111-1111-4111-8111-111111111111', auth_version=1) + session = SimpleNamespace( + sid='sid-1', + user_id=2, + subject_id=subject.subject_id, + auth_version=1, + auth_time=None, + remember_me=False, + ) + monkeypatch.setattr( + 'module_identity.service.authorization_service.AuthorizationService._completion_grant', + AsyncMock(return_value=SimpleNamespace(grant_id='grant-1')), + ) + + async def authenticate(*args: object, **kwargs: object) -> CredentialAuthenticationResult: + return result + + async def require_subject(*args: object, **kwargs: object) -> Any: + return subject + + async def create_sso(*args: object, **kwargs: object) -> tuple[str, Any]: + return 'ss1.sid-1.' + 'A' * 43, session + + monkeypatch.setattr(CredentialAuthenticationService, 'authenticate_oidc', authenticate) + monkeypatch.setattr(IdentitySubjectService, 'require_by_user_id', require_subject) + monkeypatch.setattr(SsoSessionService, 'create', create_sso) + db = _Db() + outcome = await InteractionLoginService.login( + redis, + created.interaction_id, + InteractionLoginModel(userName='alice', password='password'), + db, + created.csrf_token, + '127.0.0.1', + None, + ) + assert outcome.result is not None + assert outcome.cookie is not None + assert (await InteractionService.get_record(redis, created.interaction_id))['status'] == 'completed' + assert 'ss1.sid-1.' in outcome.cookie + + monkeypatch.setattr( + OAuthClientDao, + 'find_exact_uri', + lambda *args, **kwargs: _async(SimpleNamespace(uri='https://portal.example/callback')), + ) + monkeypatch.setattr(SsoSessionDao, 'get_active', lambda *args, **kwargs: _async(session)) + audit_events: list[str] = [] + + async def record_audit(*args: object, **kwargs: object) -> None: + audit_events.append(str(args[1])) + + monkeypatch.setattr(AuditService, 'record', record_audit) + completed = await InteractionCompletionService.complete(redis, created.interaction_id, db) + assert 'code=ac1.' in completed.location + assert audit_events == [OidcAuditEvent.AUTHORIZE_SUCCEEDED] + with pytest.raises(OidcInteractionException): + await InteractionCompletionService.complete(redis, created.interaction_id, db) + assert await redis.ttl(OidcRedisKey.interaction(f'{created.interaction_id}-completion')) > 0 + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_forced_password_login_creates_no_sso_cookie_or_active_session( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """强制改密登录只推进短时证明状态,不提前创建 SSO。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + result = CredentialAuthenticationResult( + user=_User(), + dept=None, + acr='urn:ruoyi:acr:pwd', + amr=('pwd',), + remember_me=False, + password_change_required=True, + password_change_reason='initial_password', + ) + subject = SimpleNamespace(subject_id='11111111-1111-4111-8111-111111111111', auth_version=1) + + async def authenticate(*args: object, **kwargs: object) -> CredentialAuthenticationResult: + return result + + async def require_subject(*args: object, **kwargs: object) -> Any: + return subject + + async def forbidden_sso(*args: object, **kwargs: object) -> Any: + raise AssertionError('forced password login must not create SSO') + + monkeypatch.setattr(CredentialAuthenticationService, 'authenticate_oidc', authenticate) + monkeypatch.setattr(IdentitySubjectService, 'require_by_user_id', require_subject) + monkeypatch.setattr(SsoSessionService, 'create', forbidden_sso) + outcome = await InteractionLoginService.login( + redis, + created.interaction_id, + InteractionLoginModel(userName='alice', password='password'), + _Db(), + created.csrf_token, + ) + assert outcome.result is not None + assert outcome.cookie is None + record = await InteractionService.get_record(redis, created.interaction_id) + assert record['status'] == 'password_change_required' + assert record['authenticatedSid'] is None + assert isinstance(record['credentialProofHash'], str) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +@pytest.mark.usefixtures('interaction_page_metadata') +async def test_get_interaction_requires_csrf_header() -> None: + """Interaction 页面读取也必须使用原始 CSRF Header。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + with pytest.raises(OidcInteractionException): + await get_interaction(_request(redis), created.interaction_id, _Db()) + response = await get_interaction( + _request(redis, created.csrf_token), created.interaction_id, _Db(), created.csrf_token + ) + assert response.status_code == _HTTP_OK + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +@pytest.mark.parametrize(('stored_value', 'expected'), [('true', True), ('false', False)]) +@pytest.mark.usefixtures('interaction_page_metadata') +async def test_get_interaction_captcha_enabled_reflects_system_config(stored_value: str, expected: bool) -> None: + """Interaction 页面验证码开关必须来自系统配置,而不是固定默认值。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + config_key = f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.account.captchaEnabled' + await redis.set(config_key, stored_value) + + response = await get_interaction( + _request(redis, created.csrf_token), created.interaction_id, _Db(), created.csrf_token + ) + + assert response.status_code == _HTTP_OK + assert json.loads(response.body)['data']['captchaEnabled'] is expected + + +def test_interaction_error_responses_use_the_unified_response_util_envelope() -> None: + """交互 HTTP 错误保留统一 code/msg/success/time 业务 envelope。""" + response = _failure_response('invalid interaction', _HTTP_BAD_REQUEST) + payload = json.loads(response.body) + + assert response.status_code == _HTTP_BAD_REQUEST + assert payload['code'] is not None + assert payload['msg'] == 'invalid interaction' + assert payload['success'] is False + assert 'time' in payload + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_captcha_endpoint_has_independent_oidc_limit_and_fail_closed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """验证码端点使用 Interaction/IP 摘要限流,Redis 故障不放行。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + monkeypatch.setattr(InteractionFlowService, 'captcha_enabled', lambda *_args: _async(False)) + request = _request(redis) + for _ in range(5): + assert (await captcha(request, created.interaction_id)).status_code == _HTTP_OK + limited = await captcha(request, created.interaction_id) + assert limited.status_code == _HTTP_TOO_MANY_REQUESTS + assert all('interaction-controller-id' not in key for key in redis.values) + monkeypatch.setattr( + OidcRateLimiter, + 'enforce', + lambda *_args, **_kwargs: _raise_async(RateLimitUnavailable('down')), + ) + unavailable = await captcha(request, created.interaction_id) + assert unavailable.status_code == _HTTP_SERVICE_UNAVAILABLE + + +async def _raise_async(error: Exception) -> Any: + """在测试中构造异步异常结果。""" + raise error + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_mutations_require_csrf_and_consent_cannot_expand_scope(monkeypatch: pytest.MonkeyPatch) -> None: + """登录/同意变更必须有 CSRF,提交 Scope 不能超出原请求。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload('awaiting_consent'), pepper=_PEPPER) + db = _Db() + with pytest.raises(OidcInteractionException): + await InteractionConsentService.consent( + redis, + created.interaction_id, + InteractionConsentModel(approved=True, scopes=['openid', 'profile']), + db, + None, + ) + monkeypatch.setattr( + InteractionConsentService, + 'context_from_record', + lambda *args, **kwargs: _async(_context()), + ) + monkeypatch.setattr(AuditService, 'record_independent', lambda *args, **kwargs: _async(None)) + with pytest.raises(OAuthProtocolException) as expanded: + await InteractionConsentService.consent( + redis, + created.interaction_id, + InteractionConsentModel(approved=True, scopes=['openid', 'profile', 'admin']), + db, + created.csrf_token, + ) + assert expanded.value.error == 'invalid_scope' + assert (await InteractionService.get_record(redis, created.interaction_id))['status'] == 'awaiting_consent' + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_cancel_is_csrf_protected_and_legacy_cookie_is_ignored() -> None: + """取消操作使用 CAS,且交互控制器不读取 Legacy access_token Cookie。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + request = _request(redis, created.csrf_token) + request.cookies = {'access_token': 'legacy-token'} + response = await cancel(request, created.interaction_id, _Db(), created.csrf_token) + assert response.status_code == _HTTP_OK + assert (await InteractionService.get_record(redis, created.interaction_id))['status'] == 'denied' + + +@pytest.mark.asyncio +async def test_interaction_disabled_is_local_404(monkeypatch: pytest.MonkeyPatch) -> None: + """OIDC 关闭时交互页面和写入口均返回本地 404。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', False) + response = await get_interaction(_request(FakeRedis()), 'interaction-id', _Db()) + assert response.status_code == _HTTP_NOT_FOUND + assert response.headers['cache-control'] == 'no-store' + + +def test_consent_scope_dto_rejects_duplicate_and_oversized_values() -> None: + """同意 DTO 在进入 Controller 前拒绝重复或超长 Scope。""" + with pytest.raises(ValueError): + InteractionConsentModel(approved=True, scopes=['openid', 'openid']) + with pytest.raises(ValueError): + InteractionConsentModel(approved=True, scopes=['x' * 101]) + + +@pytest.mark.asyncio +async def test_body_parser_rejects_duplicate_unknown_and_oversized_json_without_echoing_password() -> None: + """JSON 解析错误统一脱敏,且受 16KiB 总体大小限制。""" + + class BodyRequest: + def __init__(self, value: bytes) -> None: + self.headers = {'content-type': 'application/json'} + self.value = value + + async def stream(self) -> Any: + midpoint = len(self.value) // 2 + yield self.value[:midpoint] + yield self.value[midpoint:] + yield b'' + + cases = [ + ('duplicate', b'{"userName":"alice","password":"super-secret","password":"again"}'), + ('unknown', json.dumps({'userName': 'alice', 'password': 'super-secret', 'unknown': True}).encode()), + ('oversized', b'{' + b'"password":"' + b'x' * (16 * 1024) + b'"}'), + ] + for reason, value in cases: + with pytest.raises(OidcInteractionException) as raised: + await _safe_json_body(BodyRequest(value), InteractionLoginModel) + assert raised.value.error == 'invalid_request', reason + assert 'super-secret' not in str(raised.value) + + +@pytest.mark.usefixtures('oidc_enabled') +def test_interaction_routes_stream_and_normalize_malformed_json_without_password_echo() -> None: + """真实 ASGI 路由以流式上限读取请求,异常响应不回显凭据。""" + app = FastAPI() + handle_exception(app) + app.include_router(interaction_controller) + + async def db_override() -> Any: + yield _Db() + + app.dependency_overrides[get_db_session_provider(None)] = db_override + app.dependency_overrides[require_oidc_protocol_ready] = lambda: None + client = TestClient(app) + cases = [ + ('application/json', b'{"userName":"alice","password":"secret","password":"again"}'), + ('application/json', json.dumps({'userName': 'alice', 'password': 'secret', 'unknown': True}).encode()), + ('application/json', b'{"userName":"alice","password":"' + b'x' * (16 * 1024) + b'"}'), + ('text/plain', b'{"userName":"alice","password":"secret"}'), + ] + for content_type, body in cases: + response = client.post( + '/auth/interaction/interaction-id/login', + content=body, + headers={'content-type': content_type}, + ) + assert response.status_code == _HTTP_BAD_REQUEST + assert 'secret' not in response.text + assert 'again' not in response.text + assert response.headers['cache-control'] == 'no-store' + assert response.headers['pragma'] == 'no-cache' + + +@pytest.mark.asyncio +async def test_db_commit_then_redis_cas_failure_runs_compensation(monkeypatch: pytest.MonkeyPatch) -> None: + """DB 已提交但 Redis CAS 失败时必须执行补偿且保留可重试状态。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + compensated = False + + async def failed_transition(*args: object, **kwargs: object) -> object: + raise OidcInteractionException(created.interaction_id, 'CAS failed', error='server_error', status_code=503) + + async def compensate() -> None: + nonlocal compensated + compensated = True + + monkeypatch.setattr(InteractionService, 'transition', failed_transition) + with pytest.raises(OidcInteractionException): + await InteractionFlowService.commit_transition( + _Db(), + AfterCommitCoordinator(), + redis, + created.interaction_id, + 'completed', + compensate=compensate, + ) + assert compensated is True + assert (await InteractionService.get_record(redis, created.interaction_id))['status'] == 'awaiting_login' + + +@pytest.mark.asyncio +async def test_consent_commit_then_redis_cas_failure_compensates_persisted_grant( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """同意 Grant 已提交但 Interaction CAS 失败时必须持久化撤销补偿。""" + redis = FakeRedis() + created = await InteractionService.create( + redis, + _payload('awaiting_consent', userId=2, subjectId='subject-2', authVersion=1), + pepper=_PEPPER, + ) + grant = SimpleNamespace(grant_id='grant-after-commit', user_id=2, subject_id='subject-2') + persisted = OAuthGrantSnapshot( + grant_id=grant.grant_id, + user_id=grant.user_id, + subject_id=grant.subject_id, + client_pk=1, + granted_scopes=('openid', 'profile'), + granted_resources=(), + client_policy_version=1, + status='active', + consented_at=datetime.now(timezone.utc), + expires_at=None, + revoked_at=None, + revoke_reason=None, + last_used_at=None, + ) + result = ConsentResult( + approved=True, + scopes=('openid', 'profile'), + grant=grant, + persisted_grant=persisted, + ) + db = _Db() + compensated = False + + async def revoke(*_args: object, **_kwargs: object) -> bool: + return True + + monkeypatch.setattr(OAuthGrantDao, 'revoke_snapshot', revoke) + + async def compensate() -> None: + nonlocal compensated + compensated = True + await ConsentService.compensate_persisted_grant(db, result) + + async def cas_failure(*_args: object, **_kwargs: object) -> None: + raise OidcInteractionException(created.interaction_id, 'CAS failed', error='server_error', status_code=503) + + monkeypatch.setattr(InteractionService, 'transition', cas_failure) + with pytest.raises(OidcInteractionException): + await InteractionFlowService.commit_transition( + db, + AfterCommitCoordinator(), + redis, + created.interaction_id, + 'completed', + compensate=compensate, + ) + + assert compensated is True + assert db.commits >= _MIN_EXPECTED_COMMITS + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_password_change_cas_failure_runs_full_security_saga(monkeypatch: pytest.MonkeyPatch) -> None: + """改密已提交后 CAS 失败会删除 Interaction、撤销新 Session 并独立记高危审计。""" + subject_id = '11111111-1111-4111-8111-111111111111' + payload = _payload( + 'password_change_required', + userId=2, + subjectId=subject_id, + authVersion=1, + credentialProofHash=OidcUtil.credential_proof( + 'interaction-controller-id', 2, subject_id, 1, pepper=OidcConfig.oidc_token_hash_pepper + ), + ) + redis = FakeRedis() + created = await InteractionService.create(redis, payload, pepper=_PEPPER) + stored = await InteractionService.get_record(redis, created.interaction_id) + stored['status'] = 'password_change_required' + stored['credentialProofHash'] = OidcUtil.credential_proof( + created.interaction_id, 2, subject_id, 1, pepper=OidcConfig.oidc_token_hash_pepper + ) + await redis.set( + OidcRedisKey.interaction(created.interaction_id), + json.dumps(stored, separators=(',', ':')), + ex=60, + ) + user = SimpleNamespace( + user_id=2, + password=PwdUtil.get_password_hash('old-password'), + status='0', + del_flag='0', + pwd_update_date=None, + ) + subject = SimpleNamespace(subject_id=subject_id, auth_version=2) + session = SimpleNamespace(sid='new-session', user_id=2, subject_id=subject_id, auth_version=2) + db = _ScalarDb(user) + revoked: list[str] = [] + audit_events: list[str] = [] + + monkeypatch.setattr(UserService, 'validate_password_services', lambda *_args, **_kwargs: _async(None)) + monkeypatch.setattr(IdentitySubjectService, 'require_by_user_id', lambda *_args, **_kwargs: _async(subject)) + monkeypatch.setattr(IdentitySecurityEventService, 'handle_user_event', lambda *_args, **_kwargs: _async(None)) + monkeypatch.setattr( + SsoSessionService, + 'revoke_user', + lambda *_args, **_kwargs: _async(None), + ) + monkeypatch.setattr( + SsoSessionService, + 'create', + lambda *_args, **_kwargs: _async(('ss1.new-session.' + 'A' * 43, session)), + ) + + async def revoke(*_args: object, **_kwargs: object) -> None: + revoked.append(session.sid) + + async def record_independent(*args: object, **_kwargs: object) -> None: + audit_events.append(str(args[1])) + + monkeypatch.setattr(SsoSessionService, 'revoke', revoke) + monkeypatch.setattr(AuditService, 'record_independent', record_independent) + + async def cas_fail(*args: object, **kwargs: object) -> None: + raise OidcInteractionException(created.interaction_id, 'CAS failed', error='server_error', status_code=503) + + monkeypatch.setattr(InteractionService, 'transition', cas_fail) + with pytest.raises(OidcInteractionException) as raised: + await InteractionLoginService.change_password( + redis, + created.interaction_id, + ChangePasswordModel( + oldPassword='old-password', newPassword='new-password-A1!', confirmPassword='new-password-A1!' + ), + db, + created.csrf_token, + ) + assert raised.value.error == 'server_error' + assert await redis.get(OidcRedisKey.interaction(created.interaction_id)) is None + assert revoked == [session.sid] + assert audit_events == [OidcAuditEvent.LOGIN_FAILED] + assert db.commits >= _MIN_EXPECTED_COMMITS + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_password_change_saga_continues_when_interaction_delete_fails( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """补偿删除 Interaction 失败时仍继续撤销 Session 和独立审计。""" + subject_id = '11111111-1111-4111-8111-111111111111' + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + stored = await InteractionService.get_record(redis, created.interaction_id) + stored.update( + { + 'status': 'password_change_required', + 'userId': 2, + 'subjectId': subject_id, + 'authVersion': 1, + 'credentialProofHash': OidcUtil.credential_proof( + created.interaction_id, 2, subject_id, 1, pepper=OidcConfig.oidc_token_hash_pepper + ), + } + ) + await redis.set(OidcRedisKey.interaction(created.interaction_id), json.dumps(stored), ex=60) + user = SimpleNamespace( + user_id=2, + password=PwdUtil.get_password_hash('old-password'), + status='0', + del_flag='0', + pwd_update_date=None, + ) + session = SimpleNamespace(sid='new-session', user_id=2, subject_id=subject_id, auth_version=2) + db = _ScalarDb(user) + revoked: list[str] = [] + audits: list[str] = [] + + async def delete_failure(*_args: object, **_kwargs: object) -> int: + raise RuntimeError('redis unavailable') + + monkeypatch.setattr(redis, 'delete', delete_failure) + monkeypatch.setattr(UserService, 'validate_password_services', lambda *_args, **_kwargs: _async(None)) + monkeypatch.setattr( + IdentitySubjectService, + 'require_by_user_id', + lambda *_args, **_kwargs: _async(SimpleNamespace(subject_id=subject_id, auth_version=2)), + ) + monkeypatch.setattr(IdentitySecurityEventService, 'handle_user_event', lambda *_args, **_kwargs: _async(None)) + monkeypatch.setattr(SsoSessionService, 'revoke_user', lambda *_args, **_kwargs: _async(None)) + monkeypatch.setattr(SsoSessionService, 'create', lambda *_args, **_kwargs: _async(('cookie', session))) + monkeypatch.setattr(SsoSessionService, 'revoke', lambda *_args, **_kwargs: _append_async(revoked, session.sid)) + monkeypatch.setattr(AuditService, 'record_independent', lambda *args, **kwargs: _append_async(audits, str(args[1]))) + + async def cas_fail(*_args: object, **_kwargs: object) -> None: + raise OidcInteractionException(created.interaction_id, 'CAS failed', error='server_error', status_code=503) + + monkeypatch.setattr(InteractionService, 'transition', cas_fail) + with pytest.raises(OidcInteractionException): + await InteractionLoginService.change_password( + redis, + created.interaction_id, + ChangePasswordModel( + oldPassword='old-password', newPassword='new-password-A1!', confirmPassword='new-password-A1!' + ), + db, + created.csrf_token, + ) + assert revoked == [session.sid] + assert audits == [OidcAuditEvent.LOGIN_FAILED] + + +class _ScalarDb(_Db): + """支持改密测试查询用户的最小数据库替身。""" + + def __init__(self, value: object) -> None: + super().__init__() + self.value = value + + async def scalar(self, *_args: object, **_kwargs: object) -> object: + return self.value + + async def execute(self, *_args: object, **_kwargs: object) -> object: + value = self.value + + class _Result: + def scalars(self) -> object: + return self + + def first(self) -> object: + return value + + return _Result() + + +@pytest.mark.asyncio +async def test_cache_callback_failure_does_not_compensate_committed_session() -> None: + """提交后缓存回调失败与状态 CAS 失败分离,不撤销已落库 Session。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + compensated = False + + async def cache_failure() -> None: + raise RuntimeError('cache unavailable') + + async def compensate() -> None: + nonlocal compensated + compensated = True + + coordinator = AfterCommitCoordinator() + await coordinator.register(cache_failure) + with pytest.raises(OidcInteractionException): + await InteractionFlowService.commit_transition( + _Db(), + coordinator, + redis, + created.interaction_id, + 'completed', + compensate=compensate, + ) + assert compensated is False + assert (await InteractionService.get_record(redis, created.interaction_id))['status'] == 'completed' + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_code_audit_failure_invalidates_code_and_marker(monkeypatch: pytest.MonkeyPatch) -> None: + """审计或提交失败时授权码与 completion marker 均不可继续使用。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload('completed'), pepper=_PEPPER) + marker = OidcRedisKey.interaction(f'{created.interaction_id}-completion') + await redis.set(marker, 'reserved', ex=60) + session = SimpleNamespace( + sid='sid-1', + user_id=2, + subject_id='11111111-1111-4111-8111-111111111111', + auth_version=1, + auth_time=None, + ) + monkeypatch.setattr( + InteractionCompletionService, + 'active_session', + lambda *_args, **_kwargs: _async(session), + ) + monkeypatch.setattr( + 'module_identity.service.authorization_service.AuthorizationService._completion_grant', + AsyncMock(return_value=SimpleNamespace(grant_id='grant-1')), + ) + + async def audit_failure(*_args: object, **_kwargs: object) -> None: + raise RuntimeError('audit unavailable') + + monkeypatch.setattr(AuditService, 'record', audit_failure) + with pytest.raises(OAuthProtocolException) as raised: + await InteractionCompletionService._issue_code( + _Db(), redis, _payload('completed'), 'https://portal.example/callback', marker + ) + assert raised.value.error == 'server_error' + assert await redis.get(marker) is None + assert not any(key.startswith('oidc:authorization_code:') for key in redis.values) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_code_invalidate_failure_still_cleans_marker_and_rolls_back(monkeypatch: pytest.MonkeyPatch) -> None: + """授权码失效补偿自身失败时仍继续清理 marker、回滚且只返回脱敏错误。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload('completed'), pepper=_PEPPER) + marker = OidcRedisKey.interaction(f'{created.interaction_id}-completion') + await redis.set(marker, 'reserved', ex=60) + session = SimpleNamespace( + sid='sid-1', + user_id=2, + subject_id='11111111-1111-4111-8111-111111111111', + auth_version=1, + auth_time=None, + ) + monkeypatch.setattr( + InteractionCompletionService, + 'active_session', + lambda *_args, **_kwargs: _async(session), + ) + monkeypatch.setattr(AuditService, 'record', lambda *_args, **_kwargs: _raise_async(RuntimeError('audit'))) + monkeypatch.setattr( + AuthorizationCodeService, + 'invalidate', + lambda *_args, **_kwargs: _raise_async(RuntimeError('redis invalidate failed')), + ) + db = _Db() + with pytest.raises(OAuthProtocolException) as raised: + await InteractionCompletionService._issue_code( + db, redis, _payload('completed'), 'https://portal.example/callback', marker + ) + assert raised.value.error == 'server_error' + assert await redis.get(marker) is None + assert db.rollbacks == 1 + monkeypatch.setattr( + 'module_identity.service.authorization_service.AuthorizationService._completion_grant', + AsyncMock(return_value=SimpleNamespace(grant_id='grant-1')), + ) + + +def _async(value: object) -> Any: + """构造异步测试结果。""" + + async def result() -> object: + return value + + return result() + + +def _append_async(values: list[Any], value: Any) -> Any: + """构造追加测试值的异步结果。""" + + async def result() -> None: + values.append(value) + + return result() + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +@pytest.mark.parametrize('endpoint', ['', 'captcha', 'login', 'change-password', 'consent', 'cancel', 'complete']) +@pytest.mark.usefixtures('interaction_page_metadata') +async def test_interaction_http_response_matches_frontend_data_contract( + monkeypatch: pytest.MonkeyPatch, endpoint: str +) -> None: + """真实 HTTP 响应按 OpenAPI 和页面约定将交互字段放在 data 下。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + result = InteractionResultModel( + next_action='redirect', interaction_id=created.interaction_id, redirect_url='/auth/interaction/complete' + ) + outcome = InteractionLoginOutcome(result=result) + monkeypatch.setattr(InteractionLoginService, 'login', AsyncMock(return_value=outcome)) + monkeypatch.setattr(InteractionLoginService, 'change_password', AsyncMock(return_value=outcome)) + monkeypatch.setattr(InteractionConsentService, 'consent', AsyncMock(return_value=result)) + monkeypatch.setattr(InteractionConsentService, 'cancel', AsyncMock(return_value=result)) + monkeypatch.setattr( + InteractionCompletionService, 'complete', AsyncMock(return_value=SimpleNamespace(location=result.redirect_url)) + ) + await redis.set(f'{RedisInitKeyConfig.SYS_CONFIG.key}:sys.account.captchaEnabled', 'false') + app = FastAPI() + app.state.redis = redis + app.include_router(interaction_controller) + + async def db_override() -> Any: + yield _Db() + + app.dependency_overrides[get_db_session_provider(None)] = db_override + app.dependency_overrides[require_oidc_protocol_ready] = lambda: None + bodies = { + 'login': {'userName': 'alice', 'password': 'test-password'}, + 'change-password': {'oldPassword': 'old', 'newPassword': 'new', 'confirmPassword': 'new'}, + 'consent': {'approved': True, 'scopes': ['openid']}, + } + path = '/auth/interaction/' + created.interaction_id + ('/' + endpoint if endpoint else '') + with TestClient(app) as client: + response = client.request( + 'GET' if endpoint in {'', 'captcha'} else 'POST', + path, + headers={'X-CSRF-Token': created.csrf_token}, + json=bodies.get(endpoint), + ) + assert response.status_code == _HTTP_OK + payload = response.json() + model = ( + InteractionViewModel + if not endpoint + else CaptchaResponseModel + if endpoint == 'captcha' + else InteractionResultModel + ) + model.model_validate(payload['data']) + assert payload['code'] == _HTTP_OK and payload['success'] + assert 'nextAction' not in payload and 'captchaEnabled' not in payload + assert response.headers['cache-control'] == 'no-store' diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_logout_confirmation.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_logout_confirmation.py new file mode 100644 index 000000000..f13225da5 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_logout_confirmation.py @@ -0,0 +1,196 @@ +import re +from collections.abc import AsyncIterator, Iterator +from time import monotonic +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI, status +from fastapi.testclient import TestClient + +from common.aspect.db_session import get_db_session_provider +from config.env import OidcConfig +from middlewares.oidc_cors_middleware import OidcCorsMiddleware +from module_identity.controller.authorization_controller import authorization_controller +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dependencies import require_oidc_protocol_ready +from module_identity.service.infrastructure_service import OidcRateLimiter +from module_identity.service.logout_confirmation_service import LogoutConfirmationService +from module_identity.service.session_service import LogoutResult, LogoutService + + +class ConfirmationRedis: + def __init__(self) -> None: + self.values: dict[str, tuple[str, float]] = {} + + async def set(self, key: str, value: str, *, ex: int, nx: bool = False) -> bool: + if nx and key in self.values: + return False + self.values[key] = (value, monotonic() + ex) + return True + + async def eval(self, _script: str, _count: int, key: str) -> str | None: + value, expires = self.values.pop(key, (None, 0)) + return value if expires > monotonic() else None + + +@pytest.fixture +def logout_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[tuple[TestClient, AsyncMock]]: + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example') + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', 'logout-confirmation-test-' + 'x' * 32) + monkeypatch.setattr(OidcRateLimiter, 'enforce', AsyncMock()) + monkeypatch.setattr(LogoutConfirmationService, 'form_redirect_origin', AsyncMock(return_value='https://rp.example')) + app = FastAPI() + app.add_middleware(OidcCorsMiddleware) + app.state.redis = ConfirmationRedis() + app.include_router(authorization_controller) + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + + async def get_db() -> AsyncIterator[SimpleNamespace]: + yield db + + app.dependency_overrides[get_db_session_provider(None)] = get_db + app.dependency_overrides[require_oidc_protocol_ready] = lambda: None + execute = AsyncMock(return_value=LogoutResult('https://rp.example/done', 'state-1', True)) + monkeypatch.setattr(LogoutService, 'execute_logout', execute) + with TestClient(app, base_url='https://auth.example', follow_redirects=False) as client: + yield client, execute + + +@pytest.mark.parametrize('method', ['GET', 'POST']) +def test_standard_logout_requires_confirmation_then_redirects( + logout_client: tuple[TestClient, AsyncMock], method: str +) -> None: + client, execute = logout_client + parameters = { + 'id_token_hint': 'private-hint', + 'post_logout_redirect_uri': 'https://rp.example/done', + 'state': 'state-1', + } + response = ( + client.get('/oauth2/logout', params=parameters) + if method == 'GET' + else client.post('/oauth2/logout', data=parameters) + ) + assert response.status_code == status.HTTP_200_OK + assert '确认退出' in response.text + assert 'private-hint' not in response.text + assert "frame-ancestors 'none'" in response.headers['content-security-policy'] + assert "form-action 'self' https://rp.example;" in response.headers['content-security-policy'] + execute.assert_not_awaited() + token = re.search(r'name="confirmation" value="([^"]+)"', response.text).group(1) + response = client.post( + '/oauth2/logout/confirm', + data={'confirmation': token, 'decision': 'confirm'}, + headers={'Origin': 'https://auth.example'}, + ) + assert response.status_code == status.HTTP_303_SEE_OTHER + assert response.headers['location'] == 'https://rp.example/done?state=state-1' + assert execute.await_args.kwargs['confirmed'] is True + assert execute.await_args.kwargs['id_token_hint'] == 'private-hint' + execute.reset_mock() + replay = client.post( + '/oauth2/logout/confirm', + data={'confirmation': token, 'decision': 'confirm'}, + headers={'Origin': 'https://auth.example'}, + ) + assert replay.status_code == status.HTTP_400_BAD_REQUEST + execute.assert_not_awaited() + + +def test_cancel_keeps_sso_cookie_and_never_calls_revocation(logout_client: tuple[TestClient, AsyncMock]) -> None: + client, execute = logout_client + client.cookies.set(OidcConfig.oidc_sso_cookie_name, 'current-sso', domain='auth.example', path='/') + page = client.get('/oauth2/logout') + token = re.search(r'name="confirmation" value="([^"]+)"', page.text).group(1) + response = client.post( + '/oauth2/logout/confirm', + data={'confirmation': token, 'decision': 'cancel'}, + headers={'Origin': 'https://auth.example'}, + ) + assert response.status_code == status.HTTP_200_OK + assert '已取消退出' in response.text + assert OidcConfig.oidc_sso_cookie_name + '=' not in response.headers.get('set-cookie', '') + execute.assert_not_awaited() + + +@pytest.mark.parametrize('attack', ['origin', 'cookie', 'expired']) +def test_confirmation_rejects_wrong_origin_browser_or_expiry( + logout_client: tuple[TestClient, AsyncMock], attack: str +) -> None: + client, execute = logout_client + page = client.get('/oauth2/logout') + token = re.search(r'name="confirmation" value="([^"]+)"', page.text).group(1) + origin = 'https://evil.example' if attack == 'origin' else 'https://auth.example' + if attack == 'cookie': + client.cookies.clear() + if attack == 'expired': + redis = client.app.state.redis + redis.values = {key: (value, 0) for key, (value, _expires) in redis.values.items()} + response = client.post( + '/oauth2/logout/confirm', data={'confirmation': token, 'decision': 'confirm'}, headers={'Origin': origin} + ) + assert response.status_code in {400, 403} + execute.assert_not_awaited() + + +def test_cross_site_post_may_omit_lax_sso_cookie_until_confirmation( + logout_client: tuple[TestClient, AsyncMock], +) -> None: + """跨站POST不携带Lax Cookie,同源确认请求可携带会话Cookie。""" + client, execute = logout_client + page = client.post('/oauth2/logout', data={'state': 'cross-site'}) + assert page.headers['referrer-policy'] == 'strict-origin' + token = re.search(r'name="confirmation" value="([^"]+)"', page.text).group(1) + client.cookies.set(OidcConfig.oidc_sso_cookie_name, 'current-sso', domain='auth.example', path='/') + response = client.post( + '/oauth2/logout/confirm', + data={'confirmation': token, 'decision': 'confirm'}, + headers={'Origin': 'https://auth.example'}, + ) + assert response.status_code == status.HTTP_303_SEE_OTHER + assert execute.await_args.kwargs['cookie'] == 'current-sso' + + +def test_changing_an_initially_bound_sso_cookie_rejects_confirmation( + logout_client: tuple[TestClient, AsyncMock], +) -> None: + client, execute = logout_client + client.cookies.set(OidcConfig.oidc_sso_cookie_name, 'first-sso', domain='auth.example', path='/') + page = client.get('/oauth2/logout') + token = re.search(r'name="confirmation" value="([^"]+)"', page.text).group(1) + client.cookies.set(OidcConfig.oidc_sso_cookie_name, 'second-sso', domain='auth.example', path='/') + response = client.post( + '/oauth2/logout/confirm', + data={'confirmation': token, 'decision': 'confirm'}, + headers={'Origin': 'https://auth.example'}, + ) + assert response.status_code == status.HTTP_400_BAD_REQUEST + execute.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize('registered', [True, False]) +async def test_confirmation_csp_allows_only_verified_registered_redirect( + monkeypatch: pytest.MonkeyPatch, registered: bool +) -> None: + monkeypatch.setattr( + LogoutService, '_validate_id_token_hint', AsyncMock(return_value=({}, SimpleNamespace(client_pk=1))) + ) + monkeypatch.setattr( + OAuthClientDao, + 'find_exact_uri', + AsyncMock(return_value=SimpleNamespace(status='0') if registered else None), + ) + origin = await LogoutConfirmationService.form_redirect_origin( + object(), {'id_token_hint': 'valid-hint', 'post_logout_redirect_uri': 'https://rp.example:9443/done'} + ) + assert origin == ('https://rp.example:9443' if registered else None) + monkeypatch.setattr(LogoutService, '_validate_id_token_hint', AsyncMock(side_effect=ValueError('invalid hint'))) + assert ( + await LogoutConfirmationService.form_redirect_origin( + object(), {'id_token_hint': 'invalid-hint', 'post_logout_redirect_uri': 'https://rp.example/done'} + ) + is None + ) diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_audit_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_audit_controller.py new file mode 100644 index 000000000..eccf974ff --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_audit_controller.py @@ -0,0 +1,151 @@ +from datetime import datetime, timezone +from http import HTTPStatus +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI, Request +from fastapi.routing import APIRoute +from httpx import ASGITransport, AsyncClient + +from common.annotation.rate_limit_annotation import ApiRateLimit +from common.context import RequestContext +from module_admin.service.log_service import LogQueueService +from module_identity.controller.oauth_audit_controller import ( + list_oauth_audit, + oauth_audit_controller, +) +from module_identity.dao.oauth_audit_dao import OAuthAuditDao +from module_identity.entity.vo.oauth_session_vo import AuditPageQueryModel +from module_identity.service.audit_service import AuditService + + +class _Session: + """提供只读回滚边界。""" + + def __init__(self) -> None: + self.rollbacks = 0 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +def test_audit_routes_use_list_and_export_permissions() -> None: + """审计路由必须具备 PreAuth 和精确权限。""" + routes = { + (route.path, method): route + for route in oauth_audit_controller.routes + if isinstance(route, APIRoute) + for method in route.methods + } + assert 'monitor:oauthAudit:list' in [ + getattr(dep.call, 'perm', None) for dep in routes[('/monitor/oauth/audit/list', 'GET')].dependant.dependencies + ] + assert 'monitor:oauthAudit:export' in [ + getattr(dep.call, 'perm', None) + for dep in routes[('/monitor/oauth/audit/export', 'POST')].dependant.dependencies + ] + assert oauth_audit_controller.dependencies[0].dependency.__class__.__name__ == 'PreAuth' + + +def test_audit_export_uses_form_content_type() -> None: + """导出接口必须接受前端 proxy.download 发送的表单,而非 JSON body。""" + app = FastAPI() + app.include_router(oauth_audit_controller) + content = app.openapi()['paths']['/monitor/oauth/audit/export']['post']['requestBody']['content'] + assert 'application/x-www-form-urlencoded' in content + assert 'application/json' not in content + + +@pytest.mark.asyncio +async def test_audit_list_uses_real_total_and_safe_projection(monkeypatch: pytest.MonkeyPatch) -> None: + """审计列表返回服务层真实总数且不含 detail。""" + rows = [ + SimpleNamespace( + event_id=1, + event_type='token_failed', + result='failure', + risk_level='high', + client_id='c1', + resource_id=None, + user_id=1, + subject_id='s', + sid='sid', + failure_code='invalid_client', + create_time=datetime.now(timezone.utc), + detail={'access_token': 'redact'}, + ), + SimpleNamespace( + event_id=2, + event_type='login', + result='success', + risk_level='normal', + client_id='c1', + resource_id=None, + user_id=1, + subject_id='s', + sid='sid', + failure_code=None, + create_time=datetime.now(timezone.utc), + detail={'password': 'redact'}, + ), + ] + + async def list_page(*args: object, **kwargs: object) -> list[object]: + return rows + + async def count(*args: object, **kwargs: object) -> int: + return 37 + + monkeypatch.setattr(OAuthAuditDao, 'list_admin_page', list_page) + monkeypatch.setattr(OAuthAuditDao, 'count_admin', count) + response = await list_oauth_audit(AuditPageQueryModel(page_num=1, page_size=2), _Session()) + assert b'"total":37' in response.body + assert b'"detail"' not in response.body + assert b'access_token' not in response.body + + +@pytest.mark.asyncio +async def test_audit_export_runs_rate_limit_and_operation_log(monkeypatch: pytest.MonkeyPatch) -> None: + """HTTP 审计导出经过完整装饰器链,保留文件响应并记录操作日志。""" + + query = AuditPageQueryModel() + db = _Session() + user = SimpleNamespace(user=SimpleNamespace(user_id=1, user_name='admin', dept=None)) + route = next( + route + for route in oauth_audit_controller.routes + if isinstance(route, APIRoute) and route.path == '/monitor/oauth/audit/export' + ) + app = FastAPI() + app.state.redis = SimpleNamespace() + app.include_router(oauth_audit_controller) + for dependency in route.dependant.dependencies: + if dependency.name == 'query_db': + app.dependency_overrides[dependency.call] = lambda: db + else: + app.dependency_overrides[dependency.call] = lambda: None + rate_limit = AsyncMock(return_value={'allowed': True, 'remaining': 9, 'reset_at': 60}) + export = AsyncMock(return_value=b'audit-workbook') + operation_log = AsyncMock() + monkeypatch.setattr(ApiRateLimit, '_acquire_rate_limit', rate_limit) + monkeypatch.setattr(AuditService, 'export_admin', export) + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', operation_log) + token = RequestContext.set_current_user(user) + try: + async with AsyncClient(transport=ASGITransport(app=app), base_url='http://testserver') as client: + response = await client.post( + '/monitor/oauth/audit/export', data=query.model_dump(by_alias=True, exclude_none=True) + ) + finally: + RequestContext.reset_current_user(token) + assert response.status_code == HTTPStatus.OK + assert response.headers['content-disposition'] == 'attachment; filename="oauth-audit.xlsx"' + assert response.content == b'audit-workbook' + export.assert_awaited_once_with(db, query) + rate_limit.assert_awaited_once() + operation_log.assert_awaited_once() + logged_request, logged_operation, _ = operation_log.call_args.args + assert isinstance(logged_request, Request) + assert logged_request is rate_limit.call_args.args[-1] + assert logged_operation.oper_url == '/monitor/oauth/audit/export' diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_client_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_client_controller.py new file mode 100644 index 000000000..262022ffe --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_client_controller.py @@ -0,0 +1,201 @@ +from types import SimpleNamespace + +import pytest +from fastapi import Request +from fastapi.routing import APIRoute +from starlette.responses import Response + +from common.context import RequestContext +from module_admin.service.log_service import LogQueueService +from module_identity.controller.oauth_client_controller import ( + add_system_oauth_client, + delete_system_oauth_clients, + get_system_oauth_client_list, + oauth_client_controller, + rotate_system_oauth_client_secret, +) +from module_identity.entity.vo.oauth_client_vo import ClientSecretResponseModel, SecretRotationModel +from module_identity.service.oauth_management_service import ( + OAuthClientManagementError, + OAuthClientManagementService, +) +from module_identity.service.runtime_service import OidcRuntimeService + + +class _FakeSession: + """记录 Controller 的 commit/rollback 边界。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +@pytest.fixture(autouse=True) +def _request_context(monkeypatch: pytest.MonkeyPatch) -> None: + """为日志装饰器提供当前用户上下文。""" + monkeypatch.setattr(RequestContext, 'get_current_user', staticmethod(_user)) + + async def enqueue_operation_log(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', enqueue_operation_log) + + +def _user() -> SimpleNamespace: + """构造最小当前用户对象。""" + return SimpleNamespace(user=SimpleNamespace(user_name='admin', dept=SimpleNamespace(dept_name='测试部门'))) + + +def _request() -> Request: + """构造带可观察 CORS 快照的请求。""" + app = SimpleNamespace(state=SimpleNamespace(oidc_registered_cors_origins=('https://old.example',))) + + async def receive() -> dict[str, object]: + return {'type': 'http.request', 'body': b'', 'more_body': False} + + return Request( + { + 'type': 'http', + 'method': 'POST', + 'path': '/', + 'headers': [], + 'query_string': b'', + 'app': app, + }, + receive=receive, + ) + + +def _routes() -> list[APIRoute]: + """返回 Client Router 的 HTTP 路由。""" + return [route for route in oauth_client_controller.routes if isinstance(route, APIRoute)] + + +def test_client_router_paths_methods_permissions_and_pre_auth() -> None: + """Client 路由必须完整覆盖规范路径并逐端点声明权限。""" + expected = { + ('/system/oauth/client/list', 'GET'): 'system:oauthClient:list', + ('/system/oauth/client/{client_id}', 'GET'): 'system:oauthClient:query', + ('/system/oauth/client', 'POST'): 'system:oauthClient:add', + ('/system/oauth/client', 'PUT'): 'system:oauthClient:edit', + ('/system/oauth/client/{client_ids}', 'DELETE'): 'system:oauthClient:remove', + ('/system/oauth/client/changeStatus', 'PUT'): 'system:oauthClient:edit', + ('/system/oauth/client/{client_id}/secret', 'POST'): 'system:oauthClient:rotateSecret', + ('/system/oauth/client/{client_id}/secret/{secret_id}', 'DELETE'): 'system:oauthClient:rotateSecret', + ('/system/oauth/client/{client_id}/uri', 'POST'): 'system:oauthClient:edit', + ('/system/oauth/client/{client_id}/uri/{uri_id}', 'DELETE'): 'system:oauthClient:edit', + } + routes = {(route.path, method): route for route in _routes() for method in route.methods} + assert set(routes) == set(expected) + for key, permission in expected.items(): + dependency_calls = [getattr(item.call, 'perm', None) for item in routes[key].dependant.dependencies] + assert permission in dependency_calls + assert oauth_client_controller.dependencies + assert oauth_client_controller.dependencies[0].dependency.__class__.__name__ == 'PreAuth' + + +@pytest.mark.asyncio +async def test_client_create_commits_and_secret_is_only_in_rotation_response(monkeypatch: pytest.MonkeyPatch) -> None: + """成功写操作提交事务,轮换响应含明文而普通错误不包含敏感字段。""" + session = _FakeSession() + payload = SimpleNamespace() + + async def fake_create(*args: object, **kwargs: object) -> object: + return SimpleNamespace(model_dump=lambda **_: {'clientId': 'cli_test'}) + + monkeypatch.setattr(OAuthClientManagementService, 'create_client', fake_create) + response = await add_system_oauth_client(_request(), payload, session, _user()) + assert isinstance(response, Response) + assert session.commits == 0 and session.rollbacks == 0 + + secret = ClientSecretResponseModel( + client_id='cli_test', + secret_id='secret-1', + client_secret='cs1.one-time-secret', + secret_hint='...cret', + not_before='2026-01-01T00:00:00Z', + ) + + async def fake_rotate(*args: object, **kwargs: object) -> ClientSecretResponseModel: + return secret + + monkeypatch.setattr(OAuthClientManagementService, 'rotate_secret', fake_rotate) + rotation_response = await rotate_system_oauth_client_secret( + _request(), 'cli_test', session, _user(), SecretRotationModel() + ) + assert b'cs1.one-time-secret' in rotation_response.body + assert session.commits == 0 + + async def fake_fail(*args: object, **kwargs: object) -> object: + raise OAuthClientManagementError('OAuth Client request rejected') + + monkeypatch.setattr(OAuthClientManagementService, 'create_client', fake_fail) + failed = await add_system_oauth_client(_request(), payload, session, _user()) + assert b'secret_hash' not in failed.body + assert b'cs1.' not in failed.body + assert session.rollbacks == 0 + + +@pytest.mark.asyncio +async def test_client_batch_failure_rolls_back_all_items(monkeypatch: pytest.MonkeyPatch) -> None: + """批量停用任一项失败时整体回滚,不提前提交。""" + session = _FakeSession() + calls: list[str] = [] + + async def fake_disable(*args: object, **kwargs: object) -> object: + client_id = args[1] + calls.append(client_id) + if client_id == 'cli_bad': + raise OAuthClientManagementError('client not found') + return SimpleNamespace() + + monkeypatch.setattr(OAuthClientManagementService, '_soft_disable', fake_disable) + response = await delete_system_oauth_clients(_request(), 'cli_good,cli_bad', session, _user()) + assert b'false' in response.body + assert calls == ['cli_good', 'cli_bad'] + assert session.commits == 0 + assert session.rollbacks == 1 + + +@pytest.mark.asyncio +async def test_cors_snapshot_refresh_failure_fails_closed(monkeypatch: pytest.MonkeyPatch) -> None: + """提交后 CORS 快照刷新失败时清空旧快照,确保失败闭合。""" + request = _request() + + async def fail_refresh(*args: object, **kwargs: object) -> tuple[str, ...]: + raise RuntimeError('database unavailable') + + monkeypatch.setattr(OidcRuntimeService, 'refresh_cors_snapshot', fail_refresh) + await OidcRuntimeService.cors_snapshot_callback(request.scope['app'])() + assert request.app.state.oidc_registered_cors_origins == () + + +@pytest.mark.asyncio +async def test_client_list_uses_real_total_and_actor_failure_is_mapped( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """列表返回服务层全量计数,actor 缺失时稳定返回业务失败。""" + session = _FakeSession() + + async def fake_rows(*args: object, **kwargs: object) -> list[str]: + return ['row-1', 'row-2'] + + async def fake_count(*args: object, **kwargs: object) -> int: + return 37 + + monkeypatch.setattr(OAuthClientManagementService, 'list_clients', fake_rows) + monkeypatch.setattr(OAuthClientManagementService, 'count_clients', fake_count) + response = await get_system_oauth_client_list(SimpleNamespace(), session) + assert b'"total":37' in response.body + assert b'"rows":["row-1","row-2"]' in response.body + + failed = await add_system_oauth_client(_request(), SimpleNamespace(), session, SimpleNamespace(user=None)) + assert b'false' in failed.body + assert b'actor' not in failed.body + assert session.rollbacks == 0 diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_resource_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_resource_controller.py new file mode 100644 index 000000000..a0ff756dc --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_resource_controller.py @@ -0,0 +1,179 @@ +from types import SimpleNamespace + +import pytest +from fastapi import Request +from fastapi.routing import APIRoute +from starlette.responses import Response + +from common.context import RequestContext +from module_admin.service.log_service import LogQueueService +from module_identity.controller.oauth_resource_controller import ( + delete_system_oauth_resources, + delete_system_oauth_scopes, + get_system_oauth_resource_list, + get_system_oauth_scope_list, + oauth_resource_controller, + oauth_scope_controller, +) +from module_identity.service.oauth_management_service import ( + OAuthClientManagementError, + OAuthResourceManagementService, +) + + +class _FakeSession: + """记录 Controller 提交和回滚次数。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +@pytest.fixture(autouse=True) +def _request_context(monkeypatch: pytest.MonkeyPatch) -> None: + """为日志装饰器提供当前用户上下文。""" + monkeypatch.setattr(RequestContext, 'get_current_user', staticmethod(_user)) + + async def enqueue_operation_log(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', enqueue_operation_log) + + +def _user() -> SimpleNamespace: + """构造最小当前用户对象。""" + return SimpleNamespace(user=SimpleNamespace(user_name='admin', dept=SimpleNamespace(dept_name='测试部门'))) + + +def _request() -> Request: + """构造管理写请求。""" + app = SimpleNamespace(state=SimpleNamespace()) + + async def receive() -> dict[str, object]: + return {'type': 'http.request', 'body': b'', 'more_body': False} + + return Request( + { + 'type': 'http', + 'method': 'POST', + 'path': '/', + 'headers': [], + 'query_string': b'', + 'app': app, + }, + receive=receive, + ) + + +def _route_map(router: object) -> dict[tuple[str, str], APIRoute]: + """构造路由路径和方法索引。""" + routes = [route for route in router.routes if isinstance(route, APIRoute)] + return {(route.path, method): route for route in routes for method in route.methods} + + +def _assert_routes(router: object, expected: dict[tuple[str, str], str]) -> None: + """断言路由路径、方法和权限依赖。""" + routes = _route_map(router) + assert set(routes) == set(expected) + for key, permission in expected.items(): + dependency_calls = [getattr(item.call, 'perm', None) for item in routes[key].dependant.dependencies] + assert permission in dependency_calls + assert router.dependencies + assert router.dependencies[0].dependency.__class__.__name__ == 'PreAuth' + + +def test_resource_and_scope_router_paths_methods_permissions_and_pre_auth() -> None: + """Resource 与 Scope 路由必须完整覆盖规范并隔离权限前缀。""" + _assert_routes( + oauth_resource_controller, + { + ('/system/oauth/resource/list', 'GET'): 'system:oauthResource:list', + ('/system/oauth/resource/{resource_id}', 'GET'): 'system:oauthResource:list', + ('/system/oauth/resource', 'POST'): 'system:oauthResource:add', + ('/system/oauth/resource', 'PUT'): 'system:oauthResource:edit', + ('/system/oauth/resource/{resource_ids}', 'DELETE'): 'system:oauthResource:remove', + ('/system/oauth/resource/changeStatus', 'PUT'): 'system:oauthResource:edit', + }, + ) + _assert_routes( + oauth_scope_controller, + { + ('/system/oauth/scope/list', 'GET'): 'system:oauthScope:list', + ('/system/oauth/scope/{scope_code}', 'GET'): 'system:oauthScope:list', + ('/system/oauth/scope', 'POST'): 'system:oauthScope:add', + ('/system/oauth/scope', 'PUT'): 'system:oauthScope:edit', + ('/system/oauth/scope/{scope_codes}', 'DELETE'): 'system:oauthScope:remove', + ('/system/oauth/scope/changeStatus', 'PUT'): 'system:oauthScope:edit', + }, + ) + + +@pytest.mark.asyncio +async def test_resource_batch_failure_rolls_back_without_partial_commit(monkeypatch: pytest.MonkeyPatch) -> None: + """Resource 批量停用任一项失败时整体回滚。""" + session = _FakeSession() + calls: list[str] = [] + + async def fake_disable(*args: object, **kwargs: object) -> object: + resource_id = args[1] + calls.append(resource_id) + if resource_id == 'resource-bad': + raise OAuthClientManagementError('resource not found') + return SimpleNamespace() + + monkeypatch.setattr(OAuthResourceManagementService, '_soft_disable_resource', fake_disable) + response = await delete_system_oauth_resources(_request(), 'resource-good,resource-bad', session, _user()) + assert isinstance(response, Response) + assert b'false' in response.body + assert calls == ['resource-good', 'resource-bad'] + assert session.commits == 0 + assert session.rollbacks == 1 + + +@pytest.mark.asyncio +async def test_scope_batch_rejects_empty_duplicate_and_invalid_encoding() -> None: + """Scope 批量路径拒绝空项、重复项和未解析的百分号编码。""" + session = _FakeSession() + user = _user() + for value in ('scope-a,,scope-b', 'scope-a,scope-a', 'scope-a%2Fb'): + response = await delete_system_oauth_scopes(_request(), value, session, user) + assert b'false' in response.body + assert session.commits == 0 + assert session.rollbacks == 0 + + +@pytest.mark.asyncio +async def test_resource_scope_lists_use_real_total_and_actor_failure_is_mapped( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Resource/Scope 列表使用真实全量计数,缺少 actor 时不逃逸异常。""" + session = _FakeSession() + + async def fake_rows(*args: object, **kwargs: object) -> list[str]: + return ['row-1', 'row-2'] + + async def fake_count(*args: object, **kwargs: object) -> int: + return 37 + + monkeypatch.setattr(OAuthResourceManagementService, 'list_resources', fake_rows) + monkeypatch.setattr(OAuthResourceManagementService, 'count_resources', fake_count) + resource_response = await get_system_oauth_resource_list(SimpleNamespace(), session) + assert b'"total":37' in resource_response.body + assert b'"rows":["row-1","row-2"]' in resource_response.body + + monkeypatch.setattr(OAuthResourceManagementService, 'list_scopes', fake_rows) + monkeypatch.setattr(OAuthResourceManagementService, 'count_scopes', fake_count) + scope_response = await get_system_oauth_scope_list(SimpleNamespace(), session) + assert b'"total":37' in scope_response.body + assert b'"rows":["row-1","row-2"]' in scope_response.body + + failed = await delete_system_oauth_resources(_request(), 'resource-a', session, SimpleNamespace(user=None)) + assert b'false' in failed.body + assert b'actor' not in failed.body + assert session.rollbacks == 0 diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_session_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_session_controller.py new file mode 100644 index 000000000..c88406276 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oauth_session_controller.py @@ -0,0 +1,298 @@ +import json +from http import HTTPStatus +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI, Request +from fastapi.routing import APIRoute +from httpx import ASGITransport, AsyncClient + +from common.annotation.rate_limit_annotation import ApiRateLimit +from common.context import RequestContext +from module_admin.service.log_service import LogQueueService +from module_identity.controller.oauth_session_controller import ( + get_oauth_grant, + oauth_grant_controller, + oauth_session_controller, + revoke_oauth_grants, + revoke_oauth_sessions, +) +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.service.audit_service import AuditService +from module_identity.service.oauth_session_management_service import OAuthSessionManagementService +from module_identity.service.session_service import SsoSessionService + + +class _Session: + """记录 Controller 事务边界。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +def _user() -> SimpleNamespace: + """构造最小管理用户。""" + return SimpleNamespace(user=SimpleNamespace(user_name='admin', user_id=1, dept=None)) + + +def _request() -> Request: + """构造含最小 Redis 状态的管理请求。""" + app = SimpleNamespace(state=SimpleNamespace(redis=SimpleNamespace())) + + async def receive() -> dict[str, object]: + return {'type': 'http.request', 'body': b'', 'more_body': False} + + return Request( + {'type': 'http', 'method': 'DELETE', 'path': '/', 'query_string': b'', 'headers': [], 'app': app}, receive + ) + + +def _route_map(router: object) -> dict[tuple[str, str], APIRoute]: + """索引路由路径、方法。""" + return { + (route.path, method): route + for route in router.routes + if isinstance(route, APIRoute) + for method in route.methods + } + + +def test_session_and_grant_routes_have_pre_auth_and_exact_permissions() -> None: + """Session/Grant 路由必须使用规范权限。""" + expected = { + ('/system/oauth/session/list', 'GET'): 'system:oauthSession:list', + ('/system/oauth/session/{sid}', 'GET'): 'system:oauthSession:list', + ('/system/oauth/session/{sids}', 'DELETE'): 'system:oauthSession:revoke', + ('/system/oauth/session/user/{user_id}', 'DELETE'): 'system:oauthSession:revoke', + } + routes = _route_map(oauth_session_controller) + for key, perm in expected.items(): + assert perm in [getattr(dep.call, 'perm', None) for dep in routes[key].dependant.dependencies] + assert oauth_session_controller.dependencies[0].dependency.__class__.__name__ == 'PreAuth' + expected_grants = { + ('/system/oauth/grant/list', 'GET'): 'system:oauthGrant:list', + ('/system/oauth/grant/access/list', 'GET'): 'system:oauthGrant:list', + ('/system/oauth/grant/{grant_id}', 'GET'): 'system:oauthGrant:list', + ('/system/oauth/grant/{grant_ids}', 'DELETE'): 'system:oauthGrant:revoke', + ('/system/oauth/grant/user/{user_id}/client/{client_id}/access', 'PUT'): 'system:oauthGrant:revoke', + } + routes = _route_map(oauth_grant_controller) + for key, perm in expected_grants.items(): + assert perm in [getattr(dep.call, 'perm', None) for dep in routes[key].dependant.dependencies] + + +@pytest.mark.asyncio +async def test_batch_revoke_is_atomic_and_reason_does_not_leak(monkeypatch: pytest.MonkeyPatch) -> None: + """批量撤销失败时回滚且不回显内部敏感字段。""" + db = _Session() + + async def audit(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr(AuditService, 'record', audit) + + async def enqueue(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', enqueue) + + async def revoke(*args: object, **kwargs: object) -> bool: + if args[2] == 'bad': + raise ValueError('refresh_token_hash=do-not-leak') + return True + + monkeypatch.setattr(SsoSessionService, 'revoke', revoke) + token = RequestContext.set_current_user(_user()) + try: + response = await revoke_oauth_sessions.__wrapped__( + _request(), 'good,bad', SimpleNamespace(reason='管理员操作'), db, _user() + ) + finally: + RequestContext.reset_current_user(token) + assert b'false' in response.body + assert b'refresh_token_hash' not in response.body + assert db.commits == 0 and db.rollbacks == 1 + + async def grant_revoke(*args: object, **kwargs: object) -> bool: + return True + + monkeypatch.setattr(OAuthGrantDao, 'targets', AsyncMock(return_value=[(2, 1)])) + monkeypatch.setattr(OAuthAccessPolicyDao, 'lock_client', AsyncMock()) + monkeypatch.setattr(OAuthClientDao, 'id_for_pk', AsyncMock(return_value='client-2')) + monkeypatch.setattr(OAuthGrantDao, 'revoke_for_user_client', AsyncMock(return_value=['grant-1'])) + token = RequestContext.set_current_user(_user()) + try: + response = await revoke_oauth_grants.__wrapped__( + _request(), 'grant-1', SimpleNamespace(reason='撤销'), db, _user() + ) + finally: + RequestContext.reset_current_user(token) + assert b'true' in response.body + assert db.commits == 1 + + +@pytest.mark.asyncio +async def test_session_controller_injects_redis_without_passing_request_to_service( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Session 管理 Controller 向 Service 传入精确 Redis 依赖,而非 FastAPI Request。""" + db = _Session() + request = _request() + captured: dict[str, object] = {} + + async def revoke_sessions(*args: object, **kwargs: object) -> int: + captured['args'] = args + captured['kwargs'] = kwargs + return 1 + + monkeypatch.setattr(OAuthSessionManagementService, 'revoke_sessions', revoke_sessions) + + async def enqueue_operation_log(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', enqueue_operation_log) + token = RequestContext.set_current_user(_user()) + try: + response = await revoke_oauth_sessions.__wrapped__( + request, 'sid-1', SimpleNamespace(reason='管理员操作'), db, _user() + ) + finally: + RequestContext.reset_current_user(token) + assert b'true' in response.body + assert captured['args'][1] is request.app.state.redis + assert not isinstance(captured['args'][1], Request) + + +@pytest.mark.asyncio +async def test_grant_detail_never_returns_refresh_hash(monkeypatch: pytest.MonkeyPatch) -> None: + """Grant 详情只返回安全 DTO。""" + row = SimpleNamespace( + grant_id='g1', + user_id=1, + subject_id='s1', + client_pk=2, + granted_scopes=['openid'], + granted_resources=[], + status='active', + consented_at=None, + expires_at=None, + refresh_token_hash='must-not-return', + ) + + async def get_grant(*args: object, **kwargs: object) -> object: + return row + + async def execute(*args: object, **kwargs: object) -> object: + return SimpleNamespace(all=list, scalars=lambda: SimpleNamespace(all=list, first=lambda: None)) + + monkeypatch.setattr(OAuthGrantDao, 'get_by_grant_id', get_grant) + + async def client_id(*args: object, **kwargs: object) -> str: + return 'client-2' + + monkeypatch.setattr(OAuthClientDao, 'id_for_pk', client_id) + db = _Session() + db.execute = execute + response = await get_oauth_grant('g1', db) + assert b'refresh_token_hash' not in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize('grant_ids', ['grant-1', 'grant-1,grant-2']) +async def test_grant_revoke_runs_rate_limit_and_operation_log(monkeypatch: pytest.MonkeyPatch, grant_ids: str) -> None: + """HTTP 撤销请求经过完整装饰器链并记录实际请求和撤销结果。""" + + db = _Session() + user = _user() + path = f'/system/oauth/grant/{grant_ids}' + route = _route_map(oauth_grant_controller)[('/system/oauth/grant/{grant_ids}', 'DELETE')] + app = FastAPI() + app.state.redis = SimpleNamespace() + app.include_router(oauth_grant_controller) + for dependency in route.dependant.dependencies: + if dependency.name == 'query_db': + app.dependency_overrides[dependency.call] = lambda: db + elif dependency.name == 'current_user': + app.dependency_overrides[dependency.call] = lambda: user + else: + app.dependency_overrides[dependency.call] = lambda: None + rate_limit = AsyncMock(return_value={'allowed': True, 'remaining': 9, 'reset_at': 60}) + expected_ids = grant_ids.split(',') + revoke = AsyncMock(return_value=expected_ids) + audit = AsyncMock() + operation_log = AsyncMock() + monkeypatch.setattr(ApiRateLimit, '_acquire_rate_limit', rate_limit) + monkeypatch.setattr(OAuthGrantDao, 'targets', AsyncMock(return_value=[(2, 1)])) + monkeypatch.setattr(OAuthAccessPolicyDao, 'lock_client', AsyncMock()) + monkeypatch.setattr(OAuthClientDao, 'id_for_pk', AsyncMock(return_value='client-2')) + monkeypatch.setattr(OAuthGrantDao, 'revoke_for_user_client', revoke) + monkeypatch.setattr(AuditService, 'record', audit) + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', operation_log) + token = RequestContext.set_current_user(user) + try: + async with AsyncClient(transport=ASGITransport(app=app), base_url='http://testserver') as client: + response = await client.request('DELETE', path, json={'reason': '管理员撤销'}) + finally: + RequestContext.reset_current_user(token) + expected_ids = grant_ids.split(',') + result = response.json() + assert response.status_code == HTTPStatus.OK + assert result['success'] is True + assert result['data']['count'] == len(expected_ids) + revoke.assert_awaited_once_with(db, 1, 2, '管理员撤销') + assert audit.await_count == len(expected_ids) + assert db.commits == 1 and db.rollbacks == 0 + rate_limit.assert_awaited_once() + operation_log.assert_awaited_once() + logged_request, logged_operation, _ = operation_log.call_args.args + assert isinstance(logged_request, Request) + assert logged_request is rate_limit.call_args.args[-1] + assert logged_operation.oper_url == path + assert json.loads(logged_operation.json_result)['data']['count'] == len(expected_ids) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('blocked', [True, False]) +async def test_access_policy_runs_rate_limit_and_operation_log(monkeypatch: pytest.MonkeyPatch, blocked: bool) -> None: + """禁止及解除访问通过真实 HTTP 装饰器链,准确传递用户、应用和原因。""" + + db, user = _Session(), _user() + path = '/system/oauth/grant/user/7/client/business-app/access' + route = _route_map(oauth_grant_controller)[('/system/oauth/grant/user/{user_id}/client/{client_id}/access', 'PUT')] + app = FastAPI() + app.state.redis = SimpleNamespace() + app.include_router(oauth_grant_controller) + for dependency in route.dependant.dependencies: + if dependency.name == 'query_db': + app.dependency_overrides[dependency.call] = lambda: db + elif dependency.name == 'current_user': + app.dependency_overrides[dependency.call] = lambda: user + else: + app.dependency_overrides[dependency.call] = lambda: None + rate_limit = AsyncMock(return_value={'allowed': True, 'remaining': 9, 'reset_at': 60}) + save = AsyncMock(return_value=1 if blocked else 0) + operation_log = AsyncMock() + monkeypatch.setattr(ApiRateLimit, '_acquire_rate_limit', rate_limit) + monkeypatch.setattr(OAuthSessionManagementService, 'set_access', save) + monkeypatch.setattr(LogQueueService, 'enqueue_operation_log', operation_log) + token = RequestContext.set_current_user(user) + try: + async with AsyncClient(transport=ASGITransport(app=app), base_url='http://testserver') as client: + response = await client.put(path, json={'blocked': blocked, 'reason': '管理员操作'}) + finally: + RequestContext.reset_current_user(token) + assert response.status_code == HTTPStatus.OK and response.json()['success'] is True + assert response.json()['data']['accessStatus'] == ('blocked' if blocked else 'allowed') + save.assert_awaited_once_with(db, 7, 'business-app', blocked, 'admin', '管理员操作') + operation_log.assert_awaited_once() + assert operation_log.call_args.args[0] is rate_limit.call_args.args[-1] diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oidc_key_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oidc_key_controller.py new file mode 100644 index 000000000..7e9d8916f --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_oidc_key_controller.py @@ -0,0 +1,132 @@ +from datetime import datetime, timezone +from types import SimpleNamespace + +import pytest +from fastapi import Request +from fastapi.routing import APIRoute + +from exceptions.exception import ServiceException +from module_identity.controller.oidc_key_controller import ( + activate_oidc_key, + delete_oidc_key, + list_oidc_keys, + oidc_key_controller, + retire_oidc_key, +) +from module_identity.service.key_service import KeyService, KeyServiceError, OidcKeyManagementService +from module_identity.service.runtime_service import OidcReadiness, OidcRuntimeService + + +class _Session: + """记录提交和回滚。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +def _user() -> SimpleNamespace: + """构造管理用户。""" + return SimpleNamespace(user=SimpleNamespace(user_name='admin')) + + +def _request() -> Request: + """构造密钥管理写请求。""" + app = SimpleNamespace(state=SimpleNamespace()) + return Request({'type': 'http', 'method': 'PUT', 'path': '/', 'headers': [], 'app': app}) + + +def test_key_routes_use_separate_permissions_for_rotation_and_activation() -> None: + """轮换与激活分别受独立权限保护,删除也必须受退役权限保护。""" + routes = { + (route.path, method): route + for route in oidc_key_controller.routes + if isinstance(route, APIRoute) + for method in route.methods + } + assert ('/system/oauth/key/{kid}/activate', 'PUT') in routes + assert ('/system/oauth/key/{kid}/retire', 'PUT') in routes + assert ('/system/oauth/key/{kid}', 'DELETE') in routes + activate_perms = [ + getattr(dep.call, 'perm', None) + for dep in routes[('/system/oauth/key/{kid}/activate', 'PUT')].dependant.dependencies + ] + assert 'system:oauthKey:activate' in activate_perms + assert 'system:oauthKey:rotate' not in activate_perms + assert oidc_key_controller.dependencies[0].dependency.__class__.__name__ == 'PreAuth' + + +@pytest.mark.asyncio +async def test_key_write_failures_rollback_and_never_return_private_fields(monkeypatch: pytest.MonkeyPatch) -> None: + """密钥状态异常稳定回滚,响应不包含私钥引用。""" + db = _Session() + + async def fail(*args: object, **kwargs: object) -> bool: + raise ValueError('private_key_ciphertext=secret') + + monkeypatch.setattr(KeyService, 'activate_key', fail) + with pytest.raises(ServiceException) as exc_info: + await activate_oidc_key.__wrapped__.__wrapped__(_request(), 'kid-1', db, _user()) + assert 'private_key_ciphertext' not in exc_info.value.message + assert db.rollbacks == 1 and db.commits == 0 + + async def retire(*args: object, **kwargs: object) -> bool: + return True + + monkeypatch.setattr(KeyService, 'retire_key', retire) + response = await retire_oidc_key.__wrapped__.__wrapped__(_request(), 'kid-1', db, _user()) + assert b'true' in response.body + assert db.commits == 1 + + monkeypatch.setattr(KeyService, 'delete_key', retire) + response = await delete_oidc_key.__wrapped__.__wrapped__(_request(), 'kid-1', db, _user()) + assert b'true' in response.body + + +@pytest.mark.asyncio +async def test_key_activation_returns_clear_message_before_publication(monkeypatch: pytest.MonkeyPatch) -> None: + """密钥尚未公开时返回可操作提示,不暴露内部异常。""" + db = _Session() + + async def unpublished(*args: object, **kwargs: object) -> bool: + raise KeyServiceError('目标签名密钥尚未发布') + + monkeypatch.setattr(KeyService, 'activate_key', unpublished) + with pytest.raises(ServiceException) as exc_info: + await activate_oidc_key.__wrapped__.__wrapped__(_request(), 'kid-1', db, _user()) + assert '签名公钥尚未到公开时间' in exc_info.value.message + assert db.rollbacks == 1 and db.commits == 0 + + +@pytest.mark.asyncio +async def test_key_list_exposes_only_safe_enabled_capability(monkeypatch: pytest.MonkeyPatch) -> None: + """密钥列表向管理页暴露启用能力,但不阻断 OIDC 关闭时的引导。""" + + class _Db: + async def rollback(self) -> None: + pass + + async def list_page(*args: object, **kwargs: object) -> tuple[list[object], int]: + return [], 0 + + monkeypatch.setattr(OidcKeyManagementService, 'list_page', list_page) + monkeypatch.setattr( + OidcRuntimeService, + 'inspect_readiness', + lambda db: _async_readiness(OidcReadiness(False, False, 'disabled', datetime.now(timezone.utc))), + ) + response = await list_oidc_keys(_Db()) + assert b'"enabled":false' in response.body + assert b'"ready":false' in response.body + assert b'private_key' not in response.body + + +async def _async_readiness(value: OidcReadiness) -> OidcReadiness: + """返回可注入控制器的异步就绪状态。""" + return value diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_token_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_token_controller.py new file mode 100644 index 000000000..d40453583 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_token_controller.py @@ -0,0 +1,297 @@ +import json +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock + +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from fastapi import HTTPException +from jwt.algorithms import RSAAlgorithm + +from exceptions.exception import OAuthProtocolException +from module_identity.controller import token_controller as controller +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.token_service import RefreshTokenReuseDetected, TokenResult, TokenService + +_OK = 200 +_BAD_REQUEST = 400 +_UNSUPPORTED_MEDIA = 415 +_UNAUTHORIZED = 401 +_SERVER_ERROR = 500 +_EXPECTED_CONTROLLER_COMMITS = 2 + + +@pytest.fixture(autouse=True) +def _disable_rate_limit_for_controller_unit_tests(monkeypatch: pytest.MonkeyPatch) -> None: + """Controller 单元测试隔离 Redis 限流实现;限流器另行测试。""" + + async def allow(*_args: object, **_kwargs: object) -> None: + return None + + monkeypatch.setattr(controller.OidcRateLimiter, 'enforce', allow) + + +class _Db: + """记录 Controller 事务边界的最小异步会话替身。""" + + def __init__(self) -> None: + self.commits = 0 + self.rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + +class _FailCommitDb(_Db): + """提交失败并允许回滚的会话替身。""" + + async def commit(self) -> None: + self.commits += 1 + raise RuntimeError('commit failed') + + +def _request() -> SimpleNamespace: + """构造携带应用 Redis 状态的请求替身。""" + return SimpleNamespace( + headers={'authorization': 'Basic test'}, + app=SimpleNamespace(state=SimpleNamespace(redis=object())), + ) + + +@pytest.mark.asyncio +async def test_token_success_is_bare_json_and_commits(monkeypatch: pytest.MonkeyPatch) -> None: + """Token 成功响应提交事务且带 no-store。""" + db = _Db() + monkeypatch.setattr(controller, 'read_form', lambda request: _async(_form())) + monkeypatch.setattr( + TokenService, + 'authenticate_client', + lambda *args, **kwargs: _async((SimpleNamespace(), OAuthClientPrincipal('c', 'public', 'none'))), + ) + monkeypatch.setattr( + TokenService, + 'issue_token', + lambda *args, **kwargs: _async(TokenResult('access', 60, scope='openid')), + ) + response = await controller.token(_request(), db) + assert response.status_code == _OK + assert response.headers['cache-control'] == 'no-store' + assert db.commits == 1 + assert db.rollbacks == 0 + + +@pytest.mark.asyncio +async def test_token_invalid_client_has_basic_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + """Token invalid_client 必须返回 RFC Basic challenge。""" + db = _Db() + monkeypatch.setattr(controller, 'read_form', lambda request: _async(_form())) + + async def invalid(*args: object, **kwargs: object) -> object: + raise OAuthProtocolException('invalid_client', 'Client authentication failed', 401) + + monkeypatch.setattr(TokenService, 'authenticate_client', invalid) + response = await controller.token(_request(), db) + assert response.status_code == _UNAUTHORIZED + assert response.headers['www-authenticate'] == 'Basic realm="oauth2/token"' + + +@pytest.mark.asyncio +async def test_token_rejects_oversized_or_basic_plus_body_secret_before_issue( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Token Endpoint 不允许绕过 Secret 长度、位置和互斥边界。""" + db = _Db() + authenticate = AsyncMock(side_effect=OAuthProtocolException('invalid_client', 'Client authentication failed', 401)) + issue = AsyncMock() + monkeypatch.setattr(TokenService, 'authenticate_client', authenticate) + monkeypatch.setattr(TokenService, 'issue_token', issue) + monkeypatch.setattr( + controller, + 'read_form', + lambda request: _async({'grant_type': 'client_credentials', 'client_id': 'c', 'client_secret': 'x' * 513}), + ) + response = await controller.token(_request(), db) + assert response.status_code == _BAD_REQUEST + authenticate.assert_not_awaited() + + authenticate.reset_mock() + monkeypatch.setattr( + controller, + 'read_form', + lambda request: _async({'grant_type': 'client_credentials', 'client_id': 'c', 'client_secret': 'body'}), + ) + response = await controller.token(_request(), db) + assert response.status_code == _UNAUTHORIZED + authenticate.assert_awaited_once() + issue.assert_not_awaited() + + authenticate.reset_mock() + request_without_basic = _request() + request_without_basic.headers = {} + response = await controller.token(request_without_basic, db) + assert response.status_code == _UNAUTHORIZED + authenticate.assert_awaited_once() + assert authenticate.await_args.kwargs['authorization'] is None + assert authenticate.await_args.kwargs['client_secret'] == 'body' + issue.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_refresh_reuse_commits_family_before_invalid_grant(monkeypatch: pytest.MonkeyPatch) -> None: + """Refresh 重放异常必须提交 Family 撤销,而非回滚。""" + db = _Db() + monkeypatch.setattr(controller, 'read_form', lambda request: _async(_form())) + monkeypatch.setattr( + TokenService, + 'authenticate_client', + lambda *args, **kwargs: _async((SimpleNamespace(), OAuthClientPrincipal('c', 'public', 'none'))), + ) + + async def reuse(*args: object, **kwargs: object) -> TokenResult: + raise RefreshTokenReuseDetected + + monkeypatch.setattr(TokenService, 'issue_token', reuse) + response = await controller.token(_request(), db) + assert response.status_code == _BAD_REQUEST + assert response.body.startswith(b'{"error":"invalid_grant"') + assert db.commits == 1 + assert db.rollbacks == 0 + + +@pytest.mark.asyncio +async def test_refresh_reuse_commit_failure_is_server_error_and_rolls_back(monkeypatch: pytest.MonkeyPatch) -> None: + """Family 提交失败时不得伪装成已落盘的 invalid_grant。""" + db = _FailCommitDb() + monkeypatch.setattr(controller, 'read_form', lambda request: _async(_form())) + monkeypatch.setattr( + TokenService, + 'authenticate_client', + lambda *args, **kwargs: _async((SimpleNamespace(), OAuthClientPrincipal('c', 'public', 'none'))), + ) + + async def reuse(*args: object, **kwargs: object) -> TokenResult: + raise RefreshTokenReuseDetected + + monkeypatch.setattr(TokenService, 'issue_token', reuse) + response = await controller.token(_request(), db) + assert response.status_code == _SERVER_ERROR + assert b'invalid_grant' not in response.body + assert db.rollbacks == 1 + + +@pytest.mark.asyncio +async def test_introspect_rejects_public_client_with_basic_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + """Introspection 在 Controller 边界拒绝公共 Client。""" + db = _Db() + monkeypatch.setattr(controller, 'read_form', lambda request: _async({'client_id': 'c', 'token': 'x'})) + monkeypatch.setattr( + TokenService, + 'authenticate_client', + lambda *args, **kwargs: _async( + (SimpleNamespace(client_type='public'), OAuthClientPrincipal('c', 'public', 'none')) + ), + ) + response = await controller.introspect(_request(), db) + assert response.status_code == _UNAUTHORIZED + assert response.headers['www-authenticate'] == 'Basic realm="oauth2/token"' + assert db.rollbacks == 1 + + +@pytest.mark.asyncio +@pytest.mark.parametrize('status_code', [_BAD_REQUEST, _UNSUPPORTED_MEDIA]) +async def test_introspect_form_boundary_preserves_http_error(monkeypatch: pytest.MonkeyPatch, status_code: int) -> None: + """Introspection 的重复字段和媒体类型错误不得被转换为 500。""" + db = _Db() + + async def invalid_form(request: object) -> dict[str, str]: + raise HTTPException(status_code=status_code, detail='invalid input') + + monkeypatch.setattr(controller, 'read_form', invalid_form) + response = await controller.introspect(_request(), db) + assert response.status_code == status_code + assert response.body == b'{"error":"invalid_request"}' + assert db.rollbacks == 0 + + +@pytest.mark.asyncio +async def test_revoke_and_introspect_receive_database_rsa_verification_key(monkeypatch: pytest.MonkeyPatch) -> None: + """撤销和内省路由把按 kid 从数据库加载的 RSA 公钥传给服务。""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + record = SimpleNamespace( + kid='kid-route', + alg='RS256', + status='active', + publish_at=None, + remove_from_jwks_at=None, + public_jwk=json.loads(RSAAlgorithm.to_jwk(key.public_key())), + ) + + class KeyDb(_Db): + async def scalar(self, statement: object) -> object: + return record + + token = jwt.encode( + { + 'iss': 'https://issuer.example', + 'sub': 'subject', + 'aud': ['https://issuer.example/oauth2/userinfo'], + 'exp': 2, + 'iat': 1, + 'jti': 'jti', + 'client_id': 'c', + 'scope': 'openid', + }, + key, + algorithm='RS256', + headers={'kid': 'kid-route', 'typ': 'at+jwt'}, + ) + db = KeyDb() + captured: list[object] = [] + monkeypatch.setattr('module_identity.dependencies.OidcKeyDao.get_verifying', lambda *args: _async(record)) + monkeypatch.setattr(controller, 'read_form', lambda request: _async({'client_id': 'c', 'token': token})) + monkeypatch.setattr( + TokenService, + 'authenticate_client', + lambda *args, **kwargs: _async( + ( + SimpleNamespace(client_type='confidential'), + OAuthClientPrincipal('c', 'confidential', 'client_secret_basic'), + ) + ), + ) + + async def revoke_service(*args: object, **kwargs: object) -> None: + captured.append(kwargs['verification_key']) + + monkeypatch.setattr(controller.RevocationService, 'revoke', revoke_service) + await controller.revoke(_request(), db) + assert captured and captured[0].public_numbers() == key.public_key().public_numbers() + + captured.clear() + + async def introspect_service(*args: object, **kwargs: object) -> dict[str, bool]: + captured.append(kwargs['verification_key']) + return {'active': False} + + monkeypatch.setattr(controller.IntrospectionService, 'introspect', introspect_service) + await controller.introspect(_request(), db) + assert captured and captured[0].public_numbers() == key.public_key().public_numbers() + assert db.commits == _EXPECTED_CONTROLLER_COMMITS + + +def _form() -> dict[str, str]: + """构造 Controller 测试表单。""" + return {'grant_type': 'client_credentials', 'client_id': 'c'} + + +def _async(value: object) -> Any: + """构造异步测试结果。""" + + async def result() -> object: + return value + + return result() diff --git a/ruoyi-fastapi-backend/tests/module_identity/controllers/test_userinfo_controller.py b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_userinfo_controller.py new file mode 100644 index 000000000..4ccb55468 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/controllers/test_userinfo_controller.py @@ -0,0 +1,247 @@ +from types import SimpleNamespace +from typing import Any + +import pytest + +from module_admin.entity.do.user_do import SysUser +from module_identity.controller import token_controller as controller +from module_identity.dependencies import AccessTokenContext +from module_identity.service import token_protocol_service as service + +_OK = 200 +_NOT_FOUND = 404 +_UNAUTHORIZED = 401 + + +def _request(redis: object | None = None) -> SimpleNamespace: + """构造 UserInfo 请求替身。""" + return SimpleNamespace(app=SimpleNamespace(state=SimpleNamespace(redis=redis or object()))) + + +class _ActiveRedis: + """模拟未命中撤销列表的 Redis。""" + + async def exists(self, key: str) -> bool: + return False + + +@pytest.fixture(autouse=True) +def _active_authorization(monkeypatch: pytest.MonkeyPatch) -> None: + """为 Claims 输出测试提供明确有效的授权,撤销链路由真实数据库回归覆盖。""" + + grant = SimpleNamespace( + user_id=2, + subject_id='subject-1', + client_pk=1, + status='active', + client_policy_version=1, + expires_at=None, + granted_scopes=['openid', 'profile', 'roles', 'dept'], + granted_resources=[], + ) + monkeypatch.setattr(service.OAuthGrantDao, 'get_by_grant_id', lambda *args, **kwargs: _async(grant)) + monkeypatch.setattr(service.OAuthAccessPolicyDao, 'is_blocked', lambda *args, **kwargs: _async(False)) + monkeypatch.setattr(service.IntrospectionService, '_client_allows_access', lambda *args, **kwargs: _async(True)) + + +@pytest.mark.asyncio +async def test_userinfo_returns_only_minimal_claims(monkeypatch: pytest.MonkeyPatch) -> None: + """UserInfo 按当前数据库 Scope、角色和部门生成最小 Claims。""" + monkeypatch.setattr(controller.OidcConfig, 'oidc_enabled', True) + client = SimpleNamespace(client_pk=1, client_id='client-1', status='0', policy_version=1) + user = SysUser(user_id=2, user_name='alice', nick_name='数据库用户', status='0', del_flag='0') + subject = SimpleNamespace(subject_id='subject-1', user_id=2, auth_version=3) + session = SimpleNamespace(user_id=2, subject_id='subject-1', auth_version=3) + definitions = [ + SimpleNamespace(scope_pk=1, scope_code='openid', claims=['sub'], status='0'), + SimpleNamespace(scope_pk=2, scope_code='profile', claims=['name'], status='0'), + SimpleNamespace(scope_pk=3, scope_code='roles', claims=['roles'], status='0'), + SimpleNamespace(scope_pk=4, scope_code='dept', claims=['dept_id', 'dept_name'], status='0'), + ] + bindings = [ + SimpleNamespace( + scope_pk=1, + claim_filter={ + 'claims': ['sub', 'name', 'roles', 'dept_id', 'dept_name'], + 'allowed_role_keys': ['database-role'], + }, + ), + SimpleNamespace(scope_pk=2, claim_filter=['name']), + SimpleNamespace( + scope_pk=3, + claim_filter={ + 'claims': ['sub', 'name', 'roles', 'dept_id', 'dept_name'], + 'allowed_role_keys': ['database-role'], + }, + ), + SimpleNamespace( + scope_pk=4, + claim_filter={ + 'claims': ['sub', 'name', 'roles', 'dept_id', 'dept_name'], + 'allowed_role_keys': ['database-role'], + }, + ), + ] + + monkeypatch.setattr(service.OAuthClientDao, 'get_by_client_id', lambda *args, **kwargs: _async(client)) + monkeypatch.setattr(service.IdentitySubjectDao, 'get_by_subject_id', lambda *args, **kwargs: _async(subject)) + monkeypatch.setattr(service.IdentityUserDao, 'get_user', lambda *args, **kwargs: _async(user)) + monkeypatch.setattr(service.SsoSessionDao, 'get_active', lambda *args, **kwargs: _async(session)) + monkeypatch.setattr( + service.OAuthClientDao, + 'list_scope_bindings', + lambda *args, **kwargs: _async(bindings), + ) + monkeypatch.setattr( + service.OAuthClientDao, + 'list_scope_definitions', + lambda *args, **kwargs: _async(definitions), + ) + monkeypatch.setattr( + service.ClaimService, + 'load_roles_and_department', + lambda *args, **kwargs: _async((['database-role'], SimpleNamespace(dept_id=7, dept_name='数据库部门'))), + ) + context = AccessTokenContext( + 'token', + { + 'sub': 'subject-1', + 'scope': 'openid profile roles dept', + 'name': '伪造名称', + 'roles': ['伪造角色'], + 'user_id': 123, + 'sid': 'sid-1', + 'grant_id': 'grant-1', + 'client_id': 'client-1', + 'jti': 'jti-1', + 'ver': 3, + 'exp': 1, + }, + ) + response = await controller.userinfo(_request(_ActiveRedis()), context, SimpleNamespace()) + assert response.status_code == _OK + assert b'"sub":"subject-1"' in response.body + assert '数据库用户'.encode() in response.body + assert '数据库部门'.encode() in response.body + assert b'database-role' in response.body + assert '伪造名称'.encode() not in response.body + assert '伪造角色'.encode() not in response.body + assert b'user_id' not in response.body + assert b'"sid"' not in response.body + + +@pytest.mark.asyncio +async def test_userinfo_uses_current_database_claims_not_token_claims(monkeypatch: pytest.MonkeyPatch) -> None: + """Token 中篡改 name/roles 不得污染按当前数据库策略重建的 UserInfo。""" + monkeypatch.setattr(controller.OidcConfig, 'oidc_enabled', True) + client = SimpleNamespace(client_pk=1, status='0', policy_version=1) + user = SimpleNamespace(user_id=2, user_name='alice', nick_name='当前用户', status='0', del_flag='0') + subject = SimpleNamespace(subject_id='subject-1', user_id=2, auth_version=3) + session = SimpleNamespace(user_id=2, subject_id='subject-1', auth_version=3) + definitions = [ + SimpleNamespace(scope_pk=1, scope_code='openid', claims=['sub'], status='0'), + SimpleNamespace(scope_pk=2, scope_code='roles', claims=['roles'], status='0'), + ] + bindings = [ + SimpleNamespace( + scope_pk=1, + claim_filter={ + 'claims': ['sub', 'name', 'roles', 'dept_id', 'dept_name'], + 'allowed_role_keys': ['database-role'], + }, + ), + SimpleNamespace( + scope_pk=2, + claim_filter={ + 'claims': ['sub', 'name', 'roles', 'dept_id', 'dept_name'], + 'allowed_role_keys': ['database-role'], + }, + ), + ] + monkeypatch.setattr(service.OAuthClientDao, 'get_by_client_id', lambda *args, **kwargs: _async(client)) + monkeypatch.setattr(service.IdentitySubjectDao, 'get_by_subject_id', lambda *args, **kwargs: _async(subject)) + monkeypatch.setattr(service.IdentityUserDao, 'get_user', lambda *args, **kwargs: _async(user)) + monkeypatch.setattr(service.SsoSessionDao, 'get_active', lambda *args, **kwargs: _async(session)) + monkeypatch.setattr(service.OAuthClientDao, 'list_scope_bindings', lambda *args, **kwargs: _async(bindings)) + monkeypatch.setattr(service.OAuthClientDao, 'list_scope_definitions', lambda *args, **kwargs: _async(definitions)) + monkeypatch.setattr( + service.ClaimService, + 'load_roles_and_department', + lambda *args, **kwargs: _async((['database-role'], None)), + ) + context = AccessTokenContext( + 'token', + { + 'sub': 'subject-1', + 'scope': 'openid roles', + 'name': 'Forged Name', + 'roles': ['forged-role'], + 'client_id': 'client-1', + 'sid': 'sid-1', + 'grant_id': 'grant-1', + 'jti': 'jti-1', + 'ver': 3, + }, + ) + response = await controller.userinfo(_request(_ActiveRedis()), context, SimpleNamespace()) + assert b'subject-1' in response.body + assert b'Forged Name' not in response.body + assert b'database-role' in response.body + assert b'forged-role' not in response.body + + +@pytest.mark.asyncio +@pytest.mark.parametrize('failure', ['revoked', 'version', 'session', 'redis']) +async def test_userinfo_live_state_failures_are_bearer_401(monkeypatch: pytest.MonkeyPatch, failure: str) -> None: + """撤销 JTI、版本/Session 失效和 Redis 异常均统一返回 401。""" + monkeypatch.setattr(controller.OidcConfig, 'oidc_enabled', True) + client = SimpleNamespace(client_id='client-1', status='0') + subject = SimpleNamespace(subject_id='subject-1', user_id=2, auth_version=3) + user = SysUser(user_id=2, user_name='alice', nick_name='Alice', status='0', del_flag='0') + monkeypatch.setattr(service.OAuthClientDao, 'get_by_client_id', lambda *args, **kwargs: _async(client)) + monkeypatch.setattr(service.IdentitySubjectDao, 'get_by_subject_id', lambda *args, **kwargs: _async(subject)) + monkeypatch.setattr(service.IdentityUserDao, 'get_user', lambda *args, **kwargs: _async(user)) + + class _Redis: + async def exists(self, key: str) -> bool: + if failure == 'redis': + raise RuntimeError('redis unavailable') + return failure == 'revoked' + + claims = { + 'sub': 'subject-1', + 'scope': 'openid', + 'sid': 'sid-1', + 'client_id': 'client-1', + 'jti': 'jti-1', + 'ver': 4 if failure == 'version' else 3, + } + monkeypatch.setattr( + service.SsoSessionDao, + 'get_active', + lambda *args, **kwargs: _async( + None if failure == 'session' else SimpleNamespace(user_id=2, subject_id='subject-1', auth_version=3) + ), + ) + response = await controller.userinfo(_request(_Redis()), AccessTokenContext('token', claims), SimpleNamespace()) + assert response.status_code == _UNAUTHORIZED + assert response.headers['www-authenticate'] == 'Bearer error="invalid_token"' + assert b'redis unavailable' not in response.body + + +@pytest.mark.asyncio +async def test_userinfo_disabled_is_local_not_found(monkeypatch: pytest.MonkeyPatch) -> None: + """OIDC 关闭时 UserInfo 返回本地 404,不进入重定向流程。""" + monkeypatch.setattr(controller.OidcConfig, 'oidc_enabled', False) + response = await controller.userinfo(_request(), SimpleNamespace(), SimpleNamespace()) + assert response.status_code == _NOT_FOUND + assert b'not_found' in response.body + + +def _async(value: object) -> Any: + """构造异步测试结果。""" + + async def result() -> object: + return value + + return result() diff --git a/ruoyi-fastapi-backend/tests/module_identity/dao/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/dao/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/dao/test_grant_policy_dao.py b/ruoyi-fastapi-backend/tests/module_identity/dao/test_grant_policy_dao.py new file mode 100644 index 000000000..e84aa47b5 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/dao/test_grant_policy_dao.py @@ -0,0 +1,56 @@ +from datetime import datetime, timedelta, timezone + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from module_admin.entity.do.user_do import SysUser +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant + + +@pytest.mark.asyncio +async def test_valid_grant_requires_active_unexpired_and_current_client_policy( + data_session: AsyncSession, +) -> None: + """验证 Grant 过期或 Client 策略版本变化后不再视为有效。""" + now = datetime.now(timezone.utc) + client = SysOAuthClient( + client_pk=4101, + client_id='grant-policy-client', + client_name='Grant Policy Client', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['authorization_code'], + response_types=['code'], + policy_version=7, + ) + data_session.add_all( + [ + SysUser(user_id=4101, user_name='grant-user', nick_name='Grant User', status='0', del_flag='0'), + client, + SysOAuthGrant( + grant_id='grant-policy-1', + user_id=4101, + subject_id='subject-policy-1', + client_pk=4101, + granted_scopes=['openid'], + granted_resources=[], + client_policy_version=7, + status='active', + consented_at=now, + ), + ] + ) + await data_session.flush() + assert await OAuthGrantDao.get_valid_for_user_client(data_session, 4101, 4101) is not None + + grant = await OAuthGrantDao.get_active_for_user_client(data_session, 4101, 4101) + grant.expires_at = now - timedelta(seconds=1) + await data_session.flush() + assert await OAuthGrantDao.get_valid_for_user_client(data_session, 4101, 4101) is None + + grant.expires_at = None + client.policy_version = 8 + await data_session.flush() + assert await OAuthGrantDao.get_valid_for_user_client(data_session, 4101, 4101) is None diff --git a/ruoyi-fastapi-backend/tests/module_identity/dao/test_oidc_key_dao.py b/ruoyi-fastapi-backend/tests/module_identity/dao/test_oidc_key_dao.py new file mode 100644 index 000000000..30725571e --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/dao/test_oidc_key_dao.py @@ -0,0 +1,54 @@ +from datetime import datetime, timezone +from typing import Any + +import pytest +from sqlalchemy import event, select +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.dao.oidc_key_dao import OidcKeyDao +from module_identity.entity.do.oidc_key_do import SysOidcSigningKey + + +def _key(kid: str, status: str = 'pending') -> SysOidcSigningKey: + """构造带引用私钥材料的测试密钥。""" + now = datetime.now(timezone.utc) + return SysOidcSigningKey( + kid=kid, + alg='RS256', + public_jwk={'kty': 'RSA', 'kid': kid}, + private_key_ref=f'kms://{kid}', + status=status, + publish_at=now, + create_by='test', + ) + + +@pytest.mark.asyncio +async def test_activate_serializes_by_algorithm_when_no_active_key_exists(data_session: AsyncSession) -> None: + """验证无 Active Key 时先执行算法范围写锁,再激活唯一目标。""" + data_session.add_all([_key('pending-a'), _key('pending-b')]) + await data_session.flush() + statements: list[str] = [] + + def capture( + _connection: Any, + _cursor: Any, + statement: str, + _parameters: Any, + _context: Any, + _executemany: bool, + ) -> None: + if statement.lstrip().upper().startswith(('UPDATE', 'SELECT')): + statements.append(statement.upper()) + + event.listen(data_session.bind.sync_engine, 'before_cursor_execute', capture) + try: + assert await OidcKeyDao.activate(data_session, 'pending-a', 'RS256') + await data_session.flush() + finally: + event.remove(data_session.bind.sync_engine, 'before_cursor_execute', capture) + + rows = (await data_session.execute(select(SysOidcSigningKey))).scalars().all() + assert [row.status for row in rows].count('active') == 1 + assert statements and statements[0].startswith('SELECT SYS_OIDC_SIGNING_KEY') + assert 'WHERE SYS_OIDC_SIGNING_KEY.ALG' in statements[0] diff --git a/ruoyi-fastapi-backend/tests/module_identity/dao/test_refresh_token_dao.py b/ruoyi-fastapi-backend/tests/module_identity/dao/test_refresh_token_dao.py new file mode 100644 index 000000000..f30fafc8e --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/dao/test_refresh_token_dao.py @@ -0,0 +1,136 @@ +from datetime import datetime, timedelta, timezone + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from module_admin.entity.do.user_do import SysUser +from module_identity.dao.oauth_token_dao import OAuthTokenDao +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken, SysSsoSession + +_REUSED_TOKEN_COUNT = 2 + + +def _token(token_id: str, status: str, now: datetime, idle_expires_at: datetime) -> SysOAuthRefreshToken: + """构造测试用 Refresh Token 元数据。""" + return SysOAuthRefreshToken( + token_id=token_id, + token_hash=f'{token_id:0<64}', + family_id='family-1', + grant_id='grant-1', + user_id=3001, + subject_id='subject-1', + auth_version=1, + client_pk=3001, + sid='sid-1', + scopes=['openid'], + resources=[], + status=status, + issued_at=now, + idle_expires_at=idle_expires_at, + absolute_expires_at=now + timedelta(days=1), + ) + + +@pytest.mark.asyncio +async def test_refresh_family_lock_rotation_and_reuse_distribution(data_session: AsyncSession) -> None: + """验证 Family 查询、单 Token 轮换和重放后的状态分布。""" + now = datetime.now(timezone.utc) + data_session.add_all( + [ + SysUser(user_id=3001, user_name='refresh', nick_name='Refresh', status='0', del_flag='0'), + SysOAuthClient( + client_pk=3001, + client_id='refresh-client', + client_name='Refresh Client', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['authorization_code', 'refresh_token'], + response_types=['code'], + ), + SysOAuthGrant( + grant_id='grant-1', + user_id=3001, + subject_id='subject-1', + client_pk=3001, + granted_scopes=['openid'], + granted_resources=[], + client_policy_version=1, + consented_at=now, + ), + SysSsoSession( + sid='sid-1', + session_secret_hash='a' * 64, + user_id=3001, + subject_id='subject-1', + auth_version=1, + auth_time=now, + last_seen_at=now, + idle_expires_at=now + timedelta(hours=1), + absolute_expires_at=now + timedelta(days=1), + acr='pwd', + amr=['pwd'], + ), + _token('token-1', 'active', now, now + timedelta(hours=1)), + _token('token-2', 'active', now + timedelta(seconds=1), now + timedelta(hours=1)), + ] + ) + await data_session.flush() + + family = await OAuthTokenDao.lock_family(data_session, 'family-1') + assert [token.token_id for token in family] == ['token-1', 'token-2'] + injected_now = datetime(2026, 8, 24, 5, 6, 7, tzinfo=timezone.utc) + assert await OAuthTokenDao.mark_used(data_session, 'token-1', 'token-2', now=injected_now) + assert (await OAuthTokenDao.get_by_token_id(data_session, 'token-1')).last_used_at == injected_now + changed = await OAuthTokenDao.refresh_token_family_reuse(data_session, 'family-1', 'token-1', now=injected_now) + assert changed == _REUSED_TOKEN_COUNT + await data_session.flush() + statuses = {token.token_id: token.status for token in family} + assert statuses == {'token-1': 'reuse_detected', 'token-2': 'revoked'} + + +@pytest.mark.asyncio +async def test_refresh_expire_due_handles_idle_and_absolute_expiry(data_session: AsyncSession) -> None: + """验证闲置和绝对期限任一到期都会使 Active Token 过期。""" + now = datetime.now(timezone.utc) + data_session.add_all( + [ + SysUser(user_id=3001, user_name='expire', nick_name='Expire', status='0', del_flag='0'), + SysOAuthClient( + client_pk=3001, + client_id='expire-client', + client_name='Expire Client', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['refresh_token'], + response_types=['code'], + ), + SysOAuthGrant( + grant_id='grant-1', + user_id=3001, + subject_id='subject-1', + client_pk=3001, + granted_scopes=['openid'], + granted_resources=[], + client_policy_version=1, + consented_at=now, + ), + SysSsoSession( + sid='sid-1', + session_secret_hash='b' * 64, + user_id=3001, + subject_id='subject-1', + auth_version=1, + auth_time=now, + last_seen_at=now, + idle_expires_at=now + timedelta(hours=1), + absolute_expires_at=now + timedelta(days=1), + acr='pwd', + amr=['pwd'], + ), + _token('token-idle', 'active', now, now - timedelta(seconds=1)), + ] + ) + await data_session.flush() + assert await OAuthTokenDao.expire_due(data_session) == 1 + assert (await OAuthTokenDao.get_by_token_id(data_session, 'token-idle')).status == 'expired' diff --git a/ruoyi-fastapi-backend/tests/module_identity/data/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/data/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/data/test_identity_subject.py b/ruoyi-fastapi-backend/tests/module_identity/data/test_identity_subject.py new file mode 100644 index 000000000..653d4e9aa --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/data/test_identity_subject.py @@ -0,0 +1,42 @@ +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from module_admin.entity.do.user_do import SysUser +from module_identity.dao.identity_subject_dao import IdentitySubjectDao + +_NEXT_AUTH_VERSION = 2 +_BACKFILL_COUNT = 2 + + +@pytest.mark.asyncio +async def test_subject_is_stable_and_auth_version_is_atomic(data_session: AsyncSession) -> None: + """验证幂等创建、稳定 Subject 与乐观版本条件。""" + data_session.add(SysUser(user_id=1001, user_name='alice', nick_name='Alice', status='0', del_flag='0')) + await data_session.flush() + + first = await IdentitySubjectDao.create_for_user(data_session, 1001, subject_id='stable-subject') + second = await IdentitySubjectDao.create_for_user(data_session, 1001, subject_id='different-subject') + assert first.identity_id == second.identity_id + assert second.subject_id == 'stable-subject' + assert await IdentitySubjectDao.increment_auth_version(data_session, 1001, expected_version=1) + assert not await IdentitySubjectDao.increment_auth_version(data_session, 1001, expected_version=1) + loaded = await IdentitySubjectDao.get_by_subject_id(data_session, 'stable-subject') + assert loaded is not None + assert loaded.auth_version == _NEXT_AUTH_VERSION + + +@pytest.mark.asyncio +async def test_backfill_is_idempotent_and_reports_missing_users(data_session: AsyncSession) -> None: + """验证主体回填不重复创建并能发现未关联用户。""" + data_session.add_all( + [ + SysUser(user_id=1002, user_name='bob', nick_name='Bob', status='0', del_flag='0'), + SysUser(user_id=1003, user_name='carol', nick_name='Carol', status='0', del_flag='0'), + ] + ) + await data_session.flush() + assert await IdentitySubjectDao.list_missing_user_ids(data_session, [1002, 1003]) == [1002, 1003] + rows = await IdentitySubjectDao.backfill_for_users(data_session, [1002, 1003]) + assert len(rows) == _BACKFILL_COUNT + assert await IdentitySubjectDao.backfill_for_users(data_session, [1002, 1003]) == [] + assert await IdentitySubjectDao.list_missing_user_ids(data_session, [1002, 1003]) == [] diff --git a/ruoyi-fastapi-backend/tests/module_identity/models/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/models/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/models/test_identity_models.py b/ruoyi-fastapi-backend/tests/module_identity/models/test_identity_models.py new file mode 100644 index 000000000..674bb7d70 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/models/test_identity_models.py @@ -0,0 +1,123 @@ +import hashlib + +import pytest +from sqlalchemy import CheckConstraint +from sqlalchemy.ext.asyncio import AsyncSession + +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.do.oauth_client_do import SysOAuthClient, SysOAuthClientSecret, SysOAuthClientUri +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken, SysSsoSession +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.do.oidc_key_do import SysOidcSigningKey + +_HASH_HEX_LENGTH = 64 + + +def test_all_identity_tables_are_registered() -> None: + """断言统一认证第一阶段的全部表均注册到共享 MetaData。""" + expected = { + 'sys_identity_subject', + 'sys_oauth_client', + 'sys_oauth_client_secret', + 'sys_oauth_client_uri', + 'sys_oauth_resource', + 'sys_oauth_scope', + 'sys_oauth_client_scope', + 'sys_oauth_client_resource', + 'sys_oauth_grant', + 'sys_oauth_refresh_token', + 'sys_sso_session', + 'sys_oidc_signing_key', + 'sys_oauth_audit_log', + } + actual = { + model.__table__.name + for model in ( + SysIdentitySubject, + SysOAuthClient, + SysOAuthClientSecret, + SysOAuthClientUri, + SysOAuthResource, + SysOAuthScope, + SysOAuthClientScope, + SysOAuthClientResource, + SysOAuthGrant, + SysOAuthRefreshToken, + SysSsoSession, + SysOidcSigningKey, + SysOAuthAuditLog, + ) + } + assert actual == expected + + +def test_identity_constraints_and_indexes_are_security_relevant() -> None: + """断言主体、URI、Refresh Token 和密钥的唯一性/外键/索引存在。""" + subject = SysIdentitySubject.__table__ + assert {'uk_identity_subject_user', 'uk_identity_subject_subject'} <= { + constraint.name for constraint in subject.constraints if constraint.name + } + assert any(fk.name == 'fk_identity_subject_user' and fk.ondelete == 'RESTRICT' for fk in subject.foreign_keys) + + uri = SysOAuthClientUri.__table__ + assert 'uri_hash' in uri.c + assert 'uk_oauth_client_uri_hash' in {constraint.name for constraint in uri.constraints if constraint.name} + assert any(index.name == 'idx_oauth_client_uri_type' for index in uri.indexes) + + refresh = SysOAuthRefreshToken.__table__ + assert refresh.c.token_hash.type.length == _HASH_HEX_LENGTH + assert SysSsoSession.__table__.c.session_secret_hash.type.length == _HASH_HEX_LENGTH + assert SysSsoSession.__table__.c.user_agent_hash.type.length == _HASH_HEX_LENGTH + assert 'uk_oauth_refresh_token_hash' in {constraint.name for constraint in refresh.constraints if constraint.name} + assert any(index.name == 'idx_oauth_refresh_family' for index in refresh.indexes) + + key_constraints = [ + constraint for constraint in SysOidcSigningKey.__table__.constraints if isinstance(constraint, CheckConstraint) + ] + assert any(constraint.name == 'ck_oidc_signing_key_private_material' for constraint in key_constraints) + + +def test_client_update_by_is_not_datetime_updated_automatically() -> None: + """断言更新者字段保持字符串语义,更新时间由 update_time 管理。""" + assert SysOAuthClient.update_by.onupdate is None + assert SysOAuthClient.update_time.onupdate is not None + + +def test_client_scope_relationships_have_only_named_composite_primary_keys() -> None: + """断言关系表元数据与 migration 使用一致的复合主键。""" + scope_table = SysOAuthClientScope.__table__ + resource_table = SysOAuthClientResource.__table__ + assert scope_table.primary_key.name == 'pk_oauth_client_scope' + assert resource_table.primary_key.name == 'pk_oauth_client_resource' + assert 'pk_oauth_client_scope' not in {index.name for index in scope_table.indexes} + assert 'pk_oauth_client_resource' not in {index.name for index in resource_table.indexes} + + +@pytest.mark.asyncio +async def test_client_audit_defaults_are_empty_strings(data_session: AsyncSession) -> None: + """断言 Client 审计字段落库默认值是真正空字符串。""" + client = SysOAuthClient( + client_pk=9201, + client_id='default-client', + client_name='Default Client', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code'], + response_types=['code'], + ) + data_session.add(client) + await data_session.flush() + assert client.create_by == '' + assert client.update_by == '' + + +def test_client_uri_derives_hash_for_cross_database_unique_index() -> None: + """断言完整 URI 通过 SHA-256 派生固定长度摘要。""" + uri = SysOAuthClientUri(client_pk=1, uri_type='redirect', uri='https://client.example/callback') + assert uri.uri_hash == hashlib.sha256(uri.uri.encode('utf-8')).hexdigest() diff --git a/ruoyi-fastapi-backend/tests/module_identity/models/test_oidc_key_vo.py b/ruoyi-fastapi-backend/tests/module_identity/models/test_oidc_key_vo.py new file mode 100644 index 000000000..8611a8e0a --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/models/test_oidc_key_vo.py @@ -0,0 +1,36 @@ +from datetime import datetime, timezone + +import pytest +from pydantic import ValidationError + +from module_identity.entity.vo.oidc_key_vo import OidcKeyRotateModel, OidcKeyViewModel + + +def test_public_jwk_rejects_private_rsa_parameters() -> None: + with pytest.raises(ValidationError): + OidcKeyViewModel( + kid='key-1', + status='active', + publishAt='2026-01-01T00:00:00Z', + publicJwk={ + 'kty': 'RSA', + 'use': 'sig', + 'kid': 'key-1', + 'alg': 'RS256', + 'n': 'n', + 'e': 'AQAB', + 'd': 'private', + }, + ) + + +@pytest.mark.parametrize('value', ['2026-08-28 10:00:00', datetime(2026, 8, 28, 10, 0)]) +def test_rotation_request_rejects_missing_timezone(value: object) -> None: + with pytest.raises(ValidationError): + OidcKeyRotateModel(kid='key-1', publishAt=value) + + +def test_rotation_request_normalizes_offset_to_utc() -> None: + model = OidcKeyRotateModel(kid='key-1', publishAt='2026-08-28T10:00:00+08:00') + assert model.publish_at == datetime(2026, 8, 28, 2, 0, tzinfo=timezone.utc) + assert model.model_dump(mode='json', by_alias=True)['publishAt'] == '2026-08-28T02:00:00.000Z' diff --git a/ruoyi-fastapi-backend/tests/module_identity/models/test_protocol_vo.py b/ruoyi-fastapi-backend/tests/module_identity/models/test_protocol_vo.py new file mode 100644 index 000000000..de5ad6e61 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/models/test_protocol_vo.py @@ -0,0 +1,110 @@ +import pytest +from pydantic import ValidationError + +from module_identity.entity.vo.oauth_session_vo import AuditPageQueryModel +from module_identity.entity.vo.protocol_vo import ( + AuthorizeRequest, + ErrorResponse, + OAuthServerMetadata, + TokenRequest, + UserInfoResponse, +) +from module_identity.security.pkce import generate_code_challenge, generate_code_verifier + +_AUDIT_USER_ID = 1001 + + +def test_protocol_models_emit_snake_case_and_require_oidc_nonce() -> None: + challenge = generate_code_challenge(generate_code_verifier()) + with pytest.raises(ValidationError): + AuthorizeRequest( + response_type='code', + client_id='external', + redirect_uri='https://client.example/callback', + scope='openid profile', + code_challenge=challenge, + code_challenge_method='S256', + ) + model = AuthorizeRequest( + response_type='code', + client_id='external', + redirect_uri='https://client.example/callback', + scope='openid profile', + nonce='nonce', + code_challenge=challenge, + code_challenge_method='S256', + ) + assert 'client_id' in model.model_dump() + assert 'clientId' not in model.model_dump() + + +def test_protocol_enum_values_are_preserved_for_business_error_mapping() -> None: + """未知协议枚举保留给服务层生成 OAuth 标准错误,而非直接 422。""" + challenge = generate_code_challenge(generate_code_verifier()) + authorize = AuthorizeRequest( + response_type='unsupported', + client_id='external', + redirect_uri='https://client.example/callback', + scope='openid', + nonce='nonce', + code_challenge=challenge, + code_challenge_method='plain', + ) + assert authorize.response_type == 'unsupported' + assert authorize.code_challenge_method == 'plain' + token = TokenRequest(grant_type='urn:example:unknown') + assert token.grant_type == 'urn:example:unknown' + + +def test_token_request_code_requires_redirect_uri_and_userinfo_claims_are_allowed() -> None: + with pytest.raises(ValidationError): + TokenRequest(grant_type='authorization_code', code='ac1.code.secret', code_verifier='a' * 43) + with pytest.raises(ValidationError): + TokenRequest( + grant_type='authorization_code', + code='ac1.code.secret', + code_verifier='a' * 43, + redirect_uri='https://client.example/callback', + scope='openid', + ) + with pytest.raises(ValidationError): + TokenRequest(grant_type='refresh_token', refresh_token='rt1.token.secret', redirect_uri='https://example/cb') + with pytest.raises(ValidationError): + TokenRequest(grant_type='client_credentials', redirect_uri='https://example/cb') + info = UserInfoResponse(sub='subject', email='a@example.com', roles=['reader']) + assert info.email == 'a@example.com' + assert ErrorResponse(error='invalid_request').as_dict() == {'error': 'invalid_request'} + metadata = OAuthServerMetadata( + issuer='https://issuer.example', + authorization_endpoint='https://issuer.example/oauth2/authorize', + token_endpoint='https://issuer.example/oauth2/token', + jwks_uri='https://issuer.example/oauth2/jwks', + ) + assert 'userinfo_endpoint' not in metadata.model_dump() + + +def test_token_request_limits_untrusted_form_string_lengths() -> None: + """Token 表单中的 opaque 字符串必须有明确长度上限。""" + with pytest.raises(ValidationError): + TokenRequest( + grant_type='authorization_code', + code='c', + code_verifier='a' * 43, + redirect_uri='https://client.example/callback', + resource='r' * 1001, + ) + with pytest.raises(ValidationError): + TokenRequest(grant_type='refresh_token', refresh_token='r' * 4097) + + +def test_audit_query_supports_owner_and_time_filters() -> None: + query = AuditPageQueryModel( + clientId='external-client', + userId=_AUDIT_USER_ID, + startTime='2026-01-01T00:00:00Z', + endTime='2026-01-02T00:00:00Z', + ) + assert query.client_id == 'external-client' + assert query.user_id == _AUDIT_USER_ID + with pytest.raises(ValidationError): + AuditPageQueryModel(startTime='2026-01-02T00:00:00Z', endTime='2026-01-01T00:00:00Z') diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/security/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_client_auth.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_client_auth.py new file mode 100644 index 000000000..7deb08a40 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_client_auth.py @@ -0,0 +1,140 @@ +import base64 + +import pytest + +from module_identity.entity.vo.oauth_client_vo import ClientCreateModel, ClientUriModel +from module_identity.security.client_auth import ( + ClientAuthenticationError, + authenticate_client, + hash_client_secret, +) +from utils.oidc_util import OidcUtil + + +@pytest.mark.parametrize('scheme', ['Basic', 'basic', 'BASIC', 'bAsIc']) +def test_basic_credentials_decode_both_form_sides_and_colon(scheme: str) -> None: + raw = 'client+id:secret%2Bvalue%3A2' + header = scheme + ' ' + base64.b64encode(raw.encode()).decode() + + assert OidcUtil.parse_basic_credentials(header) == ('client id', 'secret+value:2') + + +def test_basic_requires_strict_standard_base64() -> None: + with pytest.raises(ValueError): + OidcUtil.parse_basic_credentials('Basic !!!not-base64!!!') + + +@pytest.mark.parametrize('prefix', ['', 'Bearer ', 'BasicX ', ' Basic ', 'Basic\t']) +def test_basic_rejects_other_schemes_and_missing_space(prefix: str) -> None: + header = prefix + base64.b64encode(b'client:SecretCase').decode() + with pytest.raises(ValueError): + OidcUtil.parse_basic_credentials(header) + + +@pytest.mark.parametrize('authorization', ['Basic !!!!', 'Bearer abc', 'Basic Y2xpZW50']) +def test_authentication_preserves_domain_error_for_malformed_basic(authorization: str) -> None: + """工具解析失败仍由认证入口转换为客户端认证失败。""" + client = {'client_id': 'client', 'client_type': 'confidential', 'status': '0'} + with pytest.raises(ClientAuthenticationError, match='客户端认证信息无效'): + authenticate_client(client, authorization) + + +def test_public_and_confidential_client_policy() -> None: + public = { + 'client_id': 'public-app', + 'client_type': 'public', + 'token_endpoint_auth_method': 'none', + 'status': '0', + } + assert authenticate_client(public, client_id='public-app').auth_method == 'none' + with pytest.raises(ClientAuthenticationError): + authenticate_client(public, client_id='public-app', client_secret='unexpected') + + secret = OidcUtil.generate_client_secret() + confidential = { + 'client_id': 'server-app', + 'client_type': 'confidential', + 'token_endpoint_auth_method': 'client_secret_basic', + 'secret_hash': hash_client_secret(secret), + 'status': '0', + } + header = 'Basic ' + base64.b64encode(f'server-app:{secret}'.encode()).decode() + assert authenticate_client(confidential, header).client_id == 'server-app' + with pytest.raises(ClientAuthenticationError): + authenticate_client(confidential, header, client_secret='body-secret') + with pytest.raises(ClientAuthenticationError): + authenticate_client(confidential, client_id='server-app', client_secret=secret) + for status in (None, '1', 'disabled'): + invalid = {**confidential, 'status': status} + with pytest.raises(ClientAuthenticationError): + authenticate_client(invalid, header) + + +def test_client_pkce_grant_and_uri_policy() -> None: + with pytest.raises(ValueError): + ClientCreateModel( + client_name='public', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code'], + require_pkce=False, + ) + with pytest.raises(ValueError): + ClientUriModel(uri_type='redirect', uri='https://client.example/callback#fragment') + with pytest.raises(ValueError): + ClientUriModel(uri_type='cors_origin', uri='https://client.example/path') + with pytest.raises(ValueError): + ClientCreateModel( + client_name='missing-callback', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + ) + with pytest.raises(ValueError): + ClientCreateModel( + client_name='refresh-only', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['refresh_token'], + response_types=[], + ) + + credentials_only = ClientCreateModel( + client_name='credentials-only', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['client_credentials'], + redirect_uris=[], + ) + assert credentials_only.grant_types == ['client_credentials'] + assert credentials_only.response_types == [] + with pytest.raises(ValueError): + ClientCreateModel( + client_name='public-credentials', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['client_credentials'], + response_types=[], + ) + with pytest.raises(ValueError): + ClientCreateModel( + client_name='wrong-code-response', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['authorization_code'], + response_types=[], + redirect_uris=['https://client.example/callback'], + ) + + valid = { + 'client_name': 'server', + 'client_type': 'confidential', + 'token_endpoint_auth_method': 'client_secret_basic', + 'redirect_uris': ['https://client.example/callback'], + 'post_logout_redirect_uris': ['https://client.example/logout'], + 'backchannel_logout_uris': ['https://client.example/backchannel'], + 'cors_origins': ['https://client.example:8443'], + } + assert ClientCreateModel(**valid).cors_origins == ['https://client.example:8443'] + for field in ('redirect_uris', 'post_logout_redirect_uris', 'backchannel_logout_uris', 'cors_origins'): + with pytest.raises(ValueError): + ClientCreateModel(**{**valid, field: ['https://user:password@client.example/callback']}) diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_identity_dependencies.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_identity_dependencies.py new file mode 100644 index 000000000..7d8d65b89 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_identity_dependencies.py @@ -0,0 +1,233 @@ +import time +from types import SimpleNamespace +from typing import Any + +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from fastapi import HTTPException +from starlette.requests import Request + +from exceptions.exception import OAuthProtocolException +from module_identity import dependencies +from module_identity.controller.token_controller import token_controller + +_BAD_REQUEST = 400 +_UNSUPPORTED_MEDIA = 415 +_UNAUTHORIZED = 401 +_PAYLOAD_TOO_LARGE = 413 +_SERVICE_UNAVAILABLE = 503 +_MAX_FORM_BYTES = 16 * 1024 + + +def _request( + body: bytes, + content_type: str = 'application/x-www-form-urlencoded', + declared_length: int | None = None, + include_content_length: bool = True, +) -> Request: + """构造最小 URL encoded 请求。""" + sent = False + + async def receive() -> dict[str, object]: + nonlocal sent + if sent: + return {'type': 'http.disconnect'} + sent = True + return {'type': 'http.request', 'body': body, 'more_body': False} + + headers = [(b'content-type', content_type.encode())] + if include_content_length: + headers.append((b'content-length', str(len(body) if declared_length is None else declared_length).encode())) + return Request( + { + 'type': 'http', + 'method': 'POST', + 'path': '/oauth2/token', + 'headers': headers, + }, + receive, + ) + + +@pytest.mark.asyncio +async def test_read_form_rejects_duplicate_and_wrong_content_type(monkeypatch: pytest.MonkeyPatch) -> None: + """重复字段与非表单请求必须 fail closed。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + with pytest.raises(HTTPException) as duplicate: + await dependencies.read_form(_request(b'client_id=a&client_id=b')) + assert duplicate.value.status_code == _BAD_REQUEST + with pytest.raises(HTTPException) as media: + await dependencies.read_form(_request(b'client_id=a', 'application/json')) + assert media.value.status_code == _UNSUPPORTED_MEDIA + + +@pytest.mark.asyncio +async def test_read_form_rejects_oversized_and_misdeclared_body(monkeypatch: pytest.MonkeyPatch) -> None: + """协议表单在读取前后均限制 16KiB,不能用伪造 Content-Length 绕过。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + oversized = _request(b'a' * (_MAX_FORM_BYTES + 1)) + with pytest.raises(HTTPException) as too_large: + await dependencies.read_form(oversized) + assert too_large.value.status_code == _PAYLOAD_TOO_LARGE + mismatched = _request(b'client_id=a', declared_length=1) + with pytest.raises(HTTPException) as mismatch: + await dependencies.read_form(mismatched) + assert mismatch.value.status_code == _PAYLOAD_TOO_LARGE + + +@pytest.mark.asyncio +async def test_read_form_accepts_chunked_body_without_content_length(monkeypatch: pytest.MonkeyPatch) -> None: + """缺少 Content-Length 的合法 chunked 表单按实际流大小读取。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + request = _request(b'client_id=chunked', include_content_length=False) + assert await dependencies.read_form(request) == {'client_id': 'chunked'} + + +@pytest.mark.asyncio +async def test_read_form_rejects_malformed_percent_escape(monkeypatch: pytest.MonkeyPatch) -> None: + """表单百分号转义必须完整且为十六进制。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + with pytest.raises(HTTPException) as raised: + await dependencies.read_form(_request(b'client_id=%ZZ')) + assert raised.value.status_code == _BAD_REQUEST + + +@pytest.mark.asyncio +async def test_access_key_helper_rejects_remote_key_header(monkeypatch: pytest.MonkeyPatch) -> None: + """kid 验证只允许本地公钥,拒绝 jku/x5u/jwk/x5c。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + token = jwt.encode( + {'iss': 'https://issuer.example', 'sub': 'subject'}, + key, + algorithm='RS256', + headers={'kid': 'kid-1', 'typ': 'at+jwt', 'jku': 'https://evil.example/jwks'}, + ) + with pytest.raises(dependencies.JwtProfileError): + await dependencies.load_access_verification_key(token, object()) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('scheme', ['Bearer', 'bearer', 'BEARER', 'bEaReR']) +async def test_userinfo_accepts_case_insensitive_scheme_without_changing_token( + monkeypatch: pytest.MonkeyPatch, scheme: str +) -> None: + """方案名大小写不影响真实签名令牌,凭据原文必须保持不变。""" + issuer = 'https://issuer.example' + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_issuer', issuer) + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + now = int(time.time()) + claims = { + 'iss': issuer, + 'sub': 'subject-1', + 'aud': f'{issuer}/oauth2/userinfo', + 'iat': now, + 'nbf': now, + 'exp': now + 600, + 'jti': 'jti-1', + 'client_id': 'client-1', + 'sid': 'session-1', + 'scope': 'openid', + 'ver': 1, + 'gty': 'authorization_code', + 'grant_id': 'grant-1', + 'client_policy_version': 1, + 'auth_time': now, + 'acr': 'pwd', + 'amr': ['pwd'], + } + token = jwt.encode(claims, key, algorithm='RS256', headers={'kid': 'kid-1', 'typ': 'at+jwt'}) + monkeypatch.setattr(dependencies, 'load_access_verification_key', lambda *args: _async(key.public_key())) + request = Request( + { + 'type': 'http', + 'method': 'GET', + 'path': '/oauth2/userinfo', + 'headers': [(b'authorization', f'{scheme} {token}'.encode())], + } + ) + context = await dependencies.get_oidc_access_token(request, object()) + assert context.token == token + assert context.claims['sub'] == 'subject-1' + + +@pytest.mark.asyncio +@pytest.mark.parametrize('prefix', ['', 'Basic ', 'BearerX ', ' Bearer ', 'Bearer\t']) +async def test_userinfo_rejects_other_schemes_and_missing_space(monkeypatch: pytest.MonkeyPatch, prefix: str) -> None: + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + request = Request( + { + 'type': 'http', + 'method': 'GET', + 'path': '/oauth2/userinfo', + 'headers': [(b'authorization', f'{prefix}header.payload.signature'.encode())], + } + ) + with pytest.raises(HTTPException) as raised: + await dependencies.get_oidc_access_token(request, object()) + assert raised.value.status_code == _UNAUTHORIZED + + +@pytest.mark.asyncio +async def test_client_dependency_invalid_client_has_basic_challenge(monkeypatch: pytest.MonkeyPatch) -> None: + """Client Authentication 依赖的 invalid_client 必须带 Basic challenge。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(dependencies, 'read_form', lambda request: _async({'client_id': 'client-1'})) + + async def invalid(*args: object, **kwargs: object) -> object: + raise OAuthProtocolException('invalid_client', 'Client authentication failed', 401) + + monkeypatch.setattr(dependencies.TokenService, 'authenticate_client', invalid) + request = SimpleNamespace(headers={}) + with pytest.raises(HTTPException) as raised: + await dependencies.get_oidc_client(request, object()) + assert raised.value.status_code == _UNAUTHORIZED + assert raised.value.headers == {'WWW-Authenticate': 'Basic realm="oauth2/token"'} + + +def test_protocol_routers_have_no_legacy_preauth_dependency() -> None: + """协议路由不得自动挂载 Legacy PreAuth。""" + assert all( + all( + getattr(dependency.dependency, '__name__', '') == 'require_oidc_protocol_ready' + for dependency in route.dependencies + ) + for route in token_controller.routes + ) + + +@pytest.mark.asyncio +async def test_protocol_readiness_dependency_returns_503_without_active_key( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """协议已开启但签名密钥未就绪时返回稳定 OAuth 503。""" + monkeypatch.setattr(dependencies.OidcConfig, 'oidc_enabled', True) + not_ready = SimpleNamespace(ready=False) + monkeypatch.setattr( + dependencies.OidcRuntimeService, + 'cached_readiness', + lambda *args: _async(not_ready), + ) + request = SimpleNamespace(app=SimpleNamespace()) + with pytest.raises(OAuthProtocolException) as raised: + await dependencies.require_oidc_protocol_ready(request, object()) + assert raised.value.error == 'temporarily_unavailable' + assert raised.value.status_code == _SERVICE_UNAVAILABLE + + +def test_invalid_token_challenge_is_standard_bearer() -> None: + """Access Token 依赖的失败响应不泄漏验签细节。""" + error = dependencies._invalid_token() + assert error.status_code == _UNAUTHORIZED + assert error.headers['WWW-Authenticate'] == 'Bearer error="invalid_token"' + + +def _async(value: object) -> Any: + """构造异步测试结果。""" + + async def result() -> object: + return value + + return result() diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_jwt_profile.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_jwt_profile.py new file mode 100644 index 000000000..55fc6de4f --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_jwt_profile.py @@ -0,0 +1,274 @@ +import math +from datetime import datetime, timezone + +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa + +from module_identity.security.jwt_profile import ( + BACKCHANNEL_LOGOUT_EVENT, + JwtProfileError, + decode_access_token, + decode_id_token, + decode_logout_token, + encode_access_token, + encode_id_token, + encode_logout_token, +) + + +@pytest.fixture() +def key_pair() -> tuple[object, object]: + private = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return private, private.public_key() + + +def _now() -> int: + return int(datetime.now(timezone.utc).timestamp()) + + +def test_access_profile_requires_at_jwt_and_checks_audience(key_pair: tuple[object, object]) -> None: + private, public = key_pair + now = _now() + claims = { + 'iss': 'https://issuer.example', + 'sub': 'subject', + 'aud': ['resource-a'], + 'exp': now + 60, + 'iat': now, + 'nbf': now, + 'jti': 'token-id', + 'client_id': 'client', + 'scope': 'openid', + 'gty': 'authorization_code', + 'sid': 'session-id', + 'ver': 1, + 'auth_time': now, + 'acr': 'pwd', + 'amr': ['pwd'], + } + token = encode_access_token(claims, signing_key=private, kid='key-1') + assert ( + decode_access_token(token, verification_key=public, issuer=claims['iss'], audience='resource-a')['sub'] + == 'subject' + ) + with pytest.raises(JwtProfileError): + decode_access_token(token, verification_key=public, issuer=claims['iss'], audience='resource-b') + with pytest.raises(JwtProfileError): + decode_access_token(token, verification_key=public, issuer='https://other.example', audience='resource-a') + + +def test_id_token_cannot_be_used_as_access_token_and_leeway_is_forwarded(key_pair: tuple[object, object]) -> None: + private, public = key_pair + now = _now() + claims = { + 'iss': 'https://issuer.example', + 'sub': 'subject', + 'aud': 'client', + 'exp': now - 61, + 'iat': now - 60, + 'auth_time': now - 60, + 'nonce': 'opaque-nonce', + 'sid': 'sid', + 'acr': 'pwd', + 'amr': ['pwd'], + } + id_token = encode_id_token(claims, signing_key=private, kid='key-1') + with pytest.raises(JwtProfileError): + decode_access_token(id_token, verification_key=public, issuer=claims['iss'], audience='client') + assert ( + decode_id_token(id_token, verification_key=public, issuer=claims['iss'], audience='client', clock_skew=120)[ + 'sub' + ] + == 'subject' + ) + with pytest.raises(JwtProfileError): + encode_id_token({key: value for key, value in claims.items() if key != 'nonce'}, private, 'key-1') + + +def test_wrong_algorithm_kid_and_numeric_types_are_rejected(key_pair: tuple[object, object]) -> None: + private, public = key_pair + now = _now() + claims = { + 'iss': 'https://issuer.example', + 'sub': 's', + 'aud': 'r', + 'exp': now + 60, + 'iat': now, + 'nbf': now, + 'jti': 'j', + 'client_id': 'c', + 'scope': 'openid', + 'gty': 'authorization_code', + 'sid': 'session-id', + 'ver': 1, + 'auth_time': now, + 'acr': 'pwd', + 'amr': ['pwd'], + } + with pytest.raises(JwtProfileError): + encode_access_token({**claims, 'exp': True}, signing_key=private, kid='key-1') + with pytest.raises(JwtProfileError): + encode_access_token({**claims, 'exp': math.nan}, signing_key=private, kid='key-1') + with pytest.raises(JwtProfileError): + encode_access_token({**claims, 'aud': []}, signing_key=private, kid='key-1') + with pytest.raises(JwtProfileError): + encode_access_token({**claims, 'aud': ['r', 'r']}, signing_key=private, kid='key-1') + token = encode_access_token(claims, signing_key=private, kid='key-1') + header = jwt.get_unverified_header(token) + assert header['typ'] == 'at+jwt' and header['alg'] == 'RS256' + with pytest.raises(JwtProfileError): + decode_access_token(token, issuer=claims['iss'], audience='r', verification_keys={'other': public}) + unsafe_header = jwt.encode( + claims, + private, + algorithm='RS256', + headers={'kid': 'key-1', 'typ': 'at+jwt', 'jku': 'https://attacker.example/jwks'}, + ) + with pytest.raises(JwtProfileError): + decode_access_token(unsafe_header, verification_key=public, issuer=claims['iss'], audience='r') + + +def test_signed_non_finite_numeric_date_is_rejected(key_pair: tuple[object, object]) -> None: + private, public = key_pair + now = _now() + claims = { + 'iss': 'https://issuer.example', + 'sub': 's', + 'aud': 'r', + 'exp': math.nan, + 'iat': now, + 'jti': 'j', + 'client_id': 'c', + 'scope': 'openid', + } + token = jwt.encode(claims, private, algorithm='RS256', headers={'kid': 'key-1', 'typ': 'at+jwt'}) + with pytest.raises(JwtProfileError): + decode_access_token(token, verification_key=public, issuer=claims['iss'], audience='r') + + +def test_access_profile_requires_gty_and_rejects_machine_user_field_confusion( + key_pair: tuple[object, object], +) -> None: + """Access Profile 强制 gty,并隔离机器与用户绑定字段。""" + private, public = key_pair + now = _now() + base = { + 'iss': 'https://issuer.example', + 'aud': 'r', + 'exp': now + 60, + 'iat': now, + 'nbf': now, + 'jti': 'j', + 'client_id': 'c', + 'scope': 'scope', + } + with pytest.raises(JwtProfileError): + encode_access_token({**base, 'sub': 'client:c', 'gty': 'client_credentials', 'sid': 'sid'}, private, 'key-1') + machine = encode_access_token({**base, 'sub': 'client:c', 'gty': 'client_credentials'}, private, 'key-1') + assert decode_access_token(machine, verification_key=public, issuer=base['iss'], audience='r')['gty'] == ( + 'client_credentials' + ) + user = encode_access_token( + { + **base, + 'sub': 'subject', + 'gty': 'authorization_code', + 'sid': 'sid', + 'ver': 1, + 'auth_time': now, + 'acr': 'pwd', + 'amr': ['pwd'], + }, + private, + 'key-1', + ) + assert decode_access_token(user, verification_key=public, issuer=base['iss'], audience='r')['gty'] == ( + 'authorization_code' + ) + + +def test_logout_profile_requires_event_and_non_empty_sid_or_sub(key_pair: tuple[object, object]) -> None: + private, public = key_pair + now = _now() + claims = { + 'iss': 'https://issuer.example', + 'aud': 'client', + 'iat': now, + 'exp': now + 120, + 'jti': 'logout-id', + 'sid': 'sid', + 'events': {BACKCHANNEL_LOGOUT_EVENT: {}}, + } + token = encode_logout_token(claims, signing_key=private, kid='key-1') + assert decode_logout_token(token, verification_key=public, issuer=claims['iss'], audience='client')['sid'] == 'sid' + forged_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + forged_token = encode_logout_token(claims, signing_key=forged_key, kid='key-1') + with pytest.raises(JwtProfileError): + decode_logout_token(forged_token, verification_key=public, issuer=claims['iss'], audience='client') + with pytest.raises(JwtProfileError): + encode_logout_token({**claims, 'sid': ''}, signing_key=private, kid='key-1') + with pytest.raises(JwtProfileError): + encode_logout_token({**claims, 'nonce': 'must-not-exist'}, signing_key=private, kid='key-1') + with pytest.raises(JwtProfileError): + encode_logout_token( + {key: value for key, value in claims.items() if key != 'iat'}, signing_key=private, kid='key-1' + ) + with pytest.raises(JwtProfileError): + encode_logout_token({**claims, 'aud': []}, signing_key=private, kid='key-1') + + +@pytest.mark.parametrize( + 'changes', + [ + {'iss': 'https://other.example'}, + {'aud': 'other'}, + {'nonce': 'not-allowed'}, + {'events': {}}, + {'events': {BACKCHANNEL_LOGOUT_EVENT: []}}, + {'sid': ''}, + {'jti': ''}, + {'iat': _now() + 3600}, + ], +) +def test_logout_profile_rejects_invalid_signed_claims( + key_pair: tuple[object, object], changes: dict[str, object] +) -> None: + """即使 RSA 签名有效,也必须拒绝错误的退出事件、身份绑定和时效声明。""" + private, public = key_pair + now = _now() + claims = { + 'iss': 'https://issuer.example', + 'aud': 'client', + 'iat': now, + 'exp': now + 120, + 'jti': 'logout-id', + 'sid': 'sid', + 'events': {BACKCHANNEL_LOGOUT_EVENT: {}}, + } + token = jwt.encode({**claims, **changes}, private, algorithm='RS256', headers={'kid': 'key-1', 'typ': 'logout+jwt'}) + with pytest.raises(JwtProfileError): + decode_logout_token(token, verification_key=public, issuer=claims['iss'], audience='client', clock_skew=0) + + +@pytest.mark.parametrize('expiry', [None, 'future', True, 0, -1]) +def test_logout_profile_rejects_missing_invalid_or_expired_exp(key_pair: tuple[object, object], expiry: object) -> None: + """接收方拒绝缺失、类型错误或过期的退出令牌。""" + + private, public = key_pair + claims = { + 'iss': 'https://issuer.example', + 'aud': 'client', + 'iat': _now() - 1, + 'jti': 'logout-expiry', + 'sid': 'sid', + 'events': {BACKCHANNEL_LOGOUT_EVENT: {}}, + } + if expiry is not None: + claims['exp'] = expiry + token = jwt.encode(claims, private, algorithm='RS256', headers={'kid': 'key-1', 'typ': 'logout+jwt'}) + with pytest.raises(JwtProfileError): + decode_logout_token(token, verification_key=public, issuer=claims['iss'], audience='client', clock_skew=0) + if expiry is None: + with pytest.raises(JwtProfileError): + encode_logout_token(claims, private, 'key-1') diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_oauth_rate_limit_metadata.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_oauth_rate_limit_metadata.py new file mode 100644 index 000000000..996a5a366 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_oauth_rate_limit_metadata.py @@ -0,0 +1,67 @@ +from collections.abc import Callable +from typing import Any + +from fastapi.routing import APIRoute + +from common.annotation.rate_limit_annotation import ApiRateLimit +from common.constant import ApiNamespace +from module_identity.controller.oauth_audit_controller import oauth_audit_controller +from module_identity.controller.oauth_client_controller import oauth_client_controller +from module_identity.controller.oauth_resource_controller import oauth_resource_controller, oauth_scope_controller +from module_identity.controller.oauth_session_controller import oauth_grant_controller, oauth_session_controller +from module_identity.controller.oidc_key_controller import oidc_key_controller + + +def _route(router: Any, path: str, method: str) -> APIRoute: + """查找指定方法的路由。""" + for route in router.routes: + if isinstance(route, APIRoute) and route.path == path and method in route.methods: + return route + raise AssertionError(f'route not found: {method} {path}') + + +def _rate_limit(endpoint: Callable[..., Any]) -> ApiRateLimit: + """从 ``functools.wraps`` 生成的限流包装器闭包读取配置。""" + pending: list[Callable[..., Any]] = [endpoint] + visited: set[int] = set() + while pending: + current = pending.pop() + if id(current) in visited: + continue + visited.add(id(current)) + for cell in getattr(current, '__closure__', None) or (): + value = cell.cell_contents + if isinstance(value, ApiRateLimit): + return value + if callable(value) and hasattr(value, '__closure__'): + pending.append(value) + raise AssertionError(f'route is missing ApiRateLimit: {endpoint!r}') + + +def test_high_risk_identity_routes_have_independent_rate_limit_namespaces() -> None: + """高风险写操作必须使用限流预设且不能共享命名空间。""" + expected = [ + (oauth_client_controller, '/system/oauth/client', 'POST', ApiNamespace.SYSTEM_OAUTH_CLIENT_CREATE), + (oauth_client_controller, '/system/oauth/client', 'PUT', ApiNamespace.SYSTEM_OAUTH_CLIENT_UPDATE), + ( + oauth_client_controller, + '/system/oauth/client/{client_id}/secret', + 'POST', + ApiNamespace.SYSTEM_OAUTH_CLIENT_SECRET_ROTATE, + ), + (oauth_session_controller, '/system/oauth/session/{sids}', 'DELETE', ApiNamespace.SYSTEM_OAUTH_SESSION_REVOKE), + (oauth_grant_controller, '/system/oauth/grant/{grant_ids}', 'DELETE', ApiNamespace.SYSTEM_OAUTH_GRANT_REVOKE), + (oidc_key_controller, '/system/oauth/key/rotate', 'POST', ApiNamespace.SYSTEM_OAUTH_KEY_ROTATE), + (oidc_key_controller, '/system/oauth/key/{kid}/retire', 'PUT', ApiNamespace.SYSTEM_OAUTH_KEY_RETIRE), + (oauth_audit_controller, '/monitor/oauth/audit/export', 'POST', ApiNamespace.MONITOR_OAUTH_AUDIT_EXPORT), + (oauth_resource_controller, '/system/oauth/resource', 'PUT', ApiNamespace.SYSTEM_OAUTH_RESOURCE_UPDATE), + (oauth_scope_controller, '/system/oauth/scope', 'PUT', ApiNamespace.SYSTEM_OAUTH_SCOPE_UPDATE), + ] + limits = [_rate_limit(_route(router, path, method).endpoint) for router, path, method, _ in expected] + assert [limit.namespace for limit in limits] == [namespace for _, _, _, namespace in expected] + assert len({limit.namespace for limit in limits}) == len(limits) + assert all( + limit.preset_name + in {'USER_COMMON_MUTATION', 'USER_SECURITY_MUTATION', 'USER_DESTRUCTIVE_MUTATION', 'USER_RESOURCE_EXPORT'} + for limit in limits + ) diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_oidc_messages.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_oidc_messages.py new file mode 100644 index 000000000..432a114de --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_oidc_messages.py @@ -0,0 +1,99 @@ +from urllib.parse import parse_qs, urlsplit + +import pytest +from fastapi import FastAPI, status +from fastapi.testclient import TestClient + +from exceptions.exception import OAuthProtocolException, OidcInteractionException +from exceptions.handle import handle_exception +from module_identity.security.backchannel_transport import PermanentBackchannelError +from module_identity.service.token_protocol_service import RevocationError + + +@pytest.mark.parametrize( + ('error', 'description', 'message'), + [ + ('invalid_client', 'Client authentication failed', '客户端认证失败'), + ('invalid_scope', 'Requested scope is not authorized', '请求的权限范围尚未获准'), + ('server_error', 'Token endpoint is unavailable', '令牌服务暂不可用'), + ], +) +def test_protocol_exception_has_chinese_diagnostics_and_compatible_wire_description( + error: str, description: str, message: str +) -> None: + exception = OAuthProtocolException(error, description) + + assert str(exception) == exception.message == message + assert exception.as_dict() == {'error': error, 'error_description': description} + + +def test_code_only_exception_keeps_optional_protocol_description_absent() -> None: + exception = OAuthProtocolException('invalid_scope') + + assert str(exception) == '请求的权限范围无效' + assert exception.as_dict() == {'error': 'invalid_scope'} + + +def test_chinese_interaction_message_keeps_redirect_protocol_ascii_and_state_unchanged() -> None: + app = FastAPI() + handle_exception(app) + interaction = OidcInteractionException( + 'interaction-id', + '权限范围策略已变更,请重新授权', + error='invalid_scope', + redirect_uri='https://client.example/callback?tenant=one', + state='opaque-state', + redirect_uri_verified=True, + ) + + @app.get('/authorize') + async def authorize() -> None: + raise interaction.as_protocol_exception() + + with TestClient(app) as client: + response = client.get('/authorize', follow_redirects=False) + + assert response.status_code == status.HTTP_303_SEE_OTHER + assert str(interaction) == '权限范围策略已变更,请重新授权' + location = response.headers['location'] + assert location.isascii() + assert parse_qs(urlsplit(location).query) == { + 'tenant': ['one'], + 'error': ['invalid_scope'], + 'error_description': ['Scope policy has changed'], + 'state': ['opaque-state'], + } + assert response.headers['cache-control'] == 'no-store' + + +@pytest.mark.parametrize('description', ['内部诊断:private-marker', 'dependency\nprivate-marker']) +def test_unknown_diagnostic_is_not_exposed_as_protocol_description(description: str) -> None: + app = FastAPI() + handle_exception(app) + + @app.get('/token') + async def token() -> None: + raise OAuthProtocolException('server_error', description, status_code=503) + + with TestClient(app) as client: + response = client.get('/token') + + assert response.status_code == status.HTTP_503_SERVICE_UNAVAILABLE + assert response.json() == {'error': 'server_error', 'error_description': 'Authorization service is unavailable'} + assert 'private-marker' not in response.text + assert response.headers['cache-control'] == 'no-store' + + +def test_revocation_exception_localizes_diagnostics_without_changing_protocol_fields() -> None: + exception = RevocationError('invalid_client', 'Client authentication failed') + + assert str(exception) == exception.message == '客户端认证失败' + assert exception.error == 'invalid_client' + assert exception.description == 'Client authentication failed' + + +def test_backchannel_exception_keeps_audit_code_independent_from_chinese_diagnostics() -> None: + exception = PermanentBackchannelError('invalid retry payload', '后端退出通知重试载荷无效') + + assert str(exception) == exception.message == '后端退出通知重试载荷无效' + assert exception.failure_code == 'invalid retry payload' diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_opaque_token.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_opaque_token.py new file mode 100644 index 000000000..10aba6124 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_opaque_token.py @@ -0,0 +1,42 @@ +from collections.abc import Callable + +import pytest + +from module_identity.security.opaque_token import ( + OpaqueTokenError, + generate_authorization_code, + generate_refresh_token, + generate_sso_cookie, + parse_opaque_token, + token_digest, + verify_token_digest, +) + + +@pytest.mark.parametrize( + ('factory', 'prefix'), + ((generate_authorization_code, 'ac1'), (generate_refresh_token, 'rt1'), (generate_sso_cookie, 'ss1')), +) +def test_typed_tokens_are_random_and_verifiable(factory: Callable[[], str], prefix: str) -> None: + pepper = 'p' * 32 + token = factory() + parsed = parse_opaque_token(token, expected_prefix=prefix) + digest = token_digest(token, pepper) + + assert parsed.prefix == prefix + assert verify_token_digest(token, digest, pepper) + assert not verify_token_digest(token.replace(prefix, 'rt1' if prefix != 'rt1' else 'ac1', 1), digest, pepper) + + +def test_token_id_alone_is_not_a_credential() -> None: + token = generate_refresh_token() + parsed = parse_opaque_token(token) + pepper = 'p' * 32 + digest = token_digest(token, pepper) + + assert not verify_token_digest(parsed.token_id, digest, pepper) + with pytest.raises(OpaqueTokenError): + parse_opaque_token('unknown.' + token.split('.', 1)[1]) + assert not verify_token_digest(token, digest, 'p' * 16) + with pytest.raises(ValueError): + token_digest(token, 'p' * 16) diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_pkce.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_pkce.py new file mode 100644 index 000000000..98aa429a0 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_pkce.py @@ -0,0 +1,31 @@ +import pytest + +from module_identity.security.pkce import ( + PkceError, + generate_code_challenge, + generate_code_verifier, + verify_code_challenge, +) + +_VERIFIER_LENGTH = 64 +_CHALLENGE_LENGTH = 43 + + +def test_s256_round_trip_and_mismatch() -> None: + verifier = generate_code_verifier() + challenge = generate_code_challenge(verifier) + + assert len(verifier) == _VERIFIER_LENGTH + assert len(challenge) == _CHALLENGE_LENGTH + assert verify_code_challenge(verifier, challenge) + assert not verify_code_challenge(f'{verifier}a', challenge) + + +def test_plain_and_invalid_lengths_are_rejected() -> None: + with pytest.raises(PkceError): + verify_code_challenge('a' * 43, 'a' * 43, method='plain') + with pytest.raises(ValueError): + generate_code_verifier(42) + with pytest.raises(PkceError): + generate_code_challenge('a' * 42) + assert not verify_code_challenge('a' * 43, '!' * 43) diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_protocol_exception_handler.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_protocol_exception_handler.py new file mode 100644 index 000000000..e7a37005c --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_protocol_exception_handler.py @@ -0,0 +1,88 @@ +from fastapi import FastAPI, status +from fastapi.testclient import TestClient + +from exceptions.exception import OAuthProtocolException, OidcInteractionException +from exceptions.handle import handle_exception + + +def _client() -> TestClient: + app = FastAPI() + handle_exception(app) + + @app.get('/invalid-client') + async def invalid_client() -> None: + raise OAuthProtocolException( + 'invalid_client', + 'Client authentication failed', + status.HTTP_401_UNAUTHORIZED, + headers={'Cache-Control': 'public, max-age=3600', 'Pragma': 'cache'}, + ) + + @app.get('/authorize-error') + async def authorize_error() -> None: + raise OAuthProtocolException( + 'access_denied', + redirect_uri=( + 'https://client.example/callback?tenant=one&tenant=two&error=old&error_description=old' + '&error_uri=https%3A%2F%2Fevil.example&code=old&state=old&iss=https%3A%2F%2Fevil.example' + ), + state='opaque-state', + redirect_uri_verified=True, + issuer='https://auth.example.com', + ) + + @app.get('/invalid-redirect') + async def invalid_redirect() -> None: + raise OAuthProtocolException( + 'access_denied', + redirect_uri='https://client.example/callback#fragment', + redirect_uri_verified=True, + ) + + @app.get('/interaction') + async def interaction() -> None: + raise OidcInteractionException('interaction-1', 'Interaction expired') + + return TestClient(app) + + +def test_invalid_client_uses_standard_oauth_json() -> None: + response = _client().get('/invalid-client') + + assert response.status_code == status.HTTP_401_UNAUTHORIZED + assert response.json()['error'] == 'invalid_client' + assert response.headers['cache-control'] == 'no-store' + assert response.headers['pragma'] == 'no-cache' + assert response.headers['www-authenticate'] == 'Basic realm="oauth2/token"' + assert 'success' not in response.json() + + +def test_verified_redirect_replaces_protocol_parameters_and_preserves_application_query() -> None: + response = _client().get('/authorize-error', follow_redirects=False) + + assert response.status_code == status.HTTP_303_SEE_OTHER + assert response.headers['location'] == ( + 'https://client.example/callback?tenant=one&tenant=two&error=access_denied&state=opaque-state' + '&iss=https%3A%2F%2Fauth.example.com' + ) + assert response.headers['cache-control'] == 'no-store' + assert response.headers['pragma'] == 'no-cache' + + +def test_invalid_validated_redirect_fails_closed_locally() -> None: + response = _client().get('/invalid-redirect', follow_redirects=False) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert 'location' not in response.headers + assert response.json()['error'] == 'server_error' + + +def test_interaction_exception_keeps_business_response_boundary() -> None: + response = _client().get('/interaction') + + assert response.status_code == status.HTTP_200_OK + assert response.json()['code'] is not None + assert response.json()['data'] == 'interaction-1' + assert 'error' not in response.json() + assert response.headers['cache-control'] == 'no-store' + assert response.headers['pragma'] == 'no-cache' diff --git a/ruoyi-fastapi-backend/tests/module_identity/security/test_redis_keys.py b/ruoyi-fastapi-backend/tests/module_identity/security/test_redis_keys.py new file mode 100644 index 000000000..29b62aa34 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/security/test_redis_keys.py @@ -0,0 +1,45 @@ +import pytest + +from module_identity.redis_keys import OidcRedisKey +from utils.oidc_util import OidcUtil + +_SHA256_HEX_LENGTH = 64 + + +def test_oidc_keys_are_isolated_from_legacy_access_token_namespace() -> None: + keys = [ + OidcRedisKey.interaction('interaction-id'), + OidcRedisKey.authorization_code('code-id'), + OidcRedisKey.sso_session('session-id'), + OidcRedisKey.user_sessions(7), + OidcRedisKey.revoked_jti('token-id'), + ] + + assert all(key.startswith('oidc:') and not key.startswith('access_token:') for key in keys) + + +def test_sensitive_identifiers_use_independent_32_byte_pepper() -> None: + pepper = 'independent-pepper-value-' + 'x' * 8 + digest = OidcUtil.hash_sensitive_identifier('user@example.com', pepper) + + assert len(digest) == _SHA256_HEX_LENGTH + assert digest in OidcRedisKey.login_user_rate_limit(digest) + assert digest in OidcRedisKey.sso_cookie(digest) + + with pytest.raises(ValueError): + OidcUtil.hash_sensitive_identifier('user@example.com', 'short') + with pytest.raises(TypeError): + OidcUtil.hash_sensitive_identifier('user@example.com', 123) # type: ignore[arg-type] + + +@pytest.mark.parametrize('value', ['../escape', 'space value', 'a' * 129]) +def test_unsafe_key_components_are_rejected(value: str) -> None: + with pytest.raises(ValueError): + OidcRedisKey.interaction(value) + + +def test_digest_components_must_be_sha256_hex() -> None: + with pytest.raises(ValueError): + OidcRedisKey.sso_cookie('not-a-digest') + with pytest.raises(TypeError): + OidcRedisKey.sso_cookie(123) # type: ignore[arg-type] diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/services/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_audit_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_audit_service.py new file mode 100644 index 000000000..70ca9ee99 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_audit_service.py @@ -0,0 +1,192 @@ +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace + +import pytest +import pytest_asyncio +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from config.database import Base +from exceptions.exception import OidcInteractionException +from module_identity.dao.oauth_audit_dao import OAuthAuditDao +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditArchive, SysOAuthAuditLog +from module_identity.entity.vo.oauth_session_vo import AuditPageQueryModel +from module_identity.service.audit_service import AuditService +from utils.oidc_util import OidcUtil + +_ADMIN_EVENT_ID = 7 +_ADMIN_PAGE_SIZE = 5 +_ADMIN_TOTAL = 1 +_EXPORT_LIMIT = 5000 +_ROLLBACK_COUNT = 2 +_SERVICE_UNAVAILABLE = 503 + + +@pytest_asyncio.fixture +async def audit_session() -> AsyncSession: + """创建只包含审计表的真实 SQLite 会话。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + async with engine.begin() as connection: + await connection.run_sync( + lambda sync_connection: Base.metadata.create_all( + sync_connection, + tables=[SysOAuthAuditLog.__table__, SysOAuthAuditArchive.__table__], + ) + ) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as session: + yield session + await engine.dispose() + + +def test_audit_redacts_nested_protocol_secrets_and_keeps_safe_fields() -> None: + """嵌套字典和列表中的协议秘密不得落入审计详情。""" + result = OidcUtil.sanitize_audit_detail( + { + 'attempt': 2, + 'client': {'name': 'web', 'client_secret': 'hidden'}, + 'nested': [ + {'code': 'authorization-code', 'label': 'safe'}, + {'opaque': 'rt1.record-id.secret-value'}, + {'opaque': 'ac1.record-id.secret-value'}, + {'opaque': 'prefix ss1.record-id.secret-value suffix'}, + {'opaque': 'prefix cs1.secret-value-with-enough-entropy suffix'}, + {'pkce_verifier': 'hidden'}, + ], + 'password': 'hidden', + } + ) + assert result == { + 'attempt': 2, + 'client': {'name': 'web'}, + 'nested': [ + {'label': 'safe'}, + {'opaque': '[REDACTED]'}, + {'opaque': '[REDACTED]'}, + {'opaque': '[REDACTED]'}, + {'opaque': '[REDACTED]'}, + {}, + ], + } + + +def test_audit_redacts_opaque_user_agent_but_keeps_plain_identifiers() -> None: + """凭据外观的 User-Agent 脱敏,普通 Client/Token ID 保持可检索。""" + event = AuditService.build_event( + event_type='authorization_code_reused', + result='failure', + client_id='cs1', + token_id='rt1', + user_agent='Mozilla/5.0 rt1.record-id.secret-value suffix', + ) + assert event.user_agent == '[REDACTED]' + assert event.client_id == 'cs1' + assert event.token_id == 'rt1' + assert event.risk_level == 'high' + + +def test_audit_rejects_invalid_event_type_length() -> None: + """事件类型超过数据库列长度时必须拒绝。""" + with pytest.raises(ValueError): + AuditService.build_event(event_type='x' * 65, result='success') + + +@pytest.mark.asyncio +async def test_audit_uses_field_whitelist_and_caller_commit_boundary(audit_session: AsyncSession) -> None: + """未知顶层字段被丢弃,服务只 flush 不 commit。""" + event = await AuditService.record( + audit_session, + 'invalid_client', + 'failure', + detail={'safe': 'value', 'access_token': 'hidden'}, + not_a_column='must-not-be-stored', + ) + assert event.risk_level == 'high' + assert event.detail == {'safe': 'value'} + assert not audit_session.new + loaded = (await audit_session.execute(select(SysOAuthAuditLog))).scalars().first() + assert loaded is event + + +@pytest.mark.asyncio +async def test_audit_retention_moves_old_events_to_archive(audit_session: AsyncSession) -> None: + """超过在线保留期的审计事件必须转入归档表后再删除。""" + now = datetime.now(timezone.utc) + await AuditService.record(audit_session, 'old_event', 'success', create_time=now - timedelta(days=181)) + await AuditService.record(audit_session, 'fresh_event', 'success', create_time=now) + + archived = await OAuthAuditDao.archive_before(audit_session, now - timedelta(days=180)) + await audit_session.flush() + + assert archived == 1 + online = (await audit_session.execute(select(SysOAuthAuditLog))).scalars().all() + archive = (await audit_session.execute(select(SysOAuthAuditArchive))).scalars().all() + assert [row.event_type for row in online] == ['fresh_event'] + assert [row.event_type for row in archive] == ['old_event'] + + +@pytest.mark.asyncio +async def test_audit_admin_reads_project_safe_models_and_export(monkeypatch: pytest.MonkeyPatch) -> None: + """管理端分页和导出共用脱敏投影与筛选条件。""" + row = SimpleNamespace( + event_id=_ADMIN_EVENT_ID, + event_type='login', + result='success', + risk_level='normal', + trace_id='trace-1', + client_id='client-1', + resource_id=None, + user_id=3, + subject_id='subject-1', + sid='sid-1', + ip_address='127.0.0.1', + failure_code=None, + create_time=datetime.now(timezone.utc), + ) + calls: list[dict[str, object]] = [] + + async def list_admin_page(*args: object, **kwargs: object) -> list[object]: + calls.append(kwargs) + return [row] + + async def count_admin(*args: object, **kwargs: object) -> int: + return 1 + + monkeypatch.setattr(OAuthAuditDao, 'list_admin_page', list_admin_page) + monkeypatch.setattr(OAuthAuditDao, 'count_admin', count_admin) + monkeypatch.setattr('module_identity.service.audit_service.export_list2excel', lambda values: b'xlsx') + query = AuditPageQueryModel(page_num=2, page_size=_ADMIN_PAGE_SIZE, client_id='client-1') + + rows, total = await AuditService.list_admin_page(SimpleNamespace(), query) + exported = await AuditService.export_admin(SimpleNamespace(), query) + + assert total == _ADMIN_TOTAL + assert rows[0]['auditId'] == _ADMIN_EVENT_ID + assert 'detail' not in rows[0] + assert exported == b'xlsx' + assert calls[0]['offset'] == _ADMIN_PAGE_SIZE and calls[0]['limit'] == _ADMIN_PAGE_SIZE + assert calls[1]['offset'] == 0 and calls[1]['limit'] == _EXPORT_LIMIT + + +@pytest.mark.asyncio +async def test_interaction_failure_audit_rolls_back_and_wraps_failure(monkeypatch: pytest.MonkeyPatch) -> None: + """交互失败审计保持双重回滚和统一异常包装语义。""" + + class _Db: + def __init__(self) -> None: + self.rollbacks = 0 + + async def rollback(self) -> None: + self.rollbacks += 1 + + async def fail(*args: object, **kwargs: object) -> object: + raise RuntimeError('audit unavailable') + + monkeypatch.setattr(AuditService, 'record_independent', fail) + db = _Db() + with pytest.raises(OidcInteractionException) as raised: + await AuditService.record_interaction_failure(db, 'login_failed', failure_code='invalid_credentials') + + assert raised.value.error == 'server_error' + assert raised.value.status_code == _SERVICE_UNAVAILABLE + assert db.rollbacks == _ROLLBACK_COUNT diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_authentication_regressions.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_authentication_regressions.py new file mode 100644 index 000000000..ee5cd9809 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_authentication_regressions.py @@ -0,0 +1,909 @@ +from datetime import timedelta +from http.cookies import SimpleCookie +from types import SimpleNamespace +from unittest.mock import AsyncMock +from urllib.parse import parse_qs, urlsplit +from uuid import uuid4 + +import jwt +import pytest +import pytest_asyncio +from cryptography.hazmat.primitives.asymmetric import rsa +from sqlalchemy import delete, func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException, ServiceException +from module_admin.entity.do.dept_do import SysDept +from module_admin.entity.do.role_do import SysRole +from module_admin.entity.do.user_do import SysUser, SysUserRole +from module_admin.service.user_service import UserService +from module_identity.controller.interaction_controller import _login_response +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_client_do import SysOAuthClient, SysOAuthClientUri +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysSsoSession +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.interaction_vo import ChangePasswordModel, InteractionConsentModel, InteractionLoginModel +from module_identity.entity.vo.oauth_resource_vo import ScopeStatusModel +from module_identity.entity.vo.oauth_session_vo import GrantPageQueryModel +from module_identity.entity.vo.protocol_vo import AuthorizeRequest +from module_identity.redis_keys import OidcRedisKey +from module_identity.security.jwt_profile import decode_access_token +from module_identity.security.pkce import generate_code_challenge +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import ( + AuthorizationCodeService, + AuthorizationService, + InteractionCompletionService, +) +from module_identity.service.consent_service import ConsentService, InteractionConsentService +from module_identity.service.identity_service import CredentialAuthenticationResult, CredentialAuthenticationService +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.interaction_service import InteractionLoginService, InteractionService +from module_identity.service.oauth_management_service import OAuthResourceManagementService +from module_identity.service.oauth_session_management_service import OAuthSessionManagementService +from module_identity.service.session_service import SsoSessionService +from module_identity.service.token_protocol_service import IntrospectionService, UserInfoService +from module_identity.service.token_service import TokenResult, TokenService +from tests.module_identity.support.redis_fakes import FakeRedis +from utils.oidc_util import OidcUtil +from utils.pwd_util import PwdUtil +from utils.time_util import TimezoneUtil + +_PEPPER = 'authentication-regression-pepper-' + 'x' * 32 +_VERIFIER = 'v' * 64 +_REMEMBER_SECONDS = 7 * 24 * 60 * 60 +_RESOURCE_SECONDS = 60 +_CLIENT_SECONDS = 600 + + +@pytest_asyncio.fixture +async def auth_flow(data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch) -> SimpleNamespace: + """配置一个业务客户端、机器客户端、资源服务和真实用户会话。""" + + now = TimezoneUtil.utc_now().replace(microsecond=0) + monkeypatch.setattr(TimezoneUtil, 'utc_now', staticmethod(lambda: now)) + for name, value in { + 'oidc_enabled': True, + 'oidc_issuer': 'https://auth.example.com', + 'oidc_token_hash_pepper': _PEPPER, + 'oidc_access_token_ttl_seconds': _CLIENT_SECONDS, + 'oidc_max_access_token_ttl_seconds': 1800, + 'oidc_id_token_ttl_seconds': 300, + 'oidc_refresh_token_idle_seconds': 3600, + 'oidc_refresh_token_absolute_seconds': 7200, + 'oidc_sso_idle_seconds': 1800, + 'oidc_sso_absolute_seconds': 8 * 60 * 60, + 'oidc_sso_remember_absolute_seconds': _REMEMBER_SECONDS, + 'oidc_sso_cookie_name': '__Host-ruoyi-sso', + 'oidc_sso_cookie_secure': True, + 'oidc_sso_cookie_samesite': 'lax', + 'oidc_sso_cookie_domain': '', + }.items(): + monkeypatch.setattr(OidcConfig, name, value) + connection = await data_session.connection() + await connection.run_sync( + lambda sync: SysUser.metadata.create_all( + sync, tables=[SysDept.__table__, SysRole.__table__, SysUserRole.__table__] + ) + ) + user = SysUser(user_id=2001, user_name='alice', nick_name='Alice', status='0', del_flag='0') + subject = SysIdentitySubject(identity_id=1, user_id=user.user_id, subject_id=str(uuid4()), auth_version=1) + app = SysOAuthClient( + client_pk=1, + client_id='business-app', + client_name='Business app', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code', 'refresh_token'], + response_types=['code'], + ) + machine = SysOAuthClient( + client_pk=2, + client_id='machine', + client_name='Machine', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['client_credentials'], + response_types=[], + access_token_ttl_seconds=_CLIENT_SECONDS, + ) + resource_client = SysOAuthClient( + client_pk=3, + client_id='resource-server', + client_name='Resource server', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['client_credentials'], + response_types=[], + ) + resource = SysOAuthResource( + resource_pk=1, + resource_id='business-api', + resource_name='Business API', + audience='https://api.example.com', + introspection_client_pk=3, + access_token_ttl_seconds=_RESOURCE_SECONDS, + allowed_claims=['sub'], + create_by='test', + update_by='test', + ) + scope = SysOAuthScope( + scope_pk=2, + scope_code='api.read', + scope_name='Read API', + scope_type='resource', + resource_pk=1, + claims=[], + create_by='test', + update_by='test', + ) + data_session.add_all( + [ + user, + subject, + app, + machine, + resource_client, + resource, + scope, + SysOAuthClientUri(client_pk=1, uri_type='redirect', uri='https://app.example.com/callback'), + SysOAuthScope( + scope_pk=1, + scope_code='openid', + scope_name='OpenID', + scope_type='identity', + claims=['sub'], + consent_required=0, + create_by='test', + update_by='test', + ), + SysOAuthScope( + scope_pk=3, + scope_code='offline_access', + scope_name='Offline access', + scope_type='identity', + claims=[], + create_by='test', + update_by='test', + ), + SysOAuthClientScope(client_pk=1, scope_pk=1), + SysOAuthClientScope(client_pk=1, scope_pk=2), + SysOAuthClientScope(client_pk=1, scope_pk=3), + SysOAuthClientScope(client_pk=2, scope_pk=2), + SysOAuthClientResource(client_pk=1, resource_pk=1), + SysOAuthClientResource(client_pk=2, resource_pk=1), + ] + ) + await data_session.commit() + redis = FakeRedis() + coordinator = AfterCommitCoordinator() + cookie, session = await SsoSessionService.create( + data_session, + redis, + user.user_id, + subject.subject_id, + subject.auth_version, + 'urn:ruoyi:acr:pwd', + ('pwd',), + pepper=_PEPPER, + now=now, + coordinator=coordinator, + ) + await coordinator.commit(data_session) + return SimpleNamespace( + db=data_session, + redis=redis, + now=now, + user=user, + subject=subject, + app=app, + machine=machine, + resource=resource, + scope=scope, + cookie=cookie, + session=session, + signer=rsa.generate_private_key(public_exponent=65537, key_size=2048), + caller=OAuthClientPrincipal(resource_client.client_id, 'confidential', 'client_secret_basic'), + ) + + +def _authorize_request(flow: SimpleNamespace, *, offline: bool = False) -> dict[str, str]: + """构造携带 PKCE 和 Resource 的授权请求。""" + + return { + 'client_id': flow.app.client_id, + 'redirect_uri': 'https://app.example.com/callback', + 'response_type': 'code', + 'scope': 'openid api.read' + (' offline_access' if offline else ''), + 'nonce': 'regression-nonce', + 'code_challenge': generate_code_challenge(_VERIFIER), + 'code_challenge_method': 'S256', + 'resource': flow.resource.audience, + } + + +async def _user_token( + flow: SimpleNamespace, *, remember: bool = False, offline: bool = False +) -> tuple[TokenResult, SysOAuthGrant | None]: + """通过真实同意和授权码流程签发用户令牌。""" + + request = AuthorizeRequest(**_authorize_request(flow, offline=offline)) + context = await AuthorizationService.validate_request(flow.db, request) + consent = await ConsentService.submit_consent( + flow.db, + context, + True, + context.scopes, + remember, + user_id=flow.user.user_id, + subject_id=flow.subject.subject_id, + ) + code = await AuthorizationCodeService.issue( + flow.redis, + { + 'clientPk': flow.app.client_pk, + 'redirectUri': request.redirect_uri, + 'userId': flow.user.user_id, + 'subjectId': flow.subject.subject_id, + 'authVersion': flow.subject.auth_version, + 'sid': flow.session.sid, + 'grantId': consent.grant.grant_id if consent.grant is not None else None, + 'scopes': list(consent.scopes), + 'resources': [flow.resource.audience], + 'nonce': request.nonce, + 'codeChallenge': request.code_challenge, + 'codeChallengeMethod': 'S256', + 'authTime': flow.now.isoformat(), + }, + pepper=_PEPPER, + ) + token = await TokenService.issue_token_request( + flow.db, + flow.redis, + { + 'grant_type': 'authorization_code', + 'client_id': flow.app.client_id, + 'code': code, + 'redirect_uri': request.redirect_uri, + 'code_verifier': _VERIFIER, + }, + client_id=flow.app.client_id, + signing_key=flow.signer, + kid='regression-key', + ) + return token, consent.grant + + +async def _machine_token(flow: SimpleNamespace) -> TokenResult: + """为已认证的机器 Client 签发 Resource 访问令牌。""" + + return await TokenService.client_credentials( + flow.db, + {'grant_type': 'client_credentials', 'scope': 'api.read', 'resource': flow.resource.audience}, + OAuthClientPrincipal(flow.machine.client_id, 'confidential', 'client_secret_basic'), + signing_key=flow.signer, + kid='regression-key', + now=flow.now, + ) + + +async def _introspect(flow: SimpleNamespace, token: str) -> dict[str, object]: + """使用独立 Resource Client 查询令牌的实时状态。""" + + return await IntrospectionService.introspect( + flow.db, + flow.redis, + token, + flow.caller, + verification_key=flow.signer.public_key(), + now=flow.now, + ) + + +def _claims(flow: SimpleNamespace, token: str) -> dict[str, object]: + """验证 Access Token 的签名、Issuer 和 Resource Audience。""" + + return decode_access_token( + token, + verification_key=flow.signer.public_key(), + issuer=OidcConfig.oidc_issuer, + audience=flow.resource.audience, + ) + + +def _sign(flow: SimpleNamespace, claims: dict[str, object]) -> str: + """签署指定 Claims,用于验证内省对授权来源的校验。""" + + return jwt.encode( + claims, + flow.signer, + algorithm='RS256', + headers={'kid': 'regression-key', 'typ': 'at+jwt'}, + ) + + +@pytest.mark.asyncio +async def test_one_time_consent_has_revocable_grant_without_remembering_consent(auth_flow: SimpleNamespace) -> None: + """一次性同意有可撤销记录,但不授予离线续期和后续免确认资格。""" + + token, grant = await _user_token(auth_flow) + assert grant is not None and token.refresh_token is None + assert grant.remembered_scopes == [] and grant.remembered_resources == [] + assert _claims(auth_flow, token.access_token)['grant_id'] == grant.grant_id + assert await auth_flow.db.scalar(select(func.count()).select_from(SysOAuthGrant)) == 1 + context = await AuthorizationService.validate_request( + auth_flow.db, AuthorizeRequest(**_authorize_request(auth_flow)) + ) + assert not ConsentService.consent_is_satisfied(context, grant) + result = await _introspect(auth_flow, token.access_token) + assert result['active'] is True + assert result['username'] == 'alice' + + +@pytest.mark.asyncio +@pytest.mark.parametrize('invalid_state', ['policy', 'scope', 'session', 'user', 'version', 'jti']) +async def test_one_time_consent_still_checks_current_security_state( + auth_flow: SimpleNamespace, + invalid_state: str, +) -> None: + """一次性同意的 Access Token 仍受当前授权策略和用户安全状态约束。""" + + token, _ = await _user_token(auth_flow) + assert (await _introspect(auth_flow, token.access_token))['active'] is True + if invalid_state == 'policy': + auth_flow.app.policy_version += 1 + elif invalid_state == 'scope': + auth_flow.scope.status = '1' + elif invalid_state == 'session': + auth_flow.session.status = 'revoked' + elif invalid_state == 'user': + auth_flow.user.status = '1' + elif invalid_state == 'version': + auth_flow.subject.auth_version += 1 + else: + await auth_flow.redis.set(OidcRedisKey.revoked_jti(_claims(auth_flow, token.access_token)['jti']), '1', ex=60) + await auth_flow.db.commit() + assert await _introspect(auth_flow, token.access_token) == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize(('remember', 'offline'), [(False, False), (True, False), (False, True)]) +async def test_saved_grant_revocation_cannot_be_bypassed_by_a_later_grant( + auth_flow: SimpleNamespace, + remember: bool, + offline: bool, +) -> None: + """管理端撤销各类授权后,新授权不能恢复旧令牌,其他用户会话不受影响。""" + + token, grant = await _user_token(auth_flow, remember=remember, offline=offline) + assert grant is not None + if offline: + assert token.refresh_token is not None + token = await TokenService.issue_token_request( + auth_flow.db, + auth_flow.redis, + { + 'grant_type': 'refresh_token', + 'client_id': auth_flow.app.client_id, + 'refresh_token': token.refresh_token, + }, + client_id=auth_flow.app.client_id, + signing_key=auth_flow.signer, + kid='regression-key', + ) + assert (await _introspect(auth_flow, token.access_token))['active'] is True + assert await UserInfoService.build(auth_flow.db, _claims(auth_flow, token.access_token), auth_flow.redis) + assert ( + await OAuthSessionManagementService.revoke_grants(auth_flow.db, [grant.grant_id], 'admin', '撤销应用授权') == 1 + ) + assert await _introspect(auth_flow, token.access_token) == {'active': False} + with pytest.raises(ValueError, match='授权已失效'): + await UserInfoService.build(auth_flow.db, _claims(auth_flow, token.access_token), auth_flow.redis) + if offline: + assert await _introspect(auth_flow, token.refresh_token) == {'active': False} + with pytest.raises(OAuthProtocolException) as error: + await TokenService.refresh_token( + auth_flow.db, + { + 'grant_type': 'refresh_token', + 'client_id': auth_flow.app.client_id, + 'refresh_token': token.refresh_token, + }, + OAuthClientPrincipal(auth_flow.app.client_id, 'public', 'none'), + signing_key=auth_flow.signer, + kid='regression-key', + now=auth_flow.now, + ) + assert error.value.error == 'invalid_grant' + await auth_flow.db.refresh(auth_flow.session) + assert auth_flow.session.status == 'active' + new_token, new_grant = await _user_token(auth_flow, remember=True, offline=offline) + assert new_grant.grant_id != grant.grant_id + assert (await _introspect(auth_flow, new_token.access_token))['active'] is True + assert await _introspect(auth_flow, token.access_token) == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize('remember', [False, True]) +async def test_legacy_user_tokens_without_grant_binding_require_reauthorization( + auth_flow: SimpleNamespace, remember: bool +) -> None: + """旧令牌没有具体 Grant 绑定时拒绝访问,即使用户有其他有效授权。""" + + token, _ = await _user_token(auth_flow, remember=remember) + claims = _claims(auth_flow, token.access_token) + claims.pop('grant_id', None) + claims.pop('client_policy_version', None) + assert await _introspect(auth_flow, _sign(auth_flow, claims)) == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'invalid_context', + ['missing_policy', 'missing_grant', 'boolean_policy', 'empty_grant', 'unknown_grant', 'offline_without_grant'], +) +async def test_incomplete_or_invalid_authorization_context_is_inactive( + auth_flow: SimpleNamespace, + invalid_context: str, +) -> None: + """授权来源 Claims 缺失或不合法时,内省返回 inactive。""" + + token, _ = await _user_token(auth_flow) + claims = _claims(auth_flow, token.access_token) + claims.update(grant_id=None, client_policy_version=auth_flow.app.policy_version) + if invalid_context == 'missing_policy': + claims.pop('client_policy_version') + elif invalid_context == 'missing_grant': + claims.pop('grant_id') + elif invalid_context == 'boolean_policy': + claims['client_policy_version'] = True + elif invalid_context == 'empty_grant': + claims['grant_id'] = '' + elif invalid_context == 'unknown_grant': + claims['grant_id'] = str(uuid4()) + else: + claims['scope'] += ' offline_access' + assert await _introspect(auth_flow, _sign(auth_flow, claims)) == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ('client_ttl', 'resource_ttl', 'expected'), + [(_CLIENT_SECONDS, _RESOURCE_SECONDS, _RESOURCE_SECONDS), (30, _RESOURCE_SECONDS, 30), (3000, 2400, 1800)], +) +async def test_machine_token_honours_client_resource_and_platform_ttl( + auth_flow: SimpleNamespace, + client_ttl: int, + resource_ttl: int, + expected: int, +) -> None: + """机器令牌的响应有效期和 JWT 有效期均遵守 Client、Resource 及平台上限。""" + + auth_flow.machine.access_token_ttl_seconds = client_ttl + auth_flow.resource.access_token_ttl_seconds = resource_ttl + await auth_flow.db.commit() + token = await _machine_token(auth_flow) + claims = _claims(auth_flow, token.access_token) + assert token.expires_in == expected + assert claims['exp'] - claims['iat'] == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize('legacy', [False, True]) +async def test_scope_disable_immediately_invalidates_machine_tokens(auth_flow: SimpleNamespace, legacy: bool) -> None: + """Scope 禁用立即使机器令牌失效,新版令牌在 Scope 恢复后仍受策略版本约束。""" + + token = (await _machine_token(auth_flow)).access_token + if legacy: + claims = _claims(auth_flow, token) + claims.pop('client_policy_version', None) + token = _sign(auth_flow, claims) + assert (await _introspect(auth_flow, token))['active'] is True + await OAuthResourceManagementService.change_scope_status( + auth_flow.db, + ScopeStatusModel(scope_code='api.read', status='1'), + actor='test', + ) + with pytest.raises(OAuthProtocolException) as exc: + await _machine_token(auth_flow) + assert exc.value.error == 'invalid_scope' + assert await _introspect(auth_flow, token) == {'active': False} + if not legacy: + await OAuthResourceManagementService.change_scope_status( + auth_flow.db, + ScopeStatusModel(scope_code='api.read', status='0'), + actor='test', + ) + assert await _introspect(auth_flow, token) == {'active': False} + assert (await _introspect(auth_flow, (await _machine_token(auth_flow)).access_token))['active'] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize('change', ['unbind', 'identity_scope', 'different_resource']) +async def test_machine_tokens_recheck_current_scope_bindings(auth_flow: SimpleNamespace, change: str) -> None: + """机器令牌内省重新检查 Scope 的 Client 绑定、类型和所属 Resource。""" + + token = await _machine_token(auth_flow) + assert (await _introspect(auth_flow, token.access_token))['active'] is True + if change == 'unbind': + await auth_flow.db.execute( + delete(SysOAuthClientScope).where(SysOAuthClientScope.client_pk == auth_flow.machine.client_pk) + ) + elif change == 'identity_scope': + auth_flow.scope.scope_type = 'identity' + else: + auth_flow.scope.resource_pk = 99 + await auth_flow.db.commit() + assert await _introspect(auth_flow, token.access_token) == {'active': False} + + +@pytest.mark.asyncio +async def test_authorize_slides_idle_timeout_but_never_extends_absolute_expiry( + auth_flow: SimpleNamespace, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """授权请求更新 SSO 空闲期限,并保留 Session 的绝对过期上限。""" + + flow = auth_flow + absolute = flow.now + timedelta(minutes=40) + flow.session.absolute_expires_at = absolute + await flow.db.commit() + for minutes in (25, 31): + current = flow.now + timedelta(minutes=minutes) + monkeypatch.setattr(TimezoneUtil, 'utc_now', staticmethod(lambda current=current: current)) + await AuthorizationService.process_authorization_request( + flow.db, + flow.redis, + _authorize_request(flow), + sso_cookie=flow.cookie, + ) + await flow.db.refresh(flow.session) + assert TimezoneUtil.to_utc(flow.session.idle_expires_at) == absolute + assert TimezoneUtil.to_utc(flow.session.absolute_expires_at) == absolute + assert TimezoneUtil.to_utc(flow.session.last_seen_at) == current + assert flow.session.status == 'active' + monkeypatch.setattr(TimezoneUtil, 'utc_now', staticmethod(lambda: absolute)) + coordinator = AfterCommitCoordinator() + assert await AuthorizationService._load_sso_session(flow.db, flow.redis, flow.cookie, coordinator) is None + await coordinator.commit(flow.db) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('remember', [False, True]) +@pytest.mark.parametrize('force_password_change', [False, True]) +async def test_login_cookie_persistence_survives_forced_password_change( + auth_flow: SimpleNamespace, + monkeypatch: pytest.MonkeyPatch, + remember: bool, + force_password_change: bool, +) -> None: + """普通登录及强制改密后的 Cookie 持久化行为均遵守保持登录选项。""" + + flow = auth_flow + context = await AuthorizationService.validate_request(flow.db, AuthorizeRequest(**_authorize_request(flow))) + created = await InteractionService.create(flow.redis, context.to_internal_payload('cookie-test'), pepper=_PEPPER) + flow.user.password = PwdUtil.get_password_hash('old-password') + await flow.db.commit() + monkeypatch.setattr( + CredentialAuthenticationService, + 'authenticate_oidc', + AsyncMock( + return_value=CredentialAuthenticationResult( + user=flow.user, + dept=None, + acr='urn:ruoyi:acr:pwd', + amr=('pwd',), + remember_me=remember, + password_change_required=force_password_change, + password_change_reason='initial_password' if force_password_change else None, + ) + ), + ) + outcome = await InteractionLoginService.login( + flow.redis, + created.interaction_id, + InteractionLoginModel(userName='alice', password='old-password', rememberMe=remember), + flow.db, + created.csrf_token, + ) + if force_password_change: + assert outcome.cookie is None + monkeypatch.setattr(UserService, 'validate_password_services', AsyncMock()) + outcome = await InteractionLoginService.change_password( + flow.redis, + created.interaction_id, + ChangePasswordModel( + oldPassword='old-password', + newPassword='new-password-A1!', + confirmPassword='new-password-A1!', + ), + flow.db, + created.csrf_token, + ) + assert outcome.failure_message is None + response = _login_response(outcome) + cookies = SimpleCookie() + cookies.load(response.headers['set-cookie']) + cookie = cookies[OidcConfig.oidc_sso_cookie_name] + assert cookie['secure'] and cookie['httponly'] and cookie['samesite'] == 'lax' and cookie['path'] == '/' + assert not cookie['domain'] + sid, _ = OidcUtil.parse_sso_cookie(cookie.value) + session = await flow.db.get(SysSsoSession, sid) + assert bool(session.remember_me) is remember + if remember: + remaining = int((TimezoneUtil.to_utc(session.absolute_expires_at) - flow.now).total_seconds()) + assert int(cookie['max-age']) == remaining == _REMEMBER_SECONDS + else: + assert not cookie['max-age'] and not cookie['expires'] + + +@pytest.mark.asyncio +@pytest.mark.parametrize('offline', [False, True]) +async def test_block_and_unblock_require_new_authorization(auth_flow: SimpleNamespace, offline: bool) -> None: + """禁止访问阻断当前令牌和授权确认,解除后只有新授权可以访问。""" + + flow = auth_flow + token, grant = await _user_token(flow, offline=offline) + client_id, user_id, grant_id = flow.app.client_id, flow.user.user_id, grant.grant_id + context = await AuthorizationService.validate_request(flow.db, AuthorizeRequest(**_authorize_request(flow))) + code_payload = { + 'clientPk': flow.app.client_pk, + 'redirectUri': context.redirect_uri, + 'userId': user_id, + 'subjectId': flow.subject.subject_id, + 'authVersion': flow.subject.auth_version, + 'sid': flow.session.sid, + 'grantId': grant_id, + 'scopes': list(context.scopes), + 'resources': list(context.resources), + 'nonce': context.nonce, + 'codeChallenge': context.code_challenge, + 'codeChallengeMethod': context.code_challenge_method, + 'authTime': flow.now.isoformat(), + } + pending_code = await AuthorizationCodeService.issue(flow.redis, code_payload, pepper=_PEPPER) + count = await OAuthSessionManagementService.set_access(flow.db, user_id, client_id, True, 'admin', '禁止访问') + assert count == 1 + assert await _introspect(flow, token.access_token) == {'active': False} + with pytest.raises(ValueError): + await UserInfoService.build(flow.db, _claims(flow, token.access_token), flow.redis) + with pytest.raises(OAuthProtocolException) as error: + await ConsentService.submit_consent( + flow.db, + context, + True, + context.scopes, + True, + user_id=user_id, + subject_id=flow.subject.subject_id, + ) + assert error.value.error == 'access_denied' + with pytest.raises(OAuthProtocolException) as error: + await TokenService.authorization_code( + flow.db, + flow.redis, + { + 'grant_type': 'authorization_code', + 'client_id': client_id, + 'code': pending_code, + 'redirect_uri': context.redirect_uri, + 'code_verifier': _VERIFIER, + }, + OAuthClientPrincipal(client_id, 'public', 'none'), + signing_key=flow.signer, + kid='regression-key', + now=flow.now, + ) + assert error.value.error == 'invalid_grant' + if offline: + assert await _introspect(flow, token.refresh_token) == {'active': False} + rows, total = await OAuthSessionManagementService.list_grants(flow.db, GrantPageQueryModel(accessStatus='blocked')) + assert total == 1 and rows[0].access_status == 'blocked' and rows[0].access_reason == '禁止访问' + assert await OAuthSessionManagementService.set_access(flow.db, user_id, client_id, False, 'admin', '解除禁止') == 0 + assert await _introspect(flow, token.access_token) == {'active': False} + await flow.db.refresh(grant) + assert grant.status == 'revoked' + assert grant.remembered_scopes == [] and grant.remembered_resources == [] + fresh, fresh_grant = await _user_token(flow, offline=offline) + assert fresh_grant.grant_id != grant_id + assert (await _introspect(flow, fresh.access_token))['active'] is True + assert await _introspect(flow, token.access_token) == {'active': False} + await flow.db.refresh(flow.session) + assert flow.session.status == 'active' + + +@pytest.mark.asyncio +async def test_block_rejects_authorize_before_showing_consent(auth_flow: SimpleNamespace) -> None: + """已有 SSO 时,禁止访问在授权入口直接拒绝,不依赖同意页或令牌端点兜底。""" + + flow = auth_flow + raw, cookie = _authorize_request(flow), flow.cookie + await OAuthSessionManagementService.set_access( + flow.db, flow.user.user_id, flow.app.client_id, True, 'admin', '禁止' + ) + with pytest.raises(OAuthProtocolException) as error: + await AuthorizationService.process_authorization_request(flow.db, flow.redis, raw, sso_cookie=cookie) + assert error.value.error == 'access_denied' and error.value.can_redirect + + +@pytest.mark.asyncio +async def test_pre_authorized_flow_also_creates_revocable_grant(auth_flow: SimpleNamespace) -> None: + """可信应用跳过确认页面时,授权码仍绑定后台可撤销的具体 Grant。""" + + flow = auth_flow + flow.app.require_consent = 0 + await flow.db.commit() + result = await AuthorizationService.process_authorization_request( + flow.db, + flow.redis, + _authorize_request(flow), + sso_cookie=flow.cookie, + ) + code = parse_qs(urlsplit(result.location).query)['code'][0] + payload = await AuthorizationCodeService.consume(flow.redis, code, pepper=_PEPPER) + grant = await OAuthGrantDao.get_by_grant_id(flow.db, payload['grantId']) + assert grant is not None and grant.user_id == flow.user.user_id and not grant.remembered_scopes + assert await OAuthSessionManagementService.revoke_grants(flow.db, [grant.grant_id], 'admin', '撤销') == 1 + + +@pytest.mark.asyncio +async def test_batch_revoke_covers_all_grants_for_selected_user_client(auth_flow: SimpleNamespace) -> None: + """同一用户应用的重复历史有效记录一并撤销,其他用户及机器访问不受影响。""" + + flow = auth_flow + token, grant = await _user_token(flow, remember=True) + duplicate = SysOAuthGrant( + grant_id=str(uuid4()), + user_id=flow.user.user_id, + subject_id=flow.subject.subject_id, + client_pk=flow.app.client_pk, + granted_scopes=list(grant.granted_scopes), + granted_resources=list(grant.granted_resources), + client_policy_version=flow.app.policy_version, + remembered_scopes=list(grant.remembered_scopes), + status='active', + ) + other_user = SysUser(user_id=2002, user_name='bob', nick_name='Bob', status='0', del_flag='0') + other_grant = SysOAuthGrant( + grant_id=str(uuid4()), + user_id=other_user.user_id, + subject_id=str(uuid4()), + client_pk=flow.app.client_pk, + granted_scopes=list(grant.granted_scopes), + granted_resources=list(grant.granted_resources), + client_policy_version=flow.app.policy_version, + status='active', + ) + flow.db.add_all([duplicate, other_user, other_grant]) + await flow.db.commit() + machine = await _machine_token(flow) + assert await OAuthSessionManagementService.revoke_grants( + flow.db, [grant.grant_id, duplicate.grant_id], 'admin', '撤销' + ) == len([grant, duplicate]) + assert await _introspect(flow, token.access_token) == {'active': False} + assert (await _introspect(flow, machine.access_token))['active'] is True + for row in (grant, duplicate, other_grant, flow.session): + await flow.db.refresh(row) + assert grant.status == duplicate.status == 'revoked' + assert other_grant.status == flow.session.status == 'active' + + +@pytest.mark.asyncio +async def test_failed_access_policy_audit_rolls_back_revocation( + auth_flow: SimpleNamespace, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """禁止操作的审计失败时,访问策略与令牌撤销必须同时回滚。""" + + flow = auth_flow + token, grant = await _user_token(flow) + user_id, client_pk, client_id = flow.user.user_id, flow.app.client_pk, flow.app.client_id + + monkeypatch.setattr(AuditService, 'record', AsyncMock(side_effect=RuntimeError('audit unavailable'))) + with pytest.raises(ServiceException): + await OAuthSessionManagementService.set_access(flow.db, user_id, client_id, True, 'admin', '禁止') + assert not await OAuthAccessPolicyDao.is_blocked(flow.db, user_id, client_pk) + await flow.db.refresh(grant) + assert grant.status == 'active' + assert (await _introspect(flow, token.access_token))['active'] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize('confirm', [False, True]) +async def test_interactive_completion_binds_a_grant(auth_flow: SimpleNamespace, confirm: bool) -> None: + """登录后直接完成与确认后完成两条交互路径都绑定可撤销授权。""" + + flow = auth_flow + flow.app.require_consent = int(confirm) + await flow.db.commit() + context = await AuthorizationService.validate_request(flow.db, AuthorizeRequest(**_authorize_request(flow))) + payload = context.to_internal_payload('pending') + payload.update( + { + 'authenticatedSid': flow.session.sid, + 'userId': flow.user.user_id, + 'subjectId': flow.subject.subject_id, + 'authVersion': flow.subject.auth_version, + } + ) + created = await InteractionService.create(flow.redis, payload, pepper=_PEPPER) + if confirm: + await InteractionConsentService.consent( + flow.redis, + created.interaction_id, + InteractionConsentModel(approved=True, scopes=list(context.scopes), rememberConsent=False), + flow.db, + created.csrf_token, + ) + result = await InteractionCompletionService.complete(flow.redis, created.interaction_id, flow.db) + code = parse_qs(urlsplit(result.location).query)['code'][0] + payload = await AuthorizationCodeService.consume(flow.redis, code, pepper=_PEPPER) + grant = await OAuthGrantDao.get_by_grant_id(flow.db, payload['grantId']) + assert grant is not None and grant.status == 'active' and grant.remembered_scopes == [] + + +@pytest.mark.asyncio +async def test_expired_grant_is_not_reactivated_by_new_consent(auth_flow: SimpleNamespace) -> None: + """授权先于 Access Token 过期后,再次同意也不能恢复旧令牌。""" + + flow = auth_flow + old_token, old_grant = await _user_token(flow) + old_grant.expires_at = flow.now - timedelta(seconds=1) + await flow.db.commit() + assert await _introspect(flow, old_token.access_token) == {'active': False} + token, grant = await _user_token(flow) + assert grant.grant_id != old_grant.grant_id + assert (await _introspect(flow, token.access_token))['active'] is True + assert await _introspect(flow, old_token.access_token) == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize('blocked', [False, True]) +async def test_revocation_between_consent_and_completion_prevents_code( + auth_flow: SimpleNamespace, + monkeypatch: pytest.MonkeyPatch, + blocked: bool, +) -> None: + """同意后、签码前发生撤销或禁止,不得继续发码或复用授权记录。""" + + flow = auth_flow + context = await AuthorizationService.validate_request(flow.db, AuthorizeRequest(**_authorize_request(flow))) + payload = context.to_internal_payload('pending') + payload.update( + { + 'authenticatedSid': flow.session.sid, + 'userId': flow.user.user_id, + 'subjectId': flow.subject.subject_id, + 'authVersion': flow.subject.auth_version, + } + ) + created = await InteractionService.create(flow.redis, payload, pepper=_PEPPER) + await InteractionConsentService.consent( + flow.redis, + created.interaction_id, + InteractionConsentModel(approved=True, scopes=list(context.scopes), rememberConsent=False), + flow.db, + created.csrf_token, + ) + record = await InteractionService.get_record(flow.redis, created.interaction_id) + if blocked: + await OAuthSessionManagementService.set_access( + flow.db, flow.user.user_id, flow.app.client_id, True, 'admin', '禁止' + ) + else: + await OAuthSessionManagementService.revoke_grants(flow.db, [record['grantId']], 'admin', '撤销') + issue = AsyncMock() + monkeypatch.setattr(AuthorizationCodeService, 'issue', issue) + with pytest.raises(OAuthProtocolException) as error: + await InteractionCompletionService.complete(flow.redis, created.interaction_id, flow.db) + assert error.value.error == 'access_denied' + issue.assert_not_awaited() diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_authorization_code_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_authorization_code_service.py new file mode 100644 index 000000000..d0b3939ae --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_authorization_code_service.py @@ -0,0 +1,161 @@ +import asyncio + +import pytest + +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.authorization_service import AuthorizationCodeReuseError, AuthorizationCodeService +from tests.module_identity.support.redis_fakes import FakeRedis + +_PEPPER = 'authorization-code-test-pepper-' + 'x' * 32 +_CHALLENGE = 'A' * 43 +_CLIENT_PK = 1001 +_LONG_CODE_TTL = 600 + + +def _payload(**overrides: object) -> dict[str, object]: + """构造服务端授权码白名单载荷。""" + value: dict[str, object] = { + 'clientPk': 1001, + 'redirectUri': 'https://portal.example/callback', + 'userId': 2001, + 'subjectId': 'subject-2001', + 'authVersion': 4, + 'sid': 'sid-2001', + 'grantId': 'grant-2001', + 'scopes': ['openid', 'profile'], + 'resources': ['https://api.example'], + 'nonce': 'nonce-2001', + 'codeChallenge': _CHALLENGE, + 'codeChallengeMethod': 'S256', + 'authTime': '2026-08-24T04:00:00+00:00', + } + value.update(overrides) + return value + + +@pytest.mark.asyncio +async def test_issue_uses_ac_key_nx_ttl_and_never_stores_plain_code() -> None: + """验证授权码使用独立 key、NX/TTL,并只保存摘要。""" + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, _payload(), ttl_seconds=90, pepper=_PEPPER) + assert code.startswith('ac1.') + key = OidcRedisKey.authorization_code(code.split('.')[1]) + stored = redis.values[key][0] + assert code not in stored + assert 'codeHash' in stored + assert redis.set_calls[-1][1] == {'ex': 90, 'nx': True} + + +@pytest.mark.asyncio +@pytest.mark.parametrize('ttl', [0, False, -1]) +async def test_issue_rejects_explicit_invalid_ttl(ttl: object) -> None: + """验证显式零值、布尔值和负数不会回退到默认 Code TTL。""" + redis = FakeRedis() + with pytest.raises(ValueError): + await AuthorizationCodeService.issue(redis, _payload(), ttl_seconds=ttl, pepper=_PEPPER) # type: ignore[arg-type] + assert not redis.values + + +@pytest.mark.asyncio +async def test_wrong_secret_does_not_consume_and_correct_secret_consumes_once() -> None: + """验证错误 Secret 不删除 Code,正确消费具备一次性语义。""" + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, _payload(), pepper=_PEPPER) + parts = code.split('.') + replacement = 'A' if parts[2][0] != 'A' else 'B' + wrong = f'{parts[0]}.{parts[1]}.{replacement}{parts[2][1:]}' + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationCodeService.consume(redis, wrong, pepper=_PEPPER) + assert raised.value.error == 'invalid_grant' + assert await redis.get(OidcRedisKey.authorization_code(parts[1])) is not None + assert await redis.get(OidcRedisKey.authorization_code_consumed(parts[1])) is None + consumed = await AuthorizationCodeService.consume(redis, code, pepper=_PEPPER) + assert consumed['clientPk'] == _CLIENT_PK + assert await redis.get(OidcRedisKey.authorization_code_consumed(parts[1])) == 'consumed' + assert (await AuthorizationCodeService.consumed_payload(redis, code, pepper=_PEPPER))['grantId'] == 'grant-2001' + with pytest.raises(OAuthProtocolException) as reused: + await AuthorizationCodeService.consume(redis, code, pepper=_PEPPER) + assert isinstance(reused.value, AuthorizationCodeReuseError) + + +@pytest.mark.asyncio +async def test_reuse_tombstone_ttl_covers_long_authorization_code_ttl( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """验证授权码 TTL 超过最小墓碑 TTL 时,重用墓碑覆盖完整 Code 生命周期。""" + monkeypatch.setattr(OidcConfig, 'oidc_authorization_code_ttl_seconds', _LONG_CODE_TTL) + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, _payload(), pepper=_PEPPER) + + await AuthorizationCodeService.consume(redis, code, pepper=_PEPPER) + + assert redis.eval_calls[-1][1][-1] == _LONG_CODE_TTL + + +@pytest.mark.asyncio +async def test_concurrent_correct_consumption_only_succeeds_once() -> None: + """验证两个并发正确消费请求只有一个能取得载荷。""" + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, _payload(), pepper=_PEPPER) + results = await asyncio.gather( + AuthorizationCodeService.consume(redis, code, pepper=_PEPPER), + AuthorizationCodeService.consume(redis, code, pepper=_PEPPER), + return_exceptions=True, + ) + assert sum(isinstance(result, dict) for result in results) == 1 + assert sum(isinstance(result, OAuthProtocolException) for result in results) == 1 + assert sum(isinstance(result, AuthorizationCodeReuseError) for result in results) == 1 + + +@pytest.mark.asyncio +async def test_invalid_payload_and_format_fail_closed_without_redis_write() -> None: + """验证缺失、未知对象和畸形 Code 均 fail closed。""" + redis = FakeRedis() + with pytest.raises(ValueError): + await AuthorizationCodeService.issue(redis, _payload(extra='unknown'), pepper=_PEPPER) + with pytest.raises(ValueError): + await AuthorizationCodeService.issue(redis, _payload(scopes=['openid'], codeChallenge='short'), pepper=_PEPPER) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationCodeService.consume(redis, 'ac1.invalid.short', pepper=_PEPPER) + assert raised.value.error == 'invalid_grant' + assert not redis.values + + +@pytest.mark.asyncio +async def test_expired_code_is_invalid_grant() -> None: + """验证 Redis TTL 到期后返回统一 invalid_grant。""" + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, _payload(), ttl_seconds=1, pepper=_PEPPER) + key = OidcRedisKey.authorization_code(code.split('.')[1]) + redis.values[key] = (redis.values[key][0], 0) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationCodeService.consume(redis, code, pepper=_PEPPER) + assert raised.value.error == 'invalid_grant' + assert await redis.get(OidcRedisKey.authorization_code_consumed(code.split('.')[1])) is None + + +@pytest.mark.asyncio +async def test_invalid_auth_time_and_duplicate_lists_are_rejected() -> None: + """验证授权码时间必须为 UTC,Scope/Resource 列表不得重复。""" + redis = FakeRedis() + with pytest.raises(ValueError): + await AuthorizationCodeService.issue(redis, _payload(authTime='2026-08-24T04:00:00'), pepper=_PEPPER) + with pytest.raises(ValueError): + await AuthorizationCodeService.issue(redis, _payload(scopes=['openid', 'openid']), pepper=_PEPPER) + with pytest.raises(ValueError): + await AuthorizationCodeService.issue(redis, _payload(resources=['api', 'api']), pepper=_PEPPER) + + +@pytest.mark.asyncio +async def test_consumed_corrupt_payload_is_invalid_grant() -> None: + """验证 Lua 已删除但载荷校验失败仍返回统一 OAuth 错误。""" + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, _payload(), pepper=_PEPPER) + key = OidcRedisKey.authorization_code(code.split('.')[1]) + value, expiry = redis.values[key] + redis.values[key] = (value.replace('"authTime":"2026-08-24T04:00:00+00:00"', '"authTime":"bad"'), expiry) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationCodeService.consume(redis, code, pepper=_PEPPER) + assert raised.value.error == 'invalid_grant' diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_authorization_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_authorization_service.py new file mode 100644 index 000000000..2137004de --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_authorization_service.py @@ -0,0 +1,371 @@ +from datetime import datetime, timedelta, timezone + +import pytest +from sqlalchemy.ext.asyncio import AsyncSession + +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_identity.entity.do.oauth_client_do import SysOAuthClient, SysOAuthClientUri +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.protocol_vo import AuthorizeRequest +from module_identity.service.authorization_service import AuthorizationService + +_CLIENT_PK = 7001 +_RESOURCE_PK = 7101 +_SECOND_RESOURCE_PK = 7102 +_REDIRECT_URI = 'https://portal.example/callback?tenant=one' +_CHALLENGE = 'A' * 43 + + +def _request(**overrides: object) -> AuthorizeRequest: + """构造最小有效授权请求。""" + values: dict[str, object] = { + 'response_type': 'code', + 'client_id': 'authorization-client', + 'redirect_uri': _REDIRECT_URI, + 'scope': 'openid profile', + 'nonce': 'nonce-value', + 'state': 'opaque-state', + 'code_challenge': _CHALLENGE, + 'code_challenge_method': 'S256', + } + values.update(overrides) + return AuthorizeRequest(**values) + + +async def _seed_authorization_data(db: AsyncSession, *, second_resource_scope: bool = False) -> None: + """写入授权服务测试所需的 Client、URI、Scope 和 Resource。""" + client = SysOAuthClient( + client_pk=_CLIENT_PK, + client_id='authorization-client', + client_name='Authorization Client', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code'], + response_types=['code'], + require_pkce=1, + require_consent=1, + trusted_client=1, + policy_version=3, + status='0', + ) + resource = SysOAuthResource( + resource_pk=_RESOURCE_PK, + resource_id='portal-api', + resource_name='Portal API', + audience='https://api.example', + allowed_claims=[], + create_by='test', + update_by='test', + status='0', + ) + scopes = [ + SysOAuthScope( + scope_pk=7201, + scope_code='openid', + scope_name='OpenID', + scope_type='identity', + claims=['sub'], + consent_required=1, + status='0', + create_by='test', + update_by='test', + ), + SysOAuthScope( + scope_pk=7202, + scope_code='profile', + scope_name='Profile', + scope_type='identity', + claims=['name'], + consent_required=1, + status='0', + create_by='test', + update_by='test', + ), + SysOAuthScope( + scope_pk=7203, + scope_code='server-required', + scope_name='Server Required', + scope_type='identity', + claims=[], + consent_required=0, + status='0', + create_by='test', + update_by='test', + ), + SysOAuthScope( + scope_pk=7204, + scope_code='portal.read', + scope_name='Portal Read', + scope_type='resource', + resource_pk=_RESOURCE_PK, + claims=['sub'], + consent_required=1, + status='0', + create_by='test', + update_by='test', + ), + ] + bindings = [ + SysOAuthClientScope( + client_pk=_CLIENT_PK, + scope_pk=scope.scope_pk, + pre_authorized=scope.scope_code == 'openid', + ) + for scope in scopes + ] + db.add_all( + [ + client, + SysOAuthClientUri(client_pk=_CLIENT_PK, uri_type='redirect', uri=_REDIRECT_URI, status='0'), + resource, + SysOAuthClientResource(client_pk=_CLIENT_PK, resource_pk=_RESOURCE_PK, is_default=0), + *scopes, + *bindings, + ] + ) + if second_resource_scope: + db.add( + SysOAuthScope( + scope_pk=7205, + scope_code='other.read', + scope_name='Other Read', + scope_type='resource', + resource_pk=_SECOND_RESOURCE_PK, + claims=[], + consent_required=1, + status='0', + create_by='test', + update_by='test', + ) + ) + db.add(SysOAuthClientScope(client_pk=_CLIENT_PK, scope_pk=7205)) + await db.flush() + + +@pytest.fixture +def oidc_enabled(monkeypatch: pytest.MonkeyPatch) -> None: + """为服务测试开启安全的 OIDC 配置。""" + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_pkce_methods', 'S256') + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com') + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +@pytest.mark.parametrize( + 'redirect_uri', + [ + 'https://PORTAL.example/callback?tenant=one', + 'https://portal.example/callback?tenant=two', + 'https://portal.example/Callback?tenant=one', + ], +) +async def test_redirect_uri_must_match_case_path_and_query_exactly( + data_session: AsyncSession, redirect_uri: str +) -> None: + """验证回调地址任一大小写、路径或查询差异都不能进入重定向错误路径。""" + await _seed_authorization_data(data_session) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationService.validate_request(data_session, _request(redirect_uri=redirect_uri)) + assert raised.value.can_redirect is False + assert raised.value.redirect_uri is None + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_scope_and_resource_policy_errors_are_redirectable_only_after_uri_validation( + data_session: AsyncSession, +) -> None: + """验证 Scope 越权和 Resource 归属错误带已验证 Redirect。""" + await _seed_authorization_data(data_session) + with pytest.raises(OAuthProtocolException) as scope_error: + await AuthorizationService.validate_request(data_session, _request(scope='openid admin')) + assert scope_error.value.error == 'invalid_scope' + assert scope_error.value.can_redirect is True + assert scope_error.value.state == 'opaque-state' + + with pytest.raises(OAuthProtocolException) as resource_error: + await AuthorizationService.validate_request( + data_session, _request(scope='openid portal.read', resource='https://other.example') + ) + assert resource_error.value.error == 'invalid_target' + assert resource_error.value.can_redirect is True + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_plain_pkce_method_is_rejected_after_redirect_validation(data_session: AsyncSession) -> None: + """验证 43 位 challenge 也不能绕过 S256 方法校验。""" + await _seed_authorization_data(data_session) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationService.validate_request(data_session, _request(code_challenge_method='plain')) + assert raised.value.error == 'invalid_request' + assert raised.value.can_redirect is True + assert raised.value.redirect_uri == _REDIRECT_URI + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_unsupported_response_type_is_a_safe_redirect_error(data_session: AsyncSession) -> None: + """验证不支持的 response_type 只在 Redirect 已验证后重定向。""" + await _seed_authorization_data(data_session) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationService.validate_request(data_session, _request(response_type='token')) + assert raised.value.error == 'unsupported_response_type' + assert raised.value.can_redirect is True + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_resource_scopes_cannot_cross_resource_ownership(data_session: AsyncSession) -> None: + """验证多个 Resource Scope 不能借一个 audience 混合授权。""" + await _seed_authorization_data(data_session, second_resource_scope=True) + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationService.validate_request( + data_session, _request(scope='openid portal.read other.read', resource='https://api.example') + ) + assert raised.value.error == 'invalid_target' + assert raised.value.can_redirect is True + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_resource_scope_without_resource_uses_only_unique_default_binding(data_session: AsyncSession) -> None: + """验证未传 resource 时只能采用唯一显式默认 Resource。""" + await _seed_authorization_data(data_session) + with pytest.raises(OAuthProtocolException) as no_default: + await AuthorizationService.validate_request(data_session, _request(scope='openid portal.read')) + assert no_default.value.error == 'invalid_target' + + binding = await data_session.get(SysOAuthClientResource, (_CLIENT_PK, _RESOURCE_PK)) + assert binding is not None + binding.is_default = 1 + await data_session.flush() + context = await AuthorizationService.validate_request(data_session, _request(scope='openid portal.read')) + assert context.resources == ('https://api.example',) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_multiple_default_resources_fail_closed(data_session: AsyncSession) -> None: + """验证多个默认 Resource 时拒绝隐式 audience 选择。""" + await _seed_authorization_data(data_session) + second = SysOAuthResource( + resource_pk=_SECOND_RESOURCE_PK, + resource_id='other-api', + resource_name='Other API', + audience='https://other-api.example', + allowed_claims=[], + create_by='test', + update_by='test', + status='0', + ) + data_session.add_all( + [ + second, + SysOAuthClientResource(client_pk=_CLIENT_PK, resource_pk=_SECOND_RESOURCE_PK, is_default=1), + ] + ) + binding = await data_session.get(SysOAuthClientResource, (_CLIENT_PK, _RESOURCE_PK)) + assert binding is not None + binding.is_default = 1 + await data_session.flush() + with pytest.raises(OAuthProtocolException) as raised: + await AuthorizationService.validate_request(data_session, _request(scope='openid portal.read')) + assert raised.value.error == 'invalid_target' + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_prompt_none_preserves_context_and_requires_valid_grant(data_session: AsyncSession) -> None: + """验证 prompt=none 不扩大静默权限,必须由有效 Grant 才能跳过同意。""" + await _seed_authorization_data(data_session) + context = await AuthorizationService.validate_request(data_session, _request(prompt='none')) + assert context.prompt == 'none' + assert context.requires_consent is True + assert AuthorizationService.consent_is_satisfied(context, None) is False + + now = datetime.now(timezone.utc) + grant = SysOAuthGrant( + grant_id='authorization-grant', + user_id=7001, + subject_id='subject-7001', + client_pk=_CLIENT_PK, + granted_scopes=['openid', 'profile'], + remembered_scopes=['openid', 'profile'], + remembered_resources=[], + granted_resources=[], + client_policy_version=3, + status='active', + consented_at=now, + expires_at=now + timedelta(minutes=5), + ) + assert AuthorizationService.consent_is_satisfied(context, grant) is True + grant.client_policy_version = 2 + assert AuthorizationService.consent_is_satisfied(context, grant) is False + grant.client_policy_version = 3 + grant.expires_at = now - timedelta(seconds=1) + assert AuthorizationService.consent_is_satisfied(context, grant) is False + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_require_consent_flag_and_prompt_override_do_not_use_trusted_client_as_expansion( + data_session: AsyncSession, +) -> None: + """验证 require_consent、trusted_client 和 prompt=consent 的三种组合。""" + await _seed_authorization_data(data_session) + client = await data_session.get(SysOAuthClient, _CLIENT_PK) + assert client is not None + + client.require_consent = 0 + context = await AuthorizationService.validate_request(data_session, _request()) + assert context.client.trusted_client is True + assert context.requires_consent is False + assert AuthorizationService.consent_is_satisfied(context, None) is True + + client.require_consent = 1 + context = await AuthorizationService.validate_request(data_session, _request()) + assert context.requires_consent is True + + client.require_consent = 0 + context = await AuthorizationService.validate_request(data_session, _request(prompt='login consent')) + assert context.requires_consent is True + assert AuthorizationService.consent_is_satisfied(context, None) is False + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_unrequested_server_required_scope_does_not_block_consent(data_session: AsyncSession) -> None: + """验证未请求的服务端必需 Scope 不会混入本次授权的不可取消集合。""" + await _seed_authorization_data(data_session) + context = await AuthorizationService.validate_request(data_session, _request(scope='openid profile')) + assert context.required_scopes == frozenset({'openid'}) + assert AuthorizationService.consent_is_satisfied(context, None) is False + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('oidc_enabled') +async def test_context_payload_excludes_state_and_nonce(data_session: AsyncSession) -> None: + """验证交互页面载荷不泄漏 state、nonce 和未定义参数。""" + await _seed_authorization_data(data_session) + context = await AuthorizationService.validate_request(data_session, _request(login_hint='not-forwarded')) + payload = context.to_interaction_payload('interaction-1') + internal_payload = context.to_internal_payload('interaction-1') + assert 'state' not in payload + assert 'nonce' not in payload + assert 'login_hint' not in payload + assert 'clientPk' not in payload + assert 'maxAge' not in payload + assert 'redirectUri' not in payload + assert internal_payload['redirectUri'] == _REDIRECT_URI + assert internal_payload['state'] == 'opaque-state' + assert internal_payload['nonce'] == 'nonce-value' + assert internal_payload['codeChallenge'] == _CHALLENGE diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_claim_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_claim_service.py new file mode 100644 index 000000000..24745e5c9 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_claim_service.py @@ -0,0 +1,79 @@ +from datetime import datetime, timezone + +import pytest + +from exceptions.exception import OAuthProtocolException +from module_identity.service.identity_service import ClaimService + + +def test_claims_require_requested_client_and_resource_intersection() -> None: + """只有三重交集中的 Claim 才能生成。""" + allowed = ClaimService.effective_claims( + 'openid profile email phone', + {'openid': True, 'profile': ['name'], 'email': ['email'], 'phone': ['phone_number']}, + ['sub', 'name', 'email'], + ) + assert allowed == {'sub', 'name', 'email'} + + +def test_claim_builder_does_not_leak_local_identity_or_unapproved_contact() -> None: + """本地 user_id、密码以及未获授权的手机号不得进入 Claims。""" + user = { + 'user_id': 99, + 'password': 'plain-or-hash', + 'user_name': 'alice', + 'nick_name': 'Alice', + 'email': 'alice@example.com', + 'phonenumber': '13800000000', + 'subject_id': 'stable-sub', + 'roles': ['admin'], + } + claims = ClaimService.build_claims( + user, + {'openid', 'profile', 'email', 'phone', 'roles'}, + { + 'openid': True, + 'profile': ['name'], + 'email': ['email'], + 'phone': ['phone_number'], + 'roles': {'claims': ['roles'], 'allowed_role_keys': ['admin', 'admin_key']}, + }, + {'allowed_claims': ['sub', 'name', 'email', 'roles']}, + ) + assert claims == {'sub': 'stable-sub', 'name': 'Alice', 'email': 'alice@example.com', 'roles': ['admin']} + assert 'user_id' not in claims + assert 'password' not in claims + assert 'phone_number' not in claims + + +@pytest.mark.parametrize( + 'user, policy, resource', + [ + ({'subject_id': None}, {'openid': True}, ['name']), + ({'subject_id': 'stable-sub'}, {'openid': ['name']}, ['name']), + ], +) +def test_openid_fails_closed_when_subject_or_sub_policy_is_missing( + user: dict[str, object], policy: dict[str, object], resource: list[str] +) -> None: + """openid 请求缺少稳定 Subject 或 sub 策略时不得静默生成 Claims。""" + with pytest.raises(OAuthProtocolException): + ClaimService.build_claims(user, {'openid'}, policy, resource) + + +def test_updated_at_is_numeric_date_and_roles_are_role_keys() -> None: + """本地更新时间输出 NumericDate,角色查询语句使用 role_key。""" + user = {'subject_id': 'stable-sub', 'update_time': datetime(1970, 1, 1, 0, 0, 1, tzinfo=timezone.utc)} + claims = ClaimService.build_claims( + user, + {'openid', 'profile', 'roles'}, + { + 'openid': True, + 'profile': ['updated_at'], + 'roles': {'claims': ['roles'], 'allowed_role_keys': ['admin', 'admin_key']}, + }, + ['sub', 'updated_at', 'roles'], + roles=['admin_key'], + ) + assert claims['updated_at'] == int(datetime(1970, 1, 1, 0, 0, 1, tzinfo=timezone.utc).timestamp()) + assert claims['roles'] == ['admin_key'] diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_consent_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_consent_service.py new file mode 100644 index 000000000..eb0adb950 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_consent_service.py @@ -0,0 +1,304 @@ +from dataclasses import replace +from datetime import datetime, timedelta, timezone + +import pytest +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker + +from common.constant import OidcAuditEvent +from exceptions.exception import OAuthProtocolException +from module_admin.entity.do.user_do import SysUser +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant +from module_identity.service.authorization_service import ( + AuthorizationContext, + ClientSnapshot, + ScopeSnapshot, +) +from module_identity.service.consent_service import ConsentResult, ConsentService + + +def _context(*, prompt: str | None = None, trusted_client: int = 0) -> AuthorizationContext: + """构造同意服务所需的已验证上下文。""" + client = ClientSnapshot( + client_pk=8001, + client_id='consent-client', + grant_types=('authorization_code',), + response_types=('code',), + policy_version=5, + require_pkce=True, + require_consent=True, + trusted_client=bool(trusted_client), + ) + scopes = ( + ScopeSnapshot( + scope_pk=8101, + scope_code='openid', + scope_type='identity', + resource_pk=None, + consent_required=True, + ), + ScopeSnapshot( + scope_pk=8102, + scope_code='profile', + scope_type='identity', + resource_pk=None, + consent_required=True, + ), + ScopeSnapshot( + scope_pk=8103, + scope_code='server-required', + scope_type='identity', + resource_pk=None, + consent_required=False, + ), + ) + return AuthorizationContext( + client=client, + redirect_uri='https://portal.example/callback', + scopes=('openid', 'profile', 'server-required'), + scope_models=scopes, + pre_authorized_scopes=frozenset(), + resource=None, + state='opaque-state', + nonce='opaque-nonce', + code_challenge='A' * 43, + code_challenge_method='S256', + prompt=prompt, + max_age=None, + required_scopes=frozenset({'openid', 'server-required'}), + ) + + +def test_submission_cannot_expand_scope_or_cancel_required_scope() -> None: + """验证提交范围只能是原请求子集且必需 Scope 不可取消。""" + context = _context() + with pytest.raises(OAuthProtocolException) as expanded: + ConsentService.validate_submission(context, True, ['openid', 'profile', 'admin']) + assert expanded.value.error == 'invalid_scope' + assert expanded.value.can_redirect is True + + with pytest.raises(OAuthProtocolException) as cancelled: + ConsentService.validate_submission(context, True, ['profile']) + assert cancelled.value.error == 'invalid_scope' + assert cancelled.value.can_redirect is True + + +def test_denial_is_a_verified_redirect_and_trusted_client_does_not_skip_consent() -> None: + """验证拒绝错误可安全重定向,trusted_client 不会自动预授权。""" + context = _context(trusted_client=1) + with pytest.raises(OAuthProtocolException) as denied: + ConsentService.validate_submission(context, False, ['openid', 'profile', 'server-required']) + assert denied.value.error == 'access_denied' + assert denied.value.can_redirect is True + assert context.requires_consent is True + + +def test_prompt_consent_disables_grant_based_skip() -> None: + """验证 prompt=consent 强制显示同意页,即使已有 Grant 也不能静默跳过。""" + context = _context(prompt='consent') + now = datetime.now(timezone.utc) + grant = SysOAuthGrant( + grant_id='consent-grant', + user_id=8001, + subject_id='subject-8001', + client_pk=8001, + granted_scopes=['openid', 'profile', 'server-required'], + granted_resources=[], + client_policy_version=5, + status='active', + consented_at=now, + expires_at=now + timedelta(minutes=5), + ) + assert ConsentService.consent_is_satisfied(context, grant) is False + + +@pytest.mark.asyncio +@pytest.mark.parametrize('remembered', [True, False]) +async def test_remember_consent_merges_grant_without_commit(data_session: AsyncSession, remembered: bool) -> None: + """验证 rememberConsent 使用 Grant DAO 合并,提交事务由调用方控制。""" + initial = _context() + context = replace( + initial, + scopes=(*initial.scopes, 'offline_access'), + scope_models=(*initial.scope_models, ScopeSnapshot(8104, 'offline_access', 'identity', None, True)), + ) + data_session.add( + SysUser(user_id=8001, user_name='consent-user', nick_name='Consent User', status='0', del_flag='0') + ) + data_session.add( + SysOAuthClient( + client_pk=8001, + client_id='consent-client', + client_name='Consent Client', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code'], + response_types=['code'], + trusted_client=1, + policy_version=5, + status='0', + ) + ) + await data_session.flush() + + result = await ConsentService.submit_consent( + data_session, + context, + approved=True, + scopes=list(context.scopes), + remember_consent=remembered, + user_id=8001, + subject_id='subject-8001', + ) + assert result.approved is True + assert result.grant is not None + assert result.grant.granted_scopes == list(context.scopes) + assert ConsentService.consent_is_satisfied(context, result.grant) is remembered + assert result.grant.remembered_scopes == (list(context.scopes) if remembered else []) + assert data_session.in_transaction() + await data_session.flush() + + # 后续记住更小的权限范围时,不得将此前仅供离线续期的权限一并记住 + smaller = replace(context, scopes=('openid', 'server-required')) + updated = await ConsentService.submit_consent( + data_session, + smaller, + approved=True, + scopes=list(smaller.scopes), + remember_consent=True, + user_id=8001, + subject_id='subject-8001', + ) + assert updated.grant.grant_id == result.grant.grant_id + assert set(updated.grant.granted_scopes) == set(context.scopes) + assert ConsentService.consent_is_satisfied(context, updated.grant) is remembered + assert ConsentService.consent_is_satisfied(smaller, updated.grant) is True + + +@pytest.mark.asyncio +async def test_compensation_does_not_restore_concurrently_revoked_grant(data_session: AsyncSession) -> None: + """补偿 CAS 不得覆盖管理员在提交后执行的 Grant 撤销。""" + now = datetime.now(timezone.utc) + user = SysUser(user_id=8002, user_name='race-user', nick_name='Race User', status='0', del_flag='0') + client = SysOAuthClient( + client_pk=8002, + client_id='race-client', + client_name='Race Client', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code'], + response_types=['code'], + policy_version=5, + status='0', + ) + grant = SysOAuthGrant( + grant_id='race-grant', + user_id=8002, + subject_id='subject-8002', + client_pk=8002, + granted_scopes=['openid'], + granted_resources=[], + client_policy_version=5, + status='active', + consented_at=now, + ) + data_session.add_all([user, client, grant]) + await data_session.commit() + + previous = OAuthGrantDao.snapshot(grant) + grant.granted_scopes = ['openid', 'profile'] + grant.consented_at = now + timedelta(seconds=1) + await data_session.flush() + persisted = OAuthGrantDao.snapshot(grant) + await data_session.commit() + + session_factory = async_sessionmaker(data_session.bind, expire_on_commit=False) + async with session_factory() as concurrent_session: + assert await OAuthGrantDao.revoke(concurrent_session, grant.grant_id, reason='administrator') is True + await concurrent_session.commit() + await ConsentService.compensate_persisted_grant( + data_session, + ConsentResult( + approved=True, + scopes=('openid', 'profile'), + grant=grant, + previous_grant=previous, + persisted_grant=persisted, + ), + ) + await data_session.refresh(grant) + assert grant.status == 'revoked' + assert grant.revoke_reason == 'administrator' + audit = ( + await data_session.execute( + select(SysOAuthAuditLog).where(SysOAuthAuditLog.failure_code == 'consent_compensation_conflict') + ) + ).scalar_one() + assert audit.event_type == OidcAuditEvent.CONSENT_GRANTED + assert audit.risk_level == 'high' + + +@pytest.mark.asyncio +async def test_compensation_does_not_revoke_concurrently_updated_new_grant(data_session: AsyncSession) -> None: + """新建 Grant 的撤销补偿不得覆盖另一会话的更新,并记录高风险审计。""" + now = datetime.now(timezone.utc) + user = SysUser(user_id=8003, user_name='new-race-user', nick_name='New Race User', status='0', del_flag='0') + client = SysOAuthClient( + client_pk=8003, + client_id='new-race-client', + client_name='New Race Client', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code'], + response_types=['code'], + policy_version=5, + status='0', + ) + grant = SysOAuthGrant( + grant_id='new-race-grant', + user_id=8003, + subject_id='subject-8003', + client_pk=8003, + granted_scopes=['openid'], + granted_resources=[], + client_policy_version=5, + status='active', + consented_at=now, + ) + data_session.add_all([user, client, grant]) + await data_session.commit() + persisted = OAuthGrantDao.snapshot(grant) + + session_factory = async_sessionmaker(data_session.bind, expire_on_commit=False) + async with session_factory() as concurrent_session: + await concurrent_session.execute( + SysOAuthGrant.__table__.update() + .where(SysOAuthGrant.grant_id == grant.grant_id) + .values(granted_scopes=['openid', 'admin'], subject_id='concurrent-subject') + ) + await concurrent_session.commit() + + await ConsentService.compensate_persisted_grant( + data_session, + ConsentResult( + approved=True, + scopes=('openid',), + grant=grant, + persisted_grant=persisted, + ), + ) + await data_session.rollback() + await data_session.refresh(grant) + assert grant.status == 'active' + assert grant.granted_scopes == ['openid', 'admin'] + assert grant.subject_id == 'concurrent-subject' + audit = ( + await data_session.execute( + select(SysOAuthAuditLog).where(SysOAuthAuditLog.failure_code == 'consent_compensation_conflict') + ) + ).scalar_one() + assert audit.event_type == OidcAuditEvent.CONSENT_GRANTED + assert audit.risk_level == 'high' diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_credential_authentication_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_credential_authentication_service.py new file mode 100644 index 000000000..001c4933b --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_credential_authentication_service.py @@ -0,0 +1,391 @@ +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest + +from common.enums import RedisInitKeyConfig +from config.env import OidcConfig +from exceptions.exception import LoginException +from module_admin.entity.vo.login_vo import UserLogin +from module_admin.service.login_service import LoginService +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.identity_service import ( + CredentialAuthenticationError, + CredentialAuthenticationService, +) +from utils.oidc_util import OidcUtil +from utils.pwd_util import PwdUtil + +_OIDC_FAILURE_SCRIPT_KEY_COUNT = 2 +_OIDC_FAILURE_SCRIPT_ARGUMENT_COUNT = 4 + + +class _Redis: + """记录认证服务访问的 Redis 最小异步替身。""" + + def __init__(self, values: dict[str, object] | None = None, eval_values: list[object] | None = None) -> None: + self.values = values or {} + self.eval_values = list(eval_values or []) + self.get_calls: list[str] = [] + self.set_calls: list[tuple[str, object, object]] = [] + self.delete_calls: list[str] = [] + self.eval_calls: list[tuple[str, int, tuple[object, ...]]] = [] + + async def get(self, key: str) -> object: + self.get_calls.append(key) + return self.values.get(key) + + async def set(self, key: str, value: object, *, ex: object) -> None: + self.values[key] = value + self.set_calls.append((key, value, ex)) + + async def delete(self, key: str) -> None: + self.values.pop(key, None) + self.delete_calls.append(key) + + async def eval(self, script: str, count: int, *keys_and_args: object) -> object: + self.eval_calls.append((script, count, keys_and_args)) + return self.eval_values.pop(0) if self.eval_values else None + + +class _AtomicFailureRedis(_Redis): + """按 OIDC 错误 Lua 脚本语义执行 INCR、锁定和锁前检查。""" + + async def eval(self, script: str, count: int, *keys_and_args: object) -> object: + self.eval_calls.append((script, count, keys_and_args)) + if "redis.call('incr'" not in script: + return self.eval_values.pop(0) if self.eval_values else None + failure_key, lock_key, threshold, _ttl = keys_and_args + if self.values.get(lock_key): + return -1 + current = int(self.values.get(failure_key, 0)) + 1 + if current > int(threshold): + self.values.pop(failure_key, None) + self.values[lock_key] = '1' + return -1 + self.values[failure_key] = current + return current + + +def _request(redis: _Redis, *, host: str = '127.0.0.1') -> SimpleNamespace: + return SimpleNamespace( + app=SimpleNamespace(state=SimpleNamespace(redis=redis)), + client=SimpleNamespace(host=host), + headers={}, + ) + + +def _user(user_name: str = 'alice', *, status: str = '0', update_date: object = None) -> SimpleNamespace: + return SimpleNamespace( + user_id=7, + user_name=user_name, + password='stored-hash', + status=status, + pwd_update_date=update_date, + ) + + +@pytest.mark.asyncio +async def test_oidc_unknown_user_and_wrong_password_are_equivalent_and_dummy_verify_runs() -> None: + redis = _Redis() + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=None), + ), + patch.object(PwdUtil, 'verify_password', return_value=False) as verify, + pytest.raises(CredentialAuthenticationError) as unknown, + ): + await CredentialAuthenticationService.authenticate_oidc(redis, object(), user_name='alice', password='wrong') + verify.assert_called_once_with('wrong', '$2b$12$ySHJfAWxzh49cIc7M5L21e5GlPyA7QhE2GkLn9XuUTqmKIRqhWIja') + redis_keys = [item[0] for item in redis.set_calls] + assert unknown.value.reason == 'invalid_credentials' + assert all('alice' not in key for key in redis_keys) + + redis = _Redis() + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(_user(), None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=False), + pytest.raises(CredentialAuthenticationError) as wrong, + ): + await CredentialAuthenticationService.authenticate_oidc(redis, object(), user_name='alice', password='wrong') + assert wrong.value.reason == unknown.value.reason + assert str(wrong.value) == str(unknown.value) == '账号或密码错误' + + +@pytest.mark.asyncio +async def test_oidc_captcha_is_atomically_consumed_and_success_does_not_write_legacy_token_key() -> None: + redis = _Redis(eval_values=['1234', None]) + user = _user() + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(user, None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=True), + ): + result = await CredentialAuthenticationService.authenticate_oidc( + redis, + object(), + user_name='alice', + password='correct', + code='1234', + uuid='captcha-id', + captcha_enabled=True, + remember_me=True, + ) + assert result.amr == ('pwd', 'captcha') + assert result.remember_me is True + with pytest.raises(CredentialAuthenticationError) as replay: + await CredentialAuthenticationService.authenticate_oidc( + redis, + object(), + user_name='alice', + password='correct', + code='1234', + uuid='captcha-id', + captcha_enabled=True, + ) + assert replay.value.reason == 'captcha_missing' + assert redis.eval_calls and redis.eval_calls[0][1] == 1 + assert not any(key.startswith(f'{RedisInitKeyConfig.ACCESS_TOKEN.key}:') for key, _, _ in redis.set_calls) + + +@pytest.mark.asyncio +async def test_oidc_disabled_user_is_rejected() -> None: + status, expected_reason = '1', 'user_disabled' + redis = _Redis() + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(_user(status=status), None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=True), + pytest.raises(CredentialAuthenticationError) as exc_info, + ): + await CredentialAuthenticationService.authenticate_oidc(redis, object(), user_name='alice', password='correct') + assert exc_info.value.reason == expected_reason + + +@pytest.mark.asyncio +async def test_oidc_initial_and_expired_password_flags_handle_naive_and_aware_dates() -> None: + redis = _Redis( + { + 'sys_config:sys.account.initPasswordModify': '1', + 'sys_config:sys.account.passwordValidateDays': '30', + } + ) + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(_user(update_date=None), None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=True), + ): + result = await CredentialAuthenticationService.authenticate_oidc( + redis, object(), user_name='alice', password='correct' + ) + assert result.password_change_required is True + assert result.password_change_reason == 'initial_password' + + +@pytest.mark.asyncio +async def test_oidc_expired_aware_password_is_reported_without_datetime_error() -> None: + redis = _Redis({'sys_config:sys.account.passwordValidateDays': '30'}) + old_date = datetime.now(timezone.utc) - timedelta(days=31) + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(_user(update_date=old_date), None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=True), + ): + result = await CredentialAuthenticationService.authenticate_oidc( + redis, object(), user_name='alice', password='correct' + ) + assert result.password_change_required is True + assert result.password_change_reason == 'password_expired' + + +@pytest.mark.asyncio +async def test_oidc_hashed_error_state_locks_without_plain_username() -> None: + redis = _Redis(eval_values=[-1]) + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(_user(), None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=False), + pytest.raises(CredentialAuthenticationError) as exc_info, + ): + digest = OidcUtil.hash_sensitive_identifier('alice', 'p' * 32) + redis.values[OidcRedisKey.login_user_rate_limit(digest)] = 5 + await CredentialAuthenticationService.authenticate_oidc(redis, object(), user_name='alice', password='wrong') + assert exc_info.value.reason == 'account_locked' + assert all('alice' not in key for key, _, _ in redis.set_calls) + + +@pytest.mark.asyncio +async def test_oidc_error_counter_uses_atomic_incr_expire_and_lock_script() -> None: + redis = _Redis(eval_values=[1]) + with ( + patch.object(OidcConfig, 'oidc_token_hash_pepper', 'p' * 32), + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(_user(), None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=False), + pytest.raises(CredentialAuthenticationError), + ): + await CredentialAuthenticationService.authenticate_oidc(redis, object(), user_name='alice', password='wrong') + script, key_count, arguments = redis.eval_calls[-1] + assert "redis.call('incr'" in script + assert "redis.call('expire'" in script + assert key_count == _OIDC_FAILURE_SCRIPT_KEY_COUNT + assert len(arguments) == _OIDC_FAILURE_SCRIPT_ARGUMENT_COUNT + assert all('alice' not in str(value) for value in arguments[:2]) + + +@pytest.mark.asyncio +async def test_oidc_failure_script_stays_locked_after_atomic_lock_transition() -> None: + """建锁后的并发等效第二次脚本调用不得重建错误计数。""" + redis = _AtomicFailureRedis({}) + failure_key = 'oidc:rate_limit:login:user:' + 'a' * 64 + lock_key = f'{failure_key}:lock' + redis.values[failure_key] = 5 + with pytest.raises(CredentialAuthenticationError) as first: + await CredentialAuthenticationService._record_oidc_password_error(redis, failure_key, lock_key) + with pytest.raises(CredentialAuthenticationError) as second: + await CredentialAuthenticationService._record_oidc_password_error(redis, failure_key, lock_key) + assert first.value.reason == 'account_locked' + assert second.value.reason == 'account_locked' + assert failure_key not in redis.values + assert redis.values[lock_key] == '1' + assert "exists', KEYS[2]" in redis.eval_calls[0][0] + + +@pytest.mark.asyncio +async def test_legacy_success_and_wrong_password_keep_legacy_keys_and_messages() -> None: + user = _user() + redis = _Redis() + with ( + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(user, None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=True), + ): + result = await LoginService.authenticate_user( + _request(redis), object(), UserLogin(userName='alice', password='correct', captchaEnabled=False) + ) + assert result[0] is user + assert f'{RedisInitKeyConfig.PASSWORD_ERROR_COUNT.key}:alice' in redis.delete_calls + assert not any(key.startswith('oidc:') for key, _, _ in redis.set_calls) + + redis = _Redis() + with ( + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=(user, None)), + ), + patch.object(PwdUtil, 'verify_password', return_value=False), + pytest.raises(LoginException) as exc_info, + ): + await LoginService.authenticate_user( + _request(redis), object(), UserLogin(userName='alice', password='wrong', captchaEnabled=False) + ) + assert exc_info.value.message == '密码错误' + assert any(key == 'password_error_count:alice' for key, _, _ in redis.set_calls) + + +@pytest.mark.asyncio +async def test_legacy_unknown_user_keeps_message_and_does_not_create_error_key() -> None: + redis = _Redis() + with ( + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(return_value=None), + ), + pytest.raises(LoginException) as exc_info, + ): + await LoginService.authenticate_user( + _request(redis), object(), UserLogin(userName='alice', password='wrong', captchaEnabled=False) + ) + assert exc_info.value.message == '用户不存在' + assert not any(key.startswith('password_error_count:') for key, _, _ in redis.set_calls) + + +@pytest.mark.asyncio +async def test_legacy_lock_captcha_and_blacklist_keep_original_messages_and_keys() -> None: + lock_redis = _Redis({'account_lock:alice': 'alice'}) + with ( + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(), + ) as login_by_account, + pytest.raises(LoginException) as locked, + ): + await LoginService.authenticate_user( + _request(lock_redis), object(), UserLogin(userName='alice', password='wrong', captchaEnabled=False) + ) + assert locked.value.message == '账号已锁定,请稍后再试' + login_by_account.assert_not_awaited() + assert 'account_lock:alice' in lock_redis.get_calls + + captcha_redis = _Redis() + with ( + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(), + ) as login_by_account, + pytest.raises(LoginException) as captcha_error, + ): + await LoginService.authenticate_user( + _request(captcha_redis), + object(), + UserLogin(userName='alice', password='wrong', code='1', uuid='u', captchaEnabled=True), + ) + assert captcha_error.value.message == '验证码已失效' + login_by_account.assert_not_awaited() + assert 'captcha_codes:u' in captcha_redis.get_calls + + ip_redis = _Redis({'sys_config:sys.login.blackIPList': '127.0.0.1'}) + with ( + patch( + 'module_identity.service.identity_service.login_by_account', + new=AsyncMock(), + ) as login_by_account, + pytest.raises(LoginException) as ip_error, + ): + await LoginService.authenticate_user( + _request(ip_redis), object(), UserLogin(userName='alice', password='wrong', captchaEnabled=False) + ) + assert ip_error.value.message == '当前IP禁止登录' + login_by_account.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_legacy_service_delegates_and_maps_original_login_exception() -> None: + login_user = UserLogin(userName='alice', password='wrong', captchaEnabled=False) + request = _request(_Redis()) + with ( + patch.object( + CredentialAuthenticationService, + 'authenticate_legacy', + new=AsyncMock(side_effect=CredentialAuthenticationError('invalid_credentials', '密码错误')), + ), + pytest.raises(LoginException) as exc_info, + ): + await LoginService.authenticate_user(request, object(), login_user) + assert exc_info.value.message == '密码错误' diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_identity_security_event_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_identity_security_event_service.py new file mode 100644 index 000000000..d89964110 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_identity_security_event_service.py @@ -0,0 +1,195 @@ +from datetime import datetime, timedelta, timezone + +import pytest +import pytest_asyncio +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from config.database import Base +from module_admin.entity.do.user_do import SysUser, SysUserRole +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.do.oauth_grant_do import SysOAuthRefreshToken, SysSsoSession +from module_identity.service.identity_service import ( + IdentitySecurityEventError, + IdentitySecurityEventService, +) + +_NOW = datetime(2026, 8, 24, 8, 0, tzinfo=timezone.utc) +_INITIAL_VERSION = 3 +_PASSWORD_REFRESH_COUNT = 2 +_CLAIM_USER_ID = 8 +_ROLE_REFRESH_COUNT = 4 + + +@pytest_asyncio.fixture +async def security_session() -> AsyncSession: + """创建身份安全事件涉及表的真实异步数据库会话。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + tables = [ + SysUser.__table__, + SysUserRole.__table__, + SysIdentitySubject.__table__, + SysSsoSession.__table__, + SysOAuthRefreshToken.__table__, + SysOAuthAuditLog.__table__, + ] + async with engine.begin() as connection: + await connection.run_sync(lambda sync_connection: Base.metadata.create_all(sync_connection, tables=tables)) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as session: + yield session + await engine.dispose() + + +async def _seed_user(session: AsyncSession, user_id: int, *, with_subject: bool = True) -> None: + """写入一个用户及其可撤销的 Session、Refresh。""" + session.add(SysUser(user_id=user_id, user_name=f'user-{user_id}', nick_name='User')) + if with_subject: + session.add( + SysIdentitySubject( + user_id=user_id, + subject_id=f'00000000-0000-4000-8000-{user_id:012d}', + auth_version=_INITIAL_VERSION, + create_time=_NOW, + ) + ) + for index in range(2): + sid = f'sid-{user_id}-{index}' + session.add( + SysSsoSession( + sid=sid, + session_secret_hash=str(index) * 64, + user_id=user_id, + subject_id=f'00000000-0000-4000-8000-{user_id:012d}', + auth_version=_INITIAL_VERSION, + auth_time=_NOW, + last_seen_at=_NOW, + idle_expires_at=_NOW + timedelta(hours=1), + absolute_expires_at=_NOW + timedelta(hours=2), + acr='urn:ruoyi:acr:pwd', + amr=['pwd'], + status='active', + create_time=_NOW, + ) + ) + session.add( + SysOAuthRefreshToken( + token_id=f'token-{user_id}-{index}', + token_hash=f'{user_id:02d}{index}' * 16, + family_id=f'family-{user_id}-{index}', + grant_id=f'grant-{user_id}', + user_id=user_id, + subject_id=f'00000000-0000-4000-8000-{user_id:012d}', + auth_version=_INITIAL_VERSION, + client_pk=100, + sid=sid, + scopes=['openid'], + resources=[], + status='active', + issued_at=_NOW, + idle_expires_at=_NOW + timedelta(days=1), + absolute_expires_at=_NOW + timedelta(days=7), + ) + ) + await session.commit() + + +@pytest.mark.asyncio +async def test_password_change_increments_version_and_keeps_only_current_session( + security_session: AsyncSession, +) -> None: + """修改密码应递增版本、撤销全部 Refresh,并可保留当前交互 Session。""" + await _seed_user(security_session, 7) + result = await IdentitySecurityEventService.handle_user_event( + security_session, + 7, + 'password_changed', + actor='admin', + exclude_sid='sid-7-0', + now=_NOW + timedelta(minutes=1), + ) + assert result.revoked_sessions == 1 + assert result.revoked_refresh_tokens == _PASSWORD_REFRESH_COUNT + assert (await security_session.get(SysIdentitySubject, 1)).auth_version == _INITIAL_VERSION + 1 + sessions = (await security_session.execute(select(SysSsoSession).order_by(SysSsoSession.sid))).scalars().all() + assert [row.status for row in sessions] == ['active', 'revoked'] + tokens = ( + (await security_session.execute(select(SysOAuthRefreshToken).order_by(SysOAuthRefreshToken.token_id))) + .scalars() + .all() + ) + assert all(row.status == 'revoked' for row in tokens) + audit = (await security_session.execute(select(SysOAuthAuditLog))).scalar_one() + assert audit.detail['event'] == 'password_changed' + assert audit.detail['actor'] == 'admin' + assert audit.risk_level == 'high' + + +@pytest.mark.asyncio +async def test_claim_change_keeps_sso_by_default_and_transaction_can_rollback( + security_session: AsyncSession, +) -> None: + """角色/部门变化默认保留 SSO,且全部安全写入服从调用方事务。""" + await _seed_user(security_session, _CLAIM_USER_ID) + await IdentitySecurityEventService.handle_user_event( + security_session, + _CLAIM_USER_ID, + 'role_assignment_changed', + now=_NOW + timedelta(minutes=2), + ) + assert all( + row.status == 'active' for row in (await security_session.execute(select(SysSsoSession))).scalars().all() + ) + assert all( + row.status == 'revoked' + for row in (await security_session.execute(select(SysOAuthRefreshToken))).scalars().all() + ) + await security_session.rollback() + subject = await security_session.scalar( + select(SysIdentitySubject).where(SysIdentitySubject.user_id == _CLAIM_USER_ID) + ) + assert subject is not None and subject.auth_version == _INITIAL_VERSION + assert all( + row.status == 'active' for row in (await security_session.execute(select(SysOAuthRefreshToken))).scalars().all() + ) + assert not (await security_session.execute(select(SysOAuthAuditLog))).scalars().all() + + +@pytest.mark.asyncio +async def test_missing_subject_fails_before_revocation(security_session: AsyncSession) -> None: + """主体映射缺失必须 fail closed,且不能先撤销部分凭据。""" + await _seed_user(security_session, 9, with_subject=False) + with pytest.raises(IdentitySecurityEventError, match='缺少用户主体映射'): + await IdentitySecurityEventService.handle_user_event(security_session, 9, 'user_disabled', now=_NOW) + assert all( + row.status == 'active' for row in (await security_session.execute(select(SysSsoSession))).scalars().all() + ) + assert all( + row.status == 'active' for row in (await security_session.execute(select(SysOAuthRefreshToken))).scalars().all() + ) + + +@pytest.mark.asyncio +async def test_role_event_updates_all_current_members(security_session: AsyncSession) -> None: + """角色停用应按成员集合批量更新版本与 Refresh。""" + await _seed_user(security_session, 10) + await _seed_user(security_session, 11) + security_session.add_all([SysUserRole(user_id=10, role_id=5), SysUserRole(user_id=11, role_id=5)]) + await security_session.commit() + result = await IdentitySecurityEventService.handle_role_event( + security_session, + 5, + 'role_disabled', + actor='admin', + now=_NOW + timedelta(minutes=3), + ) + assert result.affected_users == (10, 11) + assert result.revoked_refresh_tokens == _ROLE_REFRESH_COUNT + assert result.revoked_sessions == 0 + versions = ( + (await security_session.execute(select(SysIdentitySubject.auth_version).order_by(SysIdentitySubject.user_id))) + .scalars() + .all() + ) + assert versions == [_INITIAL_VERSION + 1, _INITIAL_VERSION + 1] diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_identity_subject_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_identity_subject_service.py new file mode 100644 index 000000000..dbf0bd3c6 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_identity_subject_service.py @@ -0,0 +1,130 @@ +from pathlib import Path + +import pytest +import pytest_asyncio +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from common.constant import OidcAuditEvent +from config.database import Base +from exceptions.exception import OAuthProtocolException +from module_admin.entity.do.user_do import SysUser +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.service.audit_service import AuditService +from module_identity.service.identity_service import IdentitySubjectService + +_MISSING_SUBJECT_USER_ID = 1101 +_ROLLBACK_USER_ID = 1105 +_REPAIRED_SUBJECT_COUNT = 2 +_NEXT_AUTH_VERSION = 2 + + +@pytest_asyncio.fixture +async def service_session(tmp_path: Path) -> AsyncSession: + """创建只包含主体服务所需表的真实 SQLite 会话。""" + engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "subject-service.sqlite3"}') + tables = [SysUser.__table__, SysIdentitySubject.__table__, SysOAuthAuditLog.__table__] + async with engine.begin() as connection: + await connection.run_sync(lambda sync_connection: Base.metadata.create_all(sync_connection, tables=tables)) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as session: + session.info['service_engine'] = engine + yield session + await engine.dispose() + + +@pytest.mark.asyncio +async def test_missing_subject_fails_closed_and_records_integrity_audit(service_session: AsyncSession) -> None: + """主体缺失时不能签发,并且必须留下高风险完整性事件。""" + service_session.add( + SysUser( + user_id=_MISSING_SUBJECT_USER_ID, + user_name='missing', + nick_name='Missing', + status='0', + del_flag='0', + ) + ) + await service_session.flush() + + with pytest.raises(OAuthProtocolException) as caught: + await IdentitySubjectService.require_by_user_id(service_session, _MISSING_SUBJECT_USER_ID) + + assert caught.value.error == 'server_error' + event = await service_session.execute( + select(SysOAuthAuditLog).where(SysOAuthAuditLog.user_id == _MISSING_SUBJECT_USER_ID) + ) + event = event.scalars().first() + assert event is not None + assert event.event_type == OidcAuditEvent.IDENTITY_SUBJECT_MISSING + assert event.risk_level == 'high' + assert event.result == 'failure' + + +@pytest.mark.asyncio +async def test_repair_single_and_batch_are_idempotent(service_session: AsyncSession) -> None: + """单用户和批量修复重复执行都不会替换 Subject 或增加记录。""" + service_session.add_all( + [ + SysUser(user_id=1102, user_name='one', nick_name='One', status='0', del_flag='0'), + SysUser(user_id=1103, user_name='two', nick_name='Two', status='0', del_flag='0'), + ] + ) + await service_session.flush() + first = await IdentitySubjectService.repair_missing_subject(service_session, 1102) + same = await IdentitySubjectService.repair_missing_subject(service_session, 1102) + assert first.identity_id == same.identity_id + rows = await IdentitySubjectService.repair_missing_subjects(service_session, [1102, 1103]) + assert len(rows) == 1 + assert await IdentitySubjectService.repair_missing_subjects(service_session, [1102, 1103]) == [] + assert len((await service_session.execute(select(SysIdentitySubject))).scalars().all()) == _REPAIRED_SUBJECT_COUNT + + +@pytest.mark.asyncio +async def test_auth_version_update_is_atomic(service_session: AsyncSession) -> None: + """认证版本条件不匹配时不得更新。""" + service_session.add(SysUser(user_id=1104, user_name='version', nick_name='Version', status='0', del_flag='0')) + await service_session.flush() + await IdentitySubjectService.repair_missing_subject(service_session, 1104) + updated = await IdentitySubjectService.increment_auth_version(service_session, 1104, expected_version=1) + assert updated.auth_version == _NEXT_AUTH_VERSION + with pytest.raises(OAuthProtocolException): + await IdentitySubjectService.increment_auth_version(service_session, 1104, expected_version=1) + + +@pytest.mark.asyncio +async def test_independent_audit_writer_survives_main_transaction_rollback(service_session: AsyncSession) -> None: + """主事务回滚时,独立审计提交仍保留主体完整性告警。""" + factory = async_sessionmaker(service_session.info['service_engine'], expire_on_commit=False) + + async def audit_writer(user_id: int) -> None: + async with factory() as audit_session: + await AuditService.record( + audit_session, + event_type=OidcAuditEvent.IDENTITY_SUBJECT_MISSING, + result='failure', + risk_level='high', + user_id=user_id, + failure_code='identity_integrity', + ) + await audit_session.commit() + + with pytest.raises(OAuthProtocolException): + await IdentitySubjectService.require_by_user_id(service_session, _ROLLBACK_USER_ID, audit_writer=audit_writer) + service_session.add( + SysUser(user_id=_ROLLBACK_USER_ID, user_name='rollback', nick_name='Rollback', status='0', del_flag='0') + ) + await service_session.flush() + await service_session.rollback() + async with factory() as audit_session: + event = ( + (await audit_session.execute(select(SysOAuthAuditLog).where(SysOAuthAuditLog.user_id == _ROLLBACK_USER_ID))) + .scalars() + .first() + ) + assert event is not None + assert event.risk_level == 'high' + assert ( + await service_session.execute(select(SysUser).where(SysUser.user_id == _ROLLBACK_USER_ID)) + ).scalars().first() is None diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_interaction_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_interaction_service.py new file mode 100644 index 000000000..349ca3ded --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_interaction_service.py @@ -0,0 +1,235 @@ +import json +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock + +import pytest + +from exceptions.exception import OidcInteractionException +from module_identity.dao.oauth_client_dao import OAuthClientDao +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.interaction_service import InteractionService +from tests.module_identity.support.redis_fakes import FakeRedis +from utils.oidc_util import OidcUtil + +_PEPPER = 'interaction-test-pepper-' + 'x' * 32 +_CHALLENGE = 'A' * 43 + + +def _payload(**overrides: object) -> dict[str, object]: + """构造 AuthorizationContext 的内部白名单载荷。""" + value: dict[str, object] = { + 'interactionId': 'context-id', + 'clientPk': 1001, + 'clientId': 'portal-client', + 'redirectUri': 'https://portal.example/callback', + 'responseType': 'code', + 'scopes': ['openid', 'profile'], + 'resources': ['https://api.example'], + 'state': 'opaque-state', + 'nonce': 'opaque-nonce', + 'codeChallenge': _CHALLENGE, + 'codeChallengeMethod': 'S256', + 'prompt': 'login', + 'maxAge': 3600, + 'consentRequired': True, + } + value.update(overrides) + return value + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('interaction_page_metadata') +async def test_create_uses_nx_ttl_and_returns_csrf_only_once() -> None: + """验证 Interaction 创建使用 NX/TTL,Redis 仅存 csrfHash。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), ttl_seconds=300, pepper=_PEPPER) + record = await InteractionService.get_record(redis, created.interaction_id) + csrf = created.csrf_token + key = OidcRedisKey.interaction(created.interaction_id) + stored = redis.values[key][0] + assert record['status'] == 'awaiting_login' + assert record['csrfHash'] not in csrf + assert 'csrfHash' in stored + assert csrf not in stored + assert redis.set_calls[-1][1] == {'ex': 300, 'nx': True} + page = await InteractionService.get(redis, created.interaction_id, object()) + assert page['expiresIn'] > 0 + + +@pytest.mark.asyncio +@pytest.mark.parametrize('ttl', [0, False, -1]) +async def test_create_rejects_explicit_invalid_ttl(ttl: object) -> None: + """验证显式零值、布尔值和负数不会回退到默认 TTL。""" + redis = FakeRedis() + with pytest.raises(ValueError): + await InteractionService.create(redis, _payload(), ttl_seconds=ttl, pepper=_PEPPER) # type: ignore[arg-type] + assert not redis.values + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('interaction_page_metadata') +async def test_page_projection_never_leaks_internal_protocol_fields() -> None: + """验证页面载荷不包含 Redirect、协议绑定、用户或 CSRF 内部字段。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + page = await InteractionService.get(redis, created.interaction_id, object()) + forbidden = { + 'redirectUri', + 'state', + 'nonce', + 'codeChallenge', + 'codeChallengeMethod', + 'csrfHash', + 'userId', + 'subjectId', + 'authVersion', + 'sid', + 'clientPk', + 'maxAge', + } + assert forbidden.isdisjoint(page) + assert set(page) == {'interactionId', 'client', 'requestedScopes', 'nextAction', 'captchaEnabled', 'expiresIn'} + assert page['nextAction'] == 'login' + assert page['client']['clientName'] == '示例门户' + assert page['client']['policyUri'] == 'https://portal.example/privacy' + assert page['requestedScopes'][1]['name'] == '基本资料' + assert page['requestedScopes'][1]['sensitive'] is True + assert page['requestedScopes'][1]['required'] is False + assert page['requestedScopes'][0]['required'] is True + assert page['requestedScopes'][1]['description'] + assert 'private-' not in json.dumps(page) + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('interaction_page_metadata') +async def test_csrf_is_constant_time_verified_and_not_rotated_by_read() -> None: + """验证 CSRF 正确匹配、错误拒绝,读取页面不会重新建立 Token。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + record = await InteractionService.get_record(redis, created.interaction_id) + csrf = created.csrf_token + assert InteractionService.verify_csrf(record, csrf, pepper=_PEPPER) is True + assert InteractionService.verify_csrf(record, csrf + 'x', pepper=_PEPPER) is False + await InteractionService.get(redis, created.interaction_id, object()) + assert InteractionService.verify_csrf(record, csrf, pepper=_PEPPER) is True + + +@pytest.mark.asyncio +async def test_transition_uses_lua_cas_and_rejects_illegal_rollback() -> None: + """验证状态流转保留 TTL,非法回退和终态回退均拒绝。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + updated = await InteractionService.transition(redis, created.interaction_id, {'awaiting_login'}, 'awaiting_consent') + assert updated['nextAction'] == 'consent' + assert redis.eval_calls + with pytest.raises(ValueError): + await InteractionService.transition(redis, created.interaction_id, {'awaiting_consent'}, 'awaiting_login') + completed = await InteractionService.transition(redis, created.interaction_id, {'awaiting_consent'}, 'completed') + assert completed['nextAction'] == 'redirect' + with pytest.raises(ValueError): + await InteractionService.transition(redis, created.interaction_id, {'completed'}, 'denied') + + +@pytest.mark.asyncio +async def test_prompt_none_fails_without_interactive_creation() -> None: + """验证 prompt=none 不登录时返回 login_required,不写入交互状态。""" + redis = FakeRedis() + with pytest.raises(OidcInteractionException) as raised: + await InteractionService.create(redis, _payload(prompt='none'), pepper=_PEPPER) + assert raised.value.error == 'login_required' + assert not redis.values + + with pytest.raises(OidcInteractionException) as raised: + await InteractionService.create(redis, _payload(prompt='none', authenticatedSid='sid-1'), pepper=_PEPPER) + assert raised.value.error == 'consent_required' + assert not redis.values + + +@pytest.mark.asyncio +async def test_initial_status_respects_sso_prompt_and_consent() -> None: + """验证已有 SSO、prompt 和同意要求共同决定初始状态。""" + redis = FakeRedis() + created = await InteractionService.create( + redis, _payload(prompt=None, authenticatedSid='sid-1', consentRequired=True), pepper=_PEPPER + ) + assert created.initial_status == 'awaiting_consent' + + login = await InteractionService.create( + redis, _payload(prompt='login', authenticatedSid='sid-2', consentRequired=False), pepper=_PEPPER + ) + assert login.initial_status == 'awaiting_login' + + silent = await InteractionService.create( + redis, _payload(prompt='none', authenticatedSid='sid-3', consentRequired=False), pepper=_PEPPER + ) + assert silent.initial_status == 'completed' + assert (await InteractionService.get_record(redis, silent.interaction_id))['status'] == 'completed' + + +@pytest.mark.asyncio +async def test_record_validation_and_missing_ttl_fail_closed() -> None: + """验证损坏记录和永久键不会进入状态机。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + key = OidcRedisKey.interaction(created.interaction_id) + value, _ = redis.values[key] + corrupted = value.replace('"version":1', '"version":-1') + redis.values[key] = (corrupted, None) + with pytest.raises(OidcInteractionException) as raised: + await InteractionService.get(redis, created.interaction_id, object()) + assert raised.value.error == 'invalid_request' + assert key not in redis.values + + +@pytest.mark.asyncio +async def test_transition_rejects_interaction_without_ttl(monkeypatch: pytest.MonkeyPatch) -> None: + """服务层将 Redis CAS 返回的无 TTL 哨兵转换为协议错误并清理键。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + key = OidcRedisKey.interaction(created.interaction_id) + value, _ = redis.values[key] + redis.values[key] = (value, None) + + async def get_record(*_args: object, **_kwargs: object) -> dict[str, object]: + return json.loads(value) + + monkeypatch.setattr(InteractionService, '_get_record', classmethod(get_record)) + with pytest.raises(OidcInteractionException, match='认证交互有效期无效'): + await InteractionService.transition(redis, created.interaction_id, {'awaiting_login'}, 'awaiting_consent') + assert key not in redis.values + + +@pytest.mark.asyncio +async def test_transition_rejects_duplicate_scope_and_incomplete_identity() -> None: + """验证创建和状态更新不接受重复 Scope 或不完整身份三元组。""" + redis = FakeRedis() + with pytest.raises(ValueError): + await InteractionService.create(redis, _payload(scopes=['openid', 'openid']), pepper=_PEPPER) + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + with pytest.raises(ValueError): + await InteractionService.transition( + redis, created.interaction_id, {'awaiting_login'}, 'awaiting_consent', {'userId': 1} + ) + + +def test_max_age_decision_uses_project_local_time() -> None: + """验证 max_age 使用项目约定的本地无时区时间。""" + now = datetime(2026, 8, 24, 13, 0, tzinfo=timezone.utc) + assert OidcUtil.requires_reauthentication(now - timedelta(seconds=10), 5, now) is True + assert OidcUtil.requires_reauthentication(now - timedelta(seconds=2), 5, now) is False + assert OidcUtil.requires_reauthentication(now, 5, now) is False + + +@pytest.mark.asyncio +@pytest.mark.usefixtures('interaction_page_metadata') +@pytest.mark.parametrize('change', ['client_disabled', 'scope_unbound']) +async def test_page_rejects_removed_application_or_scope(monkeypatch: pytest.MonkeyPatch, change: str) -> None: + """请求创建后停用应用或移除权限时,不继续展示陈旧授权选项。""" + redis = FakeRedis() + created = await InteractionService.create(redis, _payload(), pepper=_PEPPER) + if change == 'client_disabled': + monkeypatch.setattr(OAuthClientDao, 'get_by_pk', AsyncMock(return_value=None)) + else: + monkeypatch.setattr(OAuthClientDao, 'list_scopes', AsyncMock(return_value=[])) + with pytest.raises(OidcInteractionException): + await InteractionService.get(redis, created.interaction_id, object()) diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_introspection_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_introspection_service.py new file mode 100644 index 000000000..c5fed228c --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_introspection_service.py @@ -0,0 +1,424 @@ +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace + +import pytest + +from config.env import OidcConfig +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.dao.oauth_grant_dao import OAuthGrantDao +from module_identity.security.opaque_token import token_digest +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.token_protocol_service import IntrospectionService +from tests.module_identity.support.redis_fakes import FakeRedis + +_PEPPER = 'introspection-test-pepper-' + 'x' * 32 +_AUTH_VERSION = 3 + + +def _config() -> SimpleNamespace: + """构造内省所需的最小配置。""" + + return SimpleNamespace( + oidc_issuer='https://auth.example.com', + oidc_allowed_clock_skew_seconds=60, + oidc_token_hash_pepper=_PEPPER, + ) + + +@pytest.fixture(autouse=True) +def _configure_oidc(monkeypatch: pytest.MonkeyPatch) -> None: + """将全局 OIDC 配置固定为本文件测试所需的协议值。""" + + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com') + monkeypatch.setattr(OidcConfig, 'oidc_allowed_clock_skew_seconds', 60) + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', _PEPPER) + + +def _client(client_type: str = 'confidential') -> SimpleNamespace: + """构造已启用的 Client 数据快照。""" + + return SimpleNamespace( + client_pk=10, + client_id='introspector', + client_type=client_type, + status='0', + policy_version=1, + token_endpoint_auth_method='client_secret_basic' if client_type == 'confidential' else 'none', + grant_types=['authorization_code', 'refresh_token'], + ) + + +def _principal(client_type: str = 'confidential') -> OAuthClientPrincipal: + """构造完成认证的 Client Principal。""" + + return OAuthClientPrincipal('introspector', client_type, 'client_secret_basic') + + +@pytest.mark.asyncio +async def test_introspection_rejects_non_client_principal(monkeypatch: pytest.MonkeyPatch) -> None: + """任意用户主体或普通对象都不能调用内省端点。""" + + result = await IntrospectionService.introspect(object(), FakeRedis(), 'not-a-token', object()) + + assert result == {'active': False} + + +@pytest.mark.asyncio +async def test_introspection_rejects_public_or_inactive_client(monkeypatch: pytest.MonkeyPatch) -> None: + """内省调用方必须是数据库事实中的启用 Confidential Client。""" + + async def client_lookup(*_args: object, **_kwargs: object) -> SimpleNamespace: + return _client('public') + + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', client_lookup) + result = await IntrospectionService.introspect(object(), FakeRedis(), 'not-a-token', _principal()) + + assert result == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ('client_type', 'db_method', 'principal_type', 'principal_method', 'expected'), + [ + ('public', 'none', 'public', 'none', False), + ('confidential', 'client_secret_basic', 'confidential', 'client_secret_basic', True), + ('confidential', 'none', 'confidential', 'client_secret_basic', False), + ('confidential', 'client_secret_basic', 'confidential', 'client_secret_post', False), + ('confidential', 'client_secret_basic', 'public', 'none', False), + ], +) +async def test_introspection_auth_method_pair_is_exact( + monkeypatch: pytest.MonkeyPatch, + client_type: str, + db_method: str, + principal_type: str, + principal_method: str, + expected: bool, +) -> None: + """Introspection 仅接受数据库与 Principal 同时为 Confidential/Basic。""" + + client = _client(client_type) + client.token_endpoint_auth_method = db_method + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _lookup(client) + ) + + result = await IntrospectionService._resolve_caller( + object(), OAuthClientPrincipal('introspector', principal_type, principal_method) + ) + + assert (result is not None) is expected + + +@pytest.mark.asyncio +async def test_access_introspection_honours_revoked_jti_and_resource_owner( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Access JWT 必须通过签名 Profile、撤销键和 Resource 内省授权。""" + + client = _client() + + async def client_lookup(*_args: object, **_kwargs: object) -> SimpleNamespace: + return client + + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', client_lookup) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'subject', + 'client_id': 'introspector', + 'aud': ['https://auth.example.com/oauth2/userinfo', 'https://api.example'], + 'scope': 'openid api.read', + 'iat': 100, + 'exp': 200, + 'jti': 'jti-1', + }, + ) + monkeypatch.setattr( + IntrospectionService, + '_resources_owned_by_caller', + classmethod(lambda cls, *_args, **_kwargs: _async_true()), + ) + redis = FakeRedis() + await redis.set('oidc:revoked_jti:jti-1', '1', ex=60) + + result = await IntrospectionService.introspect( + object(), redis, 'signed-access', _principal(), now=datetime.now(timezone.utc) + ) + + assert result == {'active': False} + + +@pytest.mark.asyncio +async def test_business_client_a_is_introspected_by_independent_resource_client_b( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """业务 Client A 签发的 Token 可由 Resource 专用 Client B 内省。""" + + caller = _client() + issuer = _client() + issuer.client_pk = 20 + issuer.client_id = 'business-a' + issuer.grant_types = ['authorization_code', 'refresh_token'] + + async def client_lookup(_db: object, client_id: str, active_only: bool = True) -> SimpleNamespace: + return caller if client_id == 'introspector' else issuer + + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', client_lookup) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'subject-a', + 'client_id': 'business-a', + 'aud': ['https://auth.example.com/oauth2/userinfo', 'https://api.example'], + 'scope': 'api.read', + 'gty': 'authorization_code', + 'ver': 3, + 'iat': 100, + 'exp': 500, + 'jti': 'jti-a', + }, + ) + monkeypatch.setattr(IntrospectionService, '_resources_owned_by_caller', lambda *_args, **_kwargs: _async_true()) + monkeypatch.setattr(IntrospectionService, '_client_allows_access', lambda *_args, **_kwargs: _async_true()) + monkeypatch.setattr(IntrospectionService, '_access_user_state', lambda *_args, **_kwargs: _async_user()) + result = await IntrospectionService.introspect(object(), FakeRedis(), 'signed-access', _principal()) + + assert result['active'] is True + assert result['client_id'] == 'business-a' + assert result['username'] == 'alice' + assert result['gty'] == 'authorization_code' + assert result['ver'] == _AUTH_VERSION + + +@pytest.mark.asyncio +async def test_independent_resource_client_b_is_rejected_for_wrong_audience( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Resource 不属于调用方 B 时,即使 Token 签名正确也必须拒绝。""" + + caller = _client() + issuer = _client() + issuer.client_pk = 20 + issuer.client_id = 'business-a' + + async def client_lookup(_db: object, client_id: str, active_only: bool = True) -> SimpleNamespace: + return caller if client_id == 'introspector' else issuer + + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', client_lookup) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'subject-a', + 'client_id': 'business-a', + 'aud': ['https://api.example'], + 'scope': 'api.read', + 'iat': 100, + 'exp': 500, + 'jti': 'jti-a', + }, + ) + monkeypatch.setattr(IntrospectionService, '_resources_owned_by_caller', lambda *_args, **_kwargs: _async_false()) + + result = await IntrospectionService.introspect(object(), FakeRedis(), 'signed-access', _principal()) + + assert result == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize('invalid_state', ['subject', 'user', 'session', 'grant']) +async def test_access_user_state_rejects_each_real_time_failure( + monkeypatch: pytest.MonkeyPatch, invalid_state: str +) -> None: + """用户 Access Token 的 Subject、用户、版本、Session、Grant 任一失效都拒绝。""" + + subject = SimpleNamespace(subject_id='subject-a', user_id=7, auth_version=3) + user = SimpleNamespace(status='0', del_flag='0', user_name='alice') + session = SimpleNamespace(user_id=7, subject_id='subject-a', auth_version=3) + db = object() + if invalid_state == 'subject': + subject = None + elif invalid_state == 'user': + user.status = '1' + elif invalid_state == 'session': + session = None + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.IdentitySubjectDao.get_by_subject_id', _lookup(subject) + ) + monkeypatch.setattr('module_identity.service.token_protocol_service.IdentityUserDao.get_user', _lookup(user)) + monkeypatch.setattr('module_identity.service.token_protocol_service.SsoSessionDao.get_active', _lookup(session)) + monkeypatch.setattr(OAuthAccessPolicyDao, 'is_blocked', _lookup(False)) + monkeypatch.setattr(OAuthGrantDao, 'get_by_grant_id', _lookup(None)) + claims = { + 'sub': 'subject-a', + 'ver': 3, + 'sid': 'sid-a', + 'scope': 'api.read', + 'grant_id': 'grant-a', + } + + assert not await IntrospectionService._access_user_state( + db, claims, _client(), ['https://api.example'], datetime.now(timezone.utc) + ) + + +@pytest.mark.asyncio +async def test_machine_access_token_uses_client_binding_without_user_state( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """client_credentials Token 校验机器 Subject/Grant 能力,不读取用户状态。""" + + caller = _client() + issuer = _client() + issuer.client_pk = 20 + issuer.client_id = 'machine-a' + issuer.grant_types = ['client_credentials'] + + async def client_lookup(_db: object, client_id: str, active_only: bool = True) -> SimpleNamespace: + return caller if client_id == 'introspector' else issuer + + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', client_lookup) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.list_resources', + _lookup([SimpleNamespace(audience='https://api.example', status='0')]), + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'client:machine-a', + 'client_id': 'machine-a', + 'aud': ['https://api.example'], + 'scope': 'api.read', + 'gty': 'client_credentials', + 'iat': 100, + 'exp': 500, + 'jti': 'machine-jti', + }, + ) + monkeypatch.setattr(IntrospectionService, '_resources_owned_by_caller', lambda *_args, **_kwargs: _async_true()) + monkeypatch.setattr(IntrospectionService, '_client_allows_access', lambda *_args, **_kwargs: _async_true()) + + result = await IntrospectionService.introspect(object(), FakeRedis(), 'machine-access', _principal()) + + assert result['active'] is True + + +@pytest.mark.asyncio +async def test_invalid_refresh_shape_is_inactive(monkeypatch: pytest.MonkeyPatch) -> None: + """Refresh Token 内省必须验证完整 typed opaque token。""" + + async def client_lookup(*_args: object, **_kwargs: object) -> SimpleNamespace: + return _client() + + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', client_lookup) + result = await IntrospectionService.introspect(object(), FakeRedis(), 'rt1.token-id', _principal()) + + assert result == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize('session_auth_version', [3, 4]) +async def test_refresh_introspection_rechecks_identity_session_grant_and_resources( + monkeypatch: pytest.MonkeyPatch, session_auth_version: int +) -> None: + """Refresh 内省成功前重新检查用户版本、Session、Grant 和 Resource。""" + monkeypatch.setattr(OAuthAccessPolicyDao, 'is_blocked', _lookup(False)) + + caller = _client() + client = _client() + client.client_pk = 20 + client.client_id = 'business-a' + token = 'rt1.token-0003.' + ('A' * 43) + now = datetime.now(timezone.utc) + row = SimpleNamespace( + token_id='token-0003', + token_hash='', + client_pk=20, + family_id='family-3', + status='active', + user_id=7, + subject_id='subject-7', + auth_version=3, + sid='sid-7', + grant_id='grant-7', + scopes=['api.read'], + resources=['https://api.example'], + issued_at=now - timedelta(minutes=1), + idle_expires_at=now + timedelta(minutes=5), + absolute_expires_at=now + timedelta(hours=1), + ) + row.token_hash = token_digest(token, _PEPPER) + grant = SimpleNamespace( + user_id=7, + subject_id='subject-7', + client_pk=20, + status='active', + client_policy_version=1, + expires_at=None, + granted_scopes=['api.read'], + granted_resources=['https://api.example'], + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _lookup(caller) + ) + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthTokenDao.get_by_token_id', _lookup(row)) + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthClientDao.get_by_pk', _lookup(client)) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.IntrospectionService._refresh_family_active', + lambda *_args, **_kwargs: _async_true(), + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.IdentitySubjectDao.get_by_user_id', + _lookup(SimpleNamespace(subject_id='subject-7', auth_version=3)), + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.IdentityUserDao.get_user', + _lookup(SimpleNamespace(status='0', del_flag='0', user_name='alice')), + ) + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthGrantDao.get_by_grant_id', _lookup(grant)) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.SsoSessionDao.get_active', + _lookup(SimpleNamespace(subject_id='subject-7', user_id=7, auth_version=session_auth_version)), + ) + db = object() + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.IntrospectionService._resources_owned_by_caller', + lambda *_args, **_kwargs: _async_true(), + ) + + result = await IntrospectionService.introspect(db, FakeRedis(), token, _principal(), now=now) + + assert result['active'] is (session_auth_version == _AUTH_VERSION) + if session_auth_version == _AUTH_VERSION: + assert result['client_id'] == 'business-a' + + +def _lookup(value: object) -> object: + """返回固定值的异步查询桩。""" + + async def lookup(*_args: object, **_kwargs: object) -> object: + return value + + return lookup + + +async def _async_true() -> bool: + """返回异步 True,供类方法 monkeypatch 使用。""" + + return True + + +async def _async_user() -> SimpleNamespace: + """返回带当前用户名的已验证用户快照。""" + + return SimpleNamespace(user_name='alice') + + +async def _async_false() -> bool: + """返回异步 False,供失败路径 monkeypatch 使用。""" + + return False diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_key_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_key_service.py new file mode 100644 index 000000000..328ee2d64 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_key_service.py @@ -0,0 +1,696 @@ +from datetime import datetime, timedelta, timezone +from pathlib import Path +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +import pytest_asyncio +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from jwt.utils import base64url_encode +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from common.constant import OidcAuditEvent +from config.database import Base +from config.env import OidcConfig +from module_identity.dao.oidc_key_dao import OidcKeyDao +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.do.oidc_key_do import SysOidcSigningKey +from module_identity.service.key_service import KeyService, KeyServiceError, OidcKeyManagementService +from utils.oidc_util import OidcUtil + +_MIN_TEST_RSA_BITS = 2048 + + +@pytest_asyncio.fixture +async def data_session() -> AsyncSession: + """仅创建密钥和审计表,隔离项目其他 ORM 的 SQLite 方言差异。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + tables = [SysOidcSigningKey.__table__, SysOAuthAuditLog.__table__] + async with engine.begin() as connection: + await connection.run_sync(lambda sync_connection: Base.metadata.create_all(sync_connection, tables=tables)) + session_factory = async_sessionmaker(engine, expire_on_commit=False) + async with session_factory() as session: + yield session + await engine.dispose() + + +@pytest.fixture(autouse=True) +def _configure_oidc(monkeypatch: pytest.MonkeyPatch) -> None: + """将全局 OIDC 配置固定为密钥服务测试所需的基线。""" + + values = { + 'oidc_enabled': True, + 'oidc_active_kid': 'k1', + 'oidc_signing_private_key_path': '', + 'oidc_signing_key_source': 'file', + 'oidc_signing_key_encryption_key': '', + 'oidc_signing_algorithm': 'RS256', + 'oidc_key_rotation_overlap_seconds': 86400, + 'oidc_max_access_token_ttl_seconds': 1800, + 'oidc_allowed_clock_skew_seconds': 60, + 'oidc_access_token_ttl_seconds': 600, + 'oidc_id_token_ttl_seconds': 300, + } + for name, value in values.items(): + monkeypatch.setattr(OidcConfig, name, value) + + +def _config() -> OidcConfig: + """返回当前测试使用的全局 OIDC 配置。""" + + return OidcConfig + + +def _record(tmp_path: Path, *, status: str = 'active', matching: bool = True) -> SimpleNamespace: + """生成带文件私钥和公开 JWK 的内存记录。""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + numbers = key.public_key().public_numbers() + n = base64url_encode(numbers.n.to_bytes((numbers.n.bit_length() + 7) // 8, 'big')).decode() + if not matching: + other = rsa.generate_private_key(public_exponent=65537, key_size=2048).public_key().public_numbers() + n = base64url_encode(other.n.to_bytes((other.n.bit_length() + 7) // 8, 'big')).decode() + path = tmp_path / 'signing.pem' + path.write_bytes( + key.private_bytes(serialization.Encoding.PEM, serialization.PrivateFormat.PKCS8, serialization.NoEncryption()) + ) + + now = datetime.now(tz=timezone.utc) + return SimpleNamespace( + kid='k1', + alg='RS256', + key_use='sig', + status=status, + signing_start_at=now - timedelta(minutes=1), + signing_stop_at=None, + publish_at=now - timedelta(minutes=1), + remove_from_jwks_at=None, + private_key_ref=str(path), + private_key_ciphertext=None, + public_jwk={'kty': 'RSA', 'use': 'sig', 'kid': 'k1', 'alg': 'RS256', 'n': n, 'e': 'AQAB'}, + create_time=now, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('kid', ['../key', 'key/name', r'key\\name', 'key%2Fname', 'key name', '密钥']) +async def test_key_lifecycle_rejects_unsafe_kid_before_database_access(kid: str) -> None: + """激活、退役和删除入口统一拒绝无法安全放入路径的 kid。""" + for operation in (KeyService.activate_key, KeyService.retire_key, KeyService.delete_key): + with pytest.raises(KeyServiceError, match='包含不允许的字符'): + await operation(object(), kid) + + +@pytest.mark.asyncio +async def test_activate_due_processes_only_scheduled_keys_and_rolls_back_failures( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """后台只处理到期 pending,并在单项失败时回滚该项事务。""" + + class _Session: + commits = 0 + rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + session = _Session() + due = SimpleNamespace(kid='due-key') + scheduled: list[str] = [] + + async def due_keys(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return [due] + + async def activate(*args: object, **kwargs: object) -> bool: + scheduled.append(args[1]) + return True + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_due_pending', due_keys) + monkeypatch.setattr(KeyService, 'activate_key', activate) + assert await KeyService.activate_due(session) == 1 + assert scheduled == ['due-key'] + assert session.commits == 1 + assert session.rollbacks == 0 + + async def fail(*args: object, **kwargs: object) -> bool: + raise KeyServiceError('private key does not match public JWK') + + audit_events: list[tuple[object, ...]] = [] + + async def audit_writer(*args: object, **kwargs: object) -> None: + audit_events.append(args) + + monkeypatch.setattr(KeyService, 'activate_key', fail) + assert await KeyService.activate_due(session, audit_writer=audit_writer) == 0 + assert session.rollbacks == 1 + assert audit_events and audit_events[0][1:3] == (OidcAuditEvent.SIGNING_KEY_ROTATED, 'failure') + + +@pytest.mark.asyncio +async def test_activate_due_audits_first_failure_and_continues_next_key(monkeypatch: pytest.MonkeyPatch) -> None: + """首把激活失败必须独立审计,且不得阻断后续到期密钥。""" + + class _Session: + commits = 0 + rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + session = _Session() + pending = [SimpleNamespace(kid='first'), SimpleNamespace(kid='second')] + audit_events: list[dict[str, object]] = [] + + async def due_keys(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return pending + + async def activate(db: object, kid: str, **kwargs: object) -> bool: + if kid == 'first': + raise KeyServiceError('private key does not match public JWK') + return True + + async def audit_writer(*args: object, **kwargs: object) -> None: + audit_events.append(kwargs) + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_due_pending', due_keys) + monkeypatch.setattr(KeyService, 'activate_key', activate) + assert await KeyService.activate_due(session, audit_writer=audit_writer) == 1 + assert session.rollbacks == 1 and session.commits == 1 + assert audit_events[0]['failure_code'] == 'key_activation_failed' + + +@pytest.mark.asyncio +async def test_activate_due_continues_when_failure_audit_writer_fails(monkeypatch: pytest.MonkeyPatch) -> None: + """独立失败审计异常只能记录安全日志,不能阻断后续密钥。""" + pending = [SimpleNamespace(kid='first'), SimpleNamespace(kid='second')] + activated: list[str] = [] + + class _Session: + async def commit(self) -> None: + pass + + async def rollback(self) -> None: + pass + + async def due_keys(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return pending + + async def activate(db: object, kid: str, **kwargs: object) -> bool: + if kid == 'first': + raise KeyServiceError('activation failed') + activated.append(kid) + return True + + async def failed_audit(*args: object, **kwargs: object) -> None: + raise RuntimeError('audit database unavailable') + + warnings: list[str] = [] + monkeypatch.setattr( + 'module_identity.service.key_service.logger.warning', + lambda message, *args: warnings.append(message.format(*args)), + ) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_due_pending', due_keys) + monkeypatch.setattr(KeyService, 'activate_key', activate) + assert await KeyService.activate_due(_Session(), audit_writer=failed_audit) == 1 + assert activated == ['second'] + assert warnings == [ + 'OIDC 签名密钥激活审计记录写入失败,密钥标识=first', + 'OIDC 签名密钥激活失败,事件=signing_key_rotated,密钥标识=first', + ] + + +@pytest.mark.asyncio +async def test_list_due_pending_filters_schedule_and_algorithm(data_session: AsyncSession) -> None: + """到期查询只返回 RS256、pending 且两个时间点均已到达的密钥。""" + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + rows = [ + SysOidcSigningKey( + kid='due', + alg='RS256', + status='pending', + private_key_ref='test.pem', + public_jwk={'kty': 'RSA'}, + publish_at=now - timedelta(seconds=1), + signing_start_at=now - timedelta(seconds=1), + create_by='test', + ), + SysOidcSigningKey( + kid='future-publish', + alg='RS256', + status='pending', + private_key_ref='test.pem', + public_jwk={'kty': 'RSA'}, + publish_at=now + timedelta(seconds=1), + signing_start_at=now - timedelta(seconds=1), + create_by='test', + ), + SysOidcSigningKey( + kid='future-signing', + alg='RS256', + status='pending', + private_key_ref='test.pem', + public_jwk={'kty': 'RSA'}, + publish_at=now - timedelta(seconds=1), + signing_start_at=now + timedelta(seconds=1), + create_by='test', + ), + SysOidcSigningKey( + kid='wrong-alg', + alg='ES256', + status='pending', + private_key_ref='test.pem', + public_jwk={'kty': 'EC'}, + publish_at=now - timedelta(seconds=1), + signing_start_at=now - timedelta(seconds=1), + create_by='test', + ), + SysOidcSigningKey( + kid='active', + alg='RS256', + status='active', + private_key_ref='test.pem', + public_jwk={'kty': 'RSA'}, + publish_at=now - timedelta(seconds=1), + signing_start_at=now - timedelta(seconds=1), + create_by='test', + ), + ] + data_session.add_all(rows) + await data_session.commit() + result = await OidcKeyDao.list_due_pending(data_session, now=now) + assert [row.kid for row in result] == ['due'] + + +@pytest.mark.asyncio +async def test_rotation_lock_is_cross_instance_exclusive_and_releases_atomically() -> None: + """轮换锁在 Redis 中互斥,并只允许持有者删除自己的锁。""" + + class _Redis: + value: str | None = None + + async def set(self, key: str, value: str, *, nx: bool, ex: int) -> bool: + if nx and self.value is not None: + return False + self.value = value + return True + + async def eval(self, script: str, count: int, key: str, token: str) -> int: + if self.value == token: + self.value = None + return 1 + return 0 + + redis = _Redis() + lock = KeyService._rotation_lock(redis) + await lock.__aenter__() + with pytest.raises(KeyServiceError, match='正在轮换'): + async with KeyService._rotation_lock(redis): + pass + await lock.__aexit__(None, None, None) + assert redis.value is None + + +@pytest.mark.asyncio +async def test_private_key_must_match_public_jwk(tmp_path: Path) -> None: + """公私钥不匹配时启动加载必须失败。""" + with pytest.raises(KeyServiceError, match='不匹配'): + await KeyService.load_private_key_async(_record(tmp_path, matching=False)) + + +@pytest.mark.asyncio +async def test_private_key_requires_active_signing_window(tmp_path: Path) -> None: + """pending 或已停止签名的密钥不能用于签名。""" + with pytest.raises(KeyServiceError, match='预期状态为 active'): + await KeyService.load_private_key_async(_record(tmp_path, status='pending')) + + +@pytest.mark.asyncio +async def test_imported_database_kid_is_validated_before_jwks_and_signing( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """历史或导入记录中的非法 kid 不得进入 JWKS 或签名加载。""" + record = _record(tmp_path) + record.kid = 'bad/key' + record.public_jwk['kid'] = 'bad/key' + with pytest.raises(KeyServiceError, match='包含不允许的字符'): + await KeyService.load_private_key_async(record) + + with pytest.raises(ValueError, match='包含不允许的字符'): + OidcUtil.normalize_public_jwk(record) + + async def published(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return [record] + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_published', published) + with pytest.raises(KeyServiceError, match='包含不允许的字符'): + await KeyService.build_jwks(object()) + + +@pytest.mark.asyncio +async def test_management_key_bootstrap_is_available_when_oidc_disabled(data_session: AsyncSession) -> None: + """OIDC 协议关闭时仍可用独立加密材料完成密钥管理引导。""" + config = _config() + config.oidc_enabled = False + config.oidc_signing_key_encryption_key = 'e' * 32 + row = await KeyService.create_pending_key( + data_session, + kid='bootstrap-1', + publish_at=datetime.now(timezone.utc) + timedelta(minutes=1), + actor='admin', + ) + assert row.status == 'pending' + assert row.private_key_ciphertext.startswith('v2.') + + +@pytest.mark.asyncio +async def test_pending_key_times_are_stored_and_returned_in_utc(data_session: AsyncSession) -> None: + """轮换时间持久化后仍保持 UTC,并按毫秒精度返回。""" + config = _config() + config.oidc_signing_key_encryption_key = 'e' * 32 + await KeyService.create_pending_key( + data_session, + kid='timezone-normalized', + publish_at=datetime(2026, 8, 27, 17, 37, 37, tzinfo=timezone.utc), + activate_at=datetime(2026, 8, 27, 17, 38, 31, tzinfo=timezone.utc), + actor='admin', + now=datetime(2026, 8, 27, 17, 30, tzinfo=timezone.utc), + ) + await data_session.commit() + data_session.expunge_all() + + row = (await data_session.execute(select(SysOidcSigningKey))).scalar_one() + assert row.publish_at == datetime(2026, 8, 27, 17, 37, 37, tzinfo=timezone.utc) + assert row.signing_start_at == datetime(2026, 8, 27, 17, 38, 31, tzinfo=timezone.utc) + view = OidcKeyManagementService.view(row) + assert view['publishAt'] == datetime(2026, 8, 27, 17, 37, 37, tzinfo=timezone.utc) + assert view['signingStartAt'] == datetime(2026, 8, 27, 17, 38, 31, tzinfo=timezone.utc) + + +@pytest.mark.asyncio +async def test_management_bootstrap_is_idempotent_and_activates_first_key(data_session: AsyncSession) -> None: + """部署初始化重复执行时只保留一把 active 密钥。""" + config = _config() + config.oidc_enabled = False + config.oidc_signing_key_encryption_key = 'e' * 32 + now = datetime.now(timezone.utc) + + first, first_created = await OidcKeyManagementService.bootstrap( + data_session, + kid='bootstrap-primary', + actor='system:deployment', + redis=None, + now=now, + ) + second, second_created = await OidcKeyManagementService.bootstrap( + data_session, + kid='ignored-on-retry', + actor='system:deployment', + redis=None, + now=now, + ) + + assert first_created is True + assert second_created is False + assert first['kid'] == second['kid'] == 'bootstrap-primary' + assert first['status'] == second['status'] == 'active' + rows = (await data_session.execute(select(SysOidcSigningKey))).scalars().all() + assert [(row.kid, row.status) for row in rows] == [('bootstrap-primary', 'active')] + + +@pytest.mark.asyncio +async def test_runtime_active_key_uses_database_state_not_bootstrap_kid(tmp_path: Path) -> None: + """在线轮换后静态旧 OIDC_ACTIVE_KID 不应阻断新 active key。""" + record = _record(tmp_path) + record.kid = 'new' + record.public_jwk['kid'] = 'new' + config = _config() + config.oidc_active_kid = 'old' + assert (await KeyService.load_private_key_async(record)).key_size >= _MIN_TEST_RSA_BITS + + +def _orm_key(record: SimpleNamespace, status: str) -> SysOidcSigningKey: + """把内存密钥记录转换为真实 ORM 密钥。""" + return SysOidcSigningKey( + kid=record.kid, + alg=record.alg, + public_jwk=record.public_jwk, + private_key_ref=record.private_key_ref, + status=status, + publish_at=record.publish_at, + signing_start_at=record.signing_start_at, + signing_stop_at=record.signing_stop_at, + remove_from_jwks_at=record.remove_from_jwks_at, + create_by='test', + create_time=record.create_time, + ) + + +@pytest.mark.asyncio +async def test_real_session_rotation_persists_state_and_injected_time( + data_session: AsyncSession, tmp_path: Path +) -> None: + """真实 AsyncSession 轮换后恰有一个 active 且窗口和时间准确。""" + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + old = _record(tmp_path) + new = _record(tmp_path) + new.kid = 'new' + new.public_jwk['kid'] = 'new' + new.publish_at = now - timedelta(seconds=1) + new.signing_start_at = now - timedelta(seconds=1) + old.signing_start_at = now - timedelta(minutes=1) + old.publish_at = now - timedelta(minutes=1) + old_row = _orm_key(old, 'active') + new_row = _orm_key(new, 'pending') + data_session.add_all([old_row, new_row]) + await data_session.flush() + config = _config() + config.oidc_active_kid = 'old' + config.oidc_key_rotation_overlap_seconds = 10 + config.oidc_access_token_ttl_seconds = 20 + config.oidc_max_access_token_ttl_seconds = 30 + config.oidc_id_token_ttl_seconds = 100 + config.oidc_allowed_clock_skew_seconds = 5 + assert await KeyService.activate_key(data_session, 'new', now=now) + await data_session.commit() + rows = (await data_session.execute(select(SysOidcSigningKey).order_by(SysOidcSigningKey.kid))).scalars().all() + assert [(row.kid, row.status) for row in rows] == [('k1', 'retiring'), ('new', 'active')] + active = next(row for row in rows if row.status == 'active') + retiring = next(row for row in rows if row.status == 'retiring') + assert active.signing_start_at == now + assert retiring.remove_from_jwks_at == now + timedelta(seconds=105) + + +@pytest.mark.asyncio +async def test_manual_activation_can_override_future_signing_schedule( + data_session: AsyncSession, tmp_path: Path +) -> None: + """管理员明确确认后可提前启用已公开的 pending 密钥。""" + now = datetime.now(tz=timezone.utc) + target = _record(tmp_path, status='pending') + target.publish_at = now - timedelta(minutes=5) + target.signing_start_at = now + timedelta(hours=1) + data_session.add(_orm_key(target, 'pending')) + await data_session.flush() + + assert await KeyService.activate_key(data_session, target.kid, now=now) + await data_session.commit() + + row = (await data_session.execute(select(SysOidcSigningKey))).scalar_one() + assert row.status == 'active' + assert row.signing_start_at == now.replace(microsecond=now.microsecond // 1000 * 1000) + + +@pytest.mark.asyncio +async def test_real_session_rotation_rolls_back_on_key_mismatch(data_session: AsyncSession, tmp_path: Path) -> None: + """候选私钥失败时真实事务回滚,旧 active 状态保持不变。""" + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + old = _record(tmp_path) + new = _record(tmp_path, matching=False) + new.kid = 'new' + new.public_jwk['kid'] = 'new' + new.publish_at = now - timedelta(seconds=1) + new.signing_start_at = now - timedelta(seconds=1) + old.signing_start_at = now - timedelta(minutes=1) + old.publish_at = now - timedelta(minutes=1) + data_session.add_all([_orm_key(old, 'active'), _orm_key(new, 'pending')]) + await data_session.flush() + await data_session.commit() + config = _config() + config.oidc_active_kid = 'old' + with pytest.raises(KeyServiceError, match='不匹配'): + await KeyService.activate_key(data_session, 'new', now=now) + await data_session.rollback() + rows = (await data_session.execute(select(SysOidcSigningKey).order_by(SysOidcSigningKey.kid))).scalars().all() + assert [(row.kid, row.status) for row in rows] == [('k1', 'active'), ('new', 'pending')] + + +@pytest.mark.asyncio +async def test_jwks_filters_private_jwk_fields(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + """JWKS 始终只输出公开字段,即使数据库 JSON 被污染。""" + now = datetime.now(timezone.utc) + record = _record(tmp_path) + record.publish_at = now - timedelta(seconds=1) + record.public_jwk['d'] = 'secret' + + async def published(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return [record] + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_published', published) + payload = await KeyService.build_jwks(object(), now=now) + assert payload == { + 'keys': [ + { + 'kty': 'RSA', + 'use': 'sig', + 'kid': 'k1', + 'alg': 'RS256', + 'n': record.public_jwk['n'], + 'e': record.public_jwk['e'], + } + ] + } + assert 'd' not in payload['keys'][0] + + +@pytest.mark.asyncio +async def test_jwks_excludes_keys_outside_publish_window(monkeypatch: pytest.MonkeyPatch) -> None: + """JWKS 不发布尚未到 publish_at 或已过 remove_from_jwks_at 的密钥。""" + now = datetime.now(timezone.utc) + base = { + 'kty': 'RSA', + 'use': 'sig', + 'alg': 'RS256', + 'n': 'n', + 'e': 'AQAB', + } + records = [ + SimpleNamespace( + kid='future', + alg='RS256', + key_use='sig', + status='pending', + publish_at=now + timedelta(seconds=1), + remove_from_jwks_at=None, + public_jwk={**base, 'kid': 'future'}, + ), + SimpleNamespace( + kid='removed', + alg='RS256', + key_use='sig', + status='retiring', + publish_at=now - timedelta(seconds=1), + remove_from_jwks_at=now, + public_jwk={**base, 'kid': 'removed'}, + ), + ] + + async def published(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return records + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_published', published) + assert await KeyService.build_jwks(object(), now=now) == {'keys': []} + + +@pytest.mark.asyncio +async def test_jwks_rejects_malformed_rsa_jwk(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + """非法 RSA n/e 不能污染 JWKS,服务选择 fail-closed。""" + now = datetime.now(timezone.utc) + record = _record(tmp_path) + record.publish_at = now - timedelta(seconds=1) + record.public_jwk['n'] = 'not-a-rsa-modulus' + + async def published(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return [record] + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.list_published', published) + with pytest.raises(KeyServiceError, match=r'Base64URL|RSA 参数'): + await KeyService.build_jwks(object(), now=now) + + +@pytest.mark.asyncio +async def test_activation_validates_pending_private_key_before_state_change( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + """候选密钥公私钥不匹配时不得调用状态切换。""" + target = _record(tmp_path, status='pending', matching=False) + target.publish_at = datetime.now(timezone.utc) - timedelta(seconds=1) + old = SimpleNamespace(kid='old', status='active') + changed = False + + async def lock_algorithm(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return [old, target] + + async def get_target(*args: object, **kwargs: object) -> SimpleNamespace: + return target + + async def get_old(*args: object, **kwargs: object) -> SimpleNamespace: + return old + + async def activate(*args: object, **kwargs: object) -> bool: + nonlocal changed + changed = True + return True + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.lock_algorithm_for_update', lock_algorithm) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.get_by_kid_for_update', get_target) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.get_active', get_old) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.activate', activate) + with pytest.raises(KeyServiceError, match='不匹配'): + await KeyService.activate_key(object(), 'k1') + assert changed is False + + +@pytest.mark.asyncio +async def test_activation_retains_old_key_for_max_overlap_window( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path +) -> None: + """轮换将旧 active 置为 retiring,并覆盖 overlap 与 Token 安全窗口的较大值。""" + now = datetime.now(tz=timezone.utc) + target = _record(tmp_path, status='pending') + target.publish_at = now - timedelta(seconds=1) + old = SimpleNamespace(kid='old', status='active', remove_from_jwks_at=None) + config = _config() + config.oidc_key_rotation_overlap_seconds = 7200 + config.oidc_max_access_token_ttl_seconds = 1800 + config.oidc_allowed_clock_skew_seconds = 60 + + async def audit(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr('module_identity.service.key_service.AuditService.record', audit) + + async def lock_algorithm(*args: object, **kwargs: object) -> list[SimpleNamespace]: + return [old, target] + + async def get_target(*args: object, **kwargs: object) -> SimpleNamespace: + return target + + async def get_old(*args: object, **kwargs: object) -> SimpleNamespace: + return old + + async def activate(*args: object, **kwargs: object) -> bool: + return True + + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.lock_algorithm_for_update', lock_algorithm) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.get_by_kid_for_update', get_target) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.get_active', get_old) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.activate', activate) + set_retiring = AsyncMock(return_value=True) + monkeypatch.setattr('module_identity.service.key_service.OidcKeyDao.set_retiring', set_retiring) + assert await KeyService.activate_key(object(), 'k1', now=now) + set_retiring.assert_awaited_once() + assert set_retiring.await_args.args[1:] == ('old', now + timedelta(seconds=7200), now) + + +def test_etag_is_stable_and_private_fields_are_not_serialized() -> None: + """ETag 对字段顺序稳定,公开结果不含私钥字段。""" + first = {'keys': [{'kid': 'k1', 'n': 'n', 'e': 'AQAB'}]} + second = {'keys': [{'e': 'AQAB', 'n': 'n', 'kid': 'k1'}]} + assert OidcUtil.json_etag(first) == OidcUtil.json_etag(second) diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_logout.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_logout.py new file mode 100644 index 000000000..824ad11b2 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_logout.py @@ -0,0 +1,806 @@ +from datetime import datetime, timezone +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock + +import httpx +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from jwt.utils import base64url_encode +from starlette.requests import Request + +from config.env import OidcConfig +from module_identity.controller import authorization_controller as controller +from module_identity.redis_keys import OidcRedisKey +from module_identity.security.backchannel_transport import PinnedHttpxTransport +from module_identity.security.jwt_profile import ( + JwtProfileError, + decode_id_token, + decode_logout_token, + encode_id_token, +) +from module_identity.service.infrastructure_service import AfterCommitCoordinator, OidcRateLimiter +from module_identity.service.session_service import ( + LogoutResult, + LogoutService, + LogoutServiceError, +) +from tests.module_identity.support.redis_fakes import FakeRedis +from utils.oidc_util import OidcUtil + +_ISSUER = 'https://auth.example.com' +_PEPPER = 'logout-test-pepper-' + 'x' * 32 +_MAX_ATTEMPTS = 3 +_HTTP_OK = 200 +_HTTP_SERVICE_UNAVAILABLE = 503 +_HTTP_NOT_FOUND = 404 +_HTTP_SEE_OTHER = 303 + + +@pytest.fixture +def now() -> datetime: + """每个用例执行时获取时间,避免等待前序测试期间令牌过期。""" + return datetime.now(timezone.utc) + + +@pytest.fixture(autouse=True) +def _disable_logout_rate_limit(monkeypatch: pytest.MonkeyPatch) -> None: + """Logout Controller 单元测试隔离 Redis 限流器。""" + + async def allow(*_args: object, **_kwargs: object) -> None: + return None + + monkeypatch.setattr(OidcRateLimiter, 'enforce', allow) + + +def _config() -> SimpleNamespace: + """构造 Logout 测试配置。""" + + return SimpleNamespace( + oidc_enabled=True, + oidc_issuer=_ISSUER, + oidc_allowed_clock_skew_seconds=60, + oidc_token_hash_pepper=_PEPPER, + ) + + +@pytest.fixture(autouse=True) +def _configure_oidc(monkeypatch: pytest.MonkeyPatch) -> None: + """将全局 OIDC 配置固定为 Logout 测试所需的协议值。""" + + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + monkeypatch.setattr(OidcConfig, 'oidc_issuer', _ISSUER) + monkeypatch.setattr(OidcConfig, 'oidc_allowed_clock_skew_seconds', 60) + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', _PEPPER) + + +def _key_record(private_key: object, now: datetime, kid: str = 'key-1') -> SimpleNamespace: + """构造本地数据库公开 JWK 记录。""" + + numbers = private_key.public_key().public_numbers() + public_jwk = { + 'kty': 'RSA', + 'use': 'sig', + 'kid': kid, + 'alg': 'RS256', + 'n': base64url_encode(numbers.n.to_bytes((numbers.n.bit_length() + 7) // 8, 'big')).decode(), + 'e': base64url_encode(numbers.e.to_bytes((numbers.e.bit_length() + 7) // 8, 'big')).decode(), + } + return SimpleNamespace( + kid=kid, + key_use='sig', + alg='RS256', + public_jwk=public_jwk, + status='active', + publish_at=now, + remove_from_jwks_at=None, + ) + + +def _id_token(private_key: object, *, now: datetime, kid: str = 'key-1', aud: str = 'portal') -> str: + """签发用于测试的真实 ID Token。""" + + return encode_id_token( + { + 'iss': _ISSUER, + 'sub': 'subject-1', + 'aud': aud, + 'exp': int(now.timestamp()) + 300, + 'iat': int(now.timestamp()), + 'auth_time': int(now.timestamp()), + 'nonce': 'nonce-1', + 'sid': 'sid-1', + 'acr': 'urn:test', + 'amr': ['pwd'], + }, + private_key, + kid, + ) + + +@pytest.mark.asyncio +async def test_id_token_hint_uses_local_rsa_kid_and_rejects_remote_header( + monkeypatch: pytest.MonkeyPatch, now: datetime +) -> None: + """ID Token Hint 仅使用本地 DB JWK,远程 header 参数直接拒绝。""" + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + client = SimpleNamespace(client_pk=20, client_id='portal', status='0') + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.get_by_client_id', _async_return(client) + ) + monkeypatch.setattr( + 'module_identity.service.session_service.OidcKeyDao.get_verifying', + _async_return(_key_record(private_key, now)), + ) + db = SimpleNamespace() + token = _id_token(private_key, now=now, aud='portal') + + claims, resolved = await LogoutService._validate_id_token_hint(db, token, now) + + assert claims['sid'] == 'sid-1' + assert resolved.client_id == 'portal' + unsafe = jwt.encode( + jwt.decode(token, options={'verify_signature': False}) | {'nonce': 'nonce-1'}, + private_key, + algorithm='RS256', + headers={'kid': 'key-1', 'typ': 'JWT', 'jku': 'https://evil.example/jwks'}, + ) + with pytest.raises(ValueError): + await LogoutService._validate_id_token_hint(db, unsafe, now) + + +@pytest.mark.asyncio +async def test_pending_signing_key_cannot_validate_id_token_hint( + monkeypatch: pytest.MonkeyPatch, now: datetime +) -> None: + """待激活签名密钥只能发布,不能用于验证 ID Token Hint。""" + + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + client = SimpleNamespace(client_pk=20, client_id='portal', status='0') + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.get_by_client_id', _async_return(client) + ) + key = _key_record(private_key, now) + key.status = 'pending' + monkeypatch.setattr('module_identity.service.session_service.OidcKeyDao.get_verifying', _async_return(key)) + db = SimpleNamespace() + token = _id_token(private_key, now=now, aud='portal') + + with pytest.raises(LogoutServiceError): + await LogoutService._validate_id_token_hint(db, token, now) + + +@pytest.mark.asyncio +async def test_logout_redirect_and_backchannel_are_after_commit(monkeypatch: pytest.MonkeyPatch, now: datetime) -> None: + """已注册 Redirect 安全附加 state,通知只在 commit 后执行。""" + + client = SimpleNamespace(client_pk=20, client_id='portal', status='0') + session = SimpleNamespace(sid='sid-1', status='active', subject_id='subject-1') + monkeypatch.setattr( + LogoutService, '_validate_id_token_hint', _async_return(({'sid': 'sid-1', 'sub': 'subject-1'}, client)) + ) + monkeypatch.setattr('module_identity.service.session_service.SsoSessionDao.get_by_sid', _async_return(session)) + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.find_exact_uri', + _async_return(SimpleNamespace(uri='https://portal.example/logged-out', status='0')), + ) + monkeypatch.setattr(LogoutService, '_lock_session_refresh_tokens', _async_return([])) + monkeypatch.setattr(LogoutService, '_revoke_session_state', _async_noop) + sent: list[tuple[str, str]] = [] + + async def register_backchannel(*args: object, **kwargs: object) -> None: + coordinator = kwargs['coordinator'] + + async def notify() -> None: + sent.append(('https://portal.example/oidc/backchannel-logout', 'logout-token')) + + await coordinator.register(notify) + + monkeypatch.setattr(LogoutService, '_register_backchannel', register_backchannel) + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + coordinator = AfterCommitCoordinator() + + monkeypatch.setattr( + 'module_identity.service.session_service.SsoSessionDao.client_ids_for_sid', AsyncMock(return_value=[]) + ) + result = await LogoutService._logout( + db, + object(), + id_token_hint='opaque-input-not-logged', + post_logout_redirect_uri='https://portal.example/logged-out', + state='opaque-state', + now=now, + coordinator=coordinator, + confirmed=True, + ) + + assert result.redirect_uri == 'https://portal.example/logged-out' + assert result.state == 'opaque-state' + assert sent == [] + await coordinator.rollback(db) + assert sent == [] + + coordinator = AfterCommitCoordinator() + result = await LogoutService._logout( + db, + object(), + id_token_hint='opaque-input-not-logged', + post_logout_redirect_uri='https://portal.example/logged-out', + state='opaque-state', + now=now, + coordinator=coordinator, + confirmed=True, + ) + await coordinator.commit(db) + assert result.is_local is False + assert sent == [('https://portal.example/oidc/backchannel-logout', 'logout-token')] + + +@pytest.mark.asyncio +async def test_invalid_hint_with_valid_cookie_still_revokes_server_session( + monkeypatch: pytest.MonkeyPatch, now: datetime +) -> None: + """无效 Hint 不得重定向,但有效 Cookie 仍必须撤销服务端 Session。""" + + session = SimpleNamespace(sid='sid-1', status='active', subject_id='subject-1') + + async def invalid_hint(*_args: object, **_kwargs: object) -> object: + raise LogoutServiceError('invalid id_token_hint') + + monkeypatch.setattr(LogoutService, '_validate_id_token_hint', invalid_hint) + monkeypatch.setattr( + 'module_identity.service.session_service.SsoSessionService.validate_logout_cookie', _async_return(session) + ) + monkeypatch.setattr(LogoutService, '_lock_session_refresh_tokens', _async_return([])) + revoked = AsyncMock() + monkeypatch.setattr(LogoutService, '_revoke_session_state', revoked) + monkeypatch.setattr(LogoutService, '_register_backchannel', _async_noop) + coordinator = AfterCommitCoordinator() + monkeypatch.setattr( + 'module_identity.service.session_service.SsoSessionDao.client_ids_for_sid', AsyncMock(return_value=[]) + ) + result = await LogoutService._logout( + SimpleNamespace(), + object(), + id_token_hint='bad-hint', + cookie='ss1.valid-cookie', + post_logout_redirect_uri='https://portal.example/logged-out', + state='must-not-echo', + now=now, + coordinator=coordinator, + confirmed=True, + ) + + assert result.session_revoked is True + assert result.is_local is True + revoked.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_unsafe_or_unregistered_redirect_falls_back_to_local_and_logout_token_is_strict( + monkeypatch: pytest.MonkeyPatch, now: datetime +) -> None: + """未注册 URI 不开放重定向;生成的 Logout Token 使用专用 Profile。""" + + client = SimpleNamespace(client_pk=20, client_id='portal', status='0') + session = SimpleNamespace(sid='sid-1', status='active', subject_id='subject-1') + monkeypatch.setattr( + LogoutService, '_validate_id_token_hint', _async_return(({'sid': 'sid-1', 'sub': 'subject-1'}, client)) + ) + monkeypatch.setattr('module_identity.service.session_service.SsoSessionDao.get_by_sid', _async_return(session)) + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.find_exact_uri', + _async_return(SimpleNamespace(uri='https://evil.example/redirect#fragment', status='0')), + ) + monkeypatch.setattr(LogoutService, '_lock_session_refresh_tokens', _async_return([])) + monkeypatch.setattr(LogoutService, '_revoke_session_state', _async_noop) + monkeypatch.setattr(LogoutService, '_register_backchannel', _async_noop) + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + coordinator = AfterCommitCoordinator() + + monkeypatch.setattr( + 'module_identity.service.session_service.SsoSessionDao.client_ids_for_sid', AsyncMock(return_value=[]) + ) + result = await LogoutService._logout( + db, + object(), + id_token_hint='hint', + post_logout_redirect_uri='https://evil.example/redirect#fragment', + state='do-not-echo', + now=now, + coordinator=coordinator, + confirmed=True, + ) + + assert result.is_local + assert result.state is None + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + token = await LogoutService._make_logout_token( + 'portal', + 'sid-1', + now=now, + signing_key=private_key, + signing_kid='key-1', + db=db, + ) + claims = decode_logout_token(token, verification_key=private_key.public_key(), issuer=_ISSUER, audience='portal') + assert claims['sid'] == 'sid-1' + assert claims['exp'] == claims['iat'] + 120 + assert claims['events']['http://schemas.openid.net/event/backchannel-logout'] == {} + assert 'nonce' not in claims + subject_token = await LogoutService._make_logout_token( + 'portal', + 'sid-1', + subject_id='subject-1', + include_sid=False, + event_jti='event-constant', + now=now, + signing_key=private_key, + signing_kid='key-1', + db=db, + ) + subject_claims = decode_logout_token( + subject_token, verification_key=private_key.public_key(), issuer=_ISSUER, audience='portal' + ) + assert subject_claims['sub'] == 'subject-1' and 'sid' not in subject_claims + assert subject_claims['jti'] == 'event-constant' + + +@pytest.mark.asyncio +@pytest.mark.parametrize('address', ['10.0.0.9', '::1', '169.254.169.254']) +async def test_backchannel_dns_private_and_ipv6_addresses_are_rejected( + monkeypatch: pytest.MonkeyPatch, address: str +) -> None: + """Back-Channel DNS 解析到内网、IPv6 回环或 metadata 地址时拒绝发送。""" + + monkeypatch.setattr( + 'module_identity.security.uri_validator.socket.getaddrinfo', + lambda *_args, **_kwargs: [(2, 1, 6, '', (address, 443))], + ) + + assert await LogoutService._safe_backchannel_uri('https://client.example/logout') is False + + +@pytest.mark.asyncio +async def test_backchannel_retry_queue_and_structured_audit_exclude_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Back-Channel 失败有限重试、入队和审计均不携带 Logout Token。""" + + attempts = 0 + queued: list[tuple[str, str, str]] = [] + audited: list[tuple[str, str, str, str]] = [] + + async def fail(_uri: str, _token: str) -> None: + nonlocal attempts + attempts += 1 + raise RuntimeError('network down') + + async def queue( + uri: str, + client_id: str, + sid: str, + *, + event_jti: str | None = None, + subject_id: str | None = None, + include_sid: bool = True, + ) -> None: + del event_jti, subject_id, include_sid + queued.append((uri, client_id, sid)) + + async def audit(event: str, uri: str, client_id: str, sid: str, *, failure_code: str | None = None) -> None: + del failure_code + audited.append((event, uri, client_id, sid)) + + monkeypatch.setattr('module_identity.service.session_service._BACKCHANNEL_RETRY_DELAY_SECONDS', 0) + callback = LogoutService._notification_callback( + 'https://client.example/logout', + 'logout.secret.must.not.persist', + fail, + client_id='portal', + sid='sid-1', + retry_queue=queue, + audit_writer=audit, + ) + await callback() + + assert attempts == _MAX_ATTEMPTS + assert queued == [('https://client.example/logout', 'portal', 'sid-1')] + assert audited == [('backchannel_logout_failed', 'https://client.example/logout', 'portal', 'sid-1')] + assert all('logout.secret' not in repr(item) for item in (*queued, *audited)) + + +@pytest.mark.asyncio +async def test_backchannel_retry_metadata_reuses_event_jti_without_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """延迟任务只保存 event_jti 等元数据,重试不得生成新事件 ID。""" + metadata: list[dict[str, object]] = [] + + async def fail(_uri: str, _token: str) -> None: + raise RuntimeError('network down') + + async def queue( + _uri: str, + _client_id: str, + _sid: str, + **kwargs: object, + ) -> None: + metadata.append(kwargs) + + monkeypatch.setattr('module_identity.service.session_service._BACKCHANNEL_RETRY_DELAY_SECONDS', 0) + callback = LogoutService._notification_callback( + 'https://client.example/logout', + 'logout.secret.must.not.persist', + fail, + client_id='portal', + sid='sid-1', + event_jti='event-constant', + retry_queue=queue, + ) + await callback() + + assert metadata == [{'event_jti': 'event-constant', 'subject_id': None, 'include_sid': True}] + assert all('token' not in item for item in metadata[0]) + + +@pytest.mark.asyncio +async def test_backchannel_permanent_http_4xx_is_not_retried(monkeypatch: pytest.MonkeyPatch) -> None: + """HTTP 4xx(408/429 除外)直接失败审计,不进入重试队列。""" + attempts = 0 + audited: list[str] = [] + + async def permanent(_uri: str, _token: str) -> None: + nonlocal attempts + attempts += 1 + request = httpx.Request('POST', 'https://client.example/logout') + raise httpx.HTTPStatusError('bad request', request=request, response=httpx.Response(400, request=request)) + + async def audit(event: str, _uri: str, _client_id: str, _sid: str, **kwargs: object) -> None: + audited.append(f'{event}:{kwargs.get("failure_code")}') + + monkeypatch.setattr('module_identity.service.session_service._BACKCHANNEL_RETRY_DELAY_SECONDS', 0) + callback = LogoutService._notification_callback( + 'https://client.example/logout', + 'logout-token', + permanent, + client_id='portal', + sid='sid-1', + audit_writer=audit, + ) + await callback() + + assert attempts == 1 + assert audited == ['backchannel_logout_failed:http_400'] + + +@pytest.mark.asyncio +async def test_malformed_retry_payload_goes_to_safe_dead_letter_and_audit( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """损坏队列数据不静默丢弃,也不把原始正文写入 dead-letter。""" + redis = FakeRedis() + await redis.rpush(OidcRedisKey.backchannel_retry_queue(), b'{not-json') + audited: list[str] = [] + + async def audit(event: str, _uri: str, _client_id: str, _sid: str, **kwargs: object) -> None: + audited.append(f'{event}:{kwargs.get("failure_code")}') + + monkeypatch.setattr(LogoutService, '_audit_writer', lambda _db: audit) + assert await LogoutService.consume_backchannel_retry(object(), redis) == 1 + dead = await redis.lpop(OidcRedisKey.backchannel_retry_queue() + ':dead') + assert dead == '{"error":"invalid_retry_payload"}' + assert audited == ['backchannel_logout_failed:invalid retry payload'] + + +@pytest.mark.asyncio +async def test_pinned_transport_uses_approved_ip_and_preserves_origin_host(monkeypatch: pytest.MonkeyPatch) -> None: + """Pinned transport 只拨批准 IP,同时保留注册域名的 Host 和 TLS SNI。""" + + class NetworkStream: + def __init__(self) -> None: + self.writes: list[bytes] = [] + self.server_hostname: str | None = None + self.response_sent = False + + async def read(self, _max_bytes: int, timeout: float | None = None) -> bytes: + if self.response_sent: + return b'' + self.response_sent = True + return b'HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nOK' + + async def write(self, buffer: bytes, timeout: float | None = None) -> None: + self.writes.append(buffer) + + async def start_tls( + self, + ssl_context: object, + server_hostname: str | None = None, + timeout: float | None = None, + ) -> 'NetworkStream': + self.server_hostname = server_hostname + return self + + async def aclose(self) -> None: + return None + + def get_extra_info(self, _info: str) -> None: + return None + + stream = NetworkStream() + transport = PinnedHttpxTransport('client.example', {'203.0.113.7'}) + backend = transport._pool._network_backend + connect = AsyncMock(return_value=stream) + monkeypatch.setattr(backend._backend, 'connect_tcp', connect) + + async with httpx.AsyncClient(transport=transport) as client: + response = await client.post('https://client.example/logout', data={'logout_token': 'opaque'}) + + assert response.content == b'OK' + assert connect.await_args.args[:2] == ('203.0.113.7', 443) + assert stream.server_hostname == 'client.example' + assert b'host: client.example\r\n' in b''.join(stream.writes).lower() + + +@pytest.mark.asyncio +async def test_pinned_transport_stops_oversized_response_stream() -> None: + """远端超大响应在流读取时主动中止,不在内存中无限累积。""" + + class Stream: + async def __aiter__(self) -> Any: + yield b'x' * (70 * 1024) + + async def aclose(self) -> None: + return None + + response = SimpleNamespace(status=200, headers=[], stream=Stream(), extensions={}, aclose=AsyncMock()) + transport = PinnedHttpxTransport('client.example', {'203.0.113.7'}) + transport._pool.handle_async_request = AsyncMock(return_value=response) + result = await transport.handle_async_request(httpx.Request('POST', 'https://client.example/logout')) + with pytest.raises(OSError, match='响应大小超过限制'): + async for _chunk in result.aiter_bytes(): + pass + response.aclose.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_backchannel_retry_consumer_revalidates_uri_and_reissues_short_lived_token( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """重试消费者不保存 Token,重验注册 URI 后重新签发并成功出队。""" + redis = FakeRedis() + queue_key = 'oidc:backchannel:retry' + await redis.rpush( + queue_key, + '{"uri":"https://client.example/logout","client_id":"portal","sid":"sid-1",' + '"event_jti":"event-1","include_sid":true,"attempt":1}', + ) + client = SimpleNamespace(client_pk=20, client_id='portal', status='0') + registered = SimpleNamespace(uri='https://client.example/logout', status='0') + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.get_by_client_id', _async_return(client) + ) + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.find_exact_uri', _async_return(registered) + ) + monkeypatch.setattr(LogoutService, '_safe_backchannel_uri', _async_return(True)) + monkeypatch.setattr(LogoutService, '_make_logout_token', _async_return('fresh.logout.token')) + sent: list[tuple[str, str]] = [] + + async def notify(uri: str, token: str) -> None: + sent.append((uri, token)) + + result = await LogoutService.consume_backchannel_retry(object(), redis, notifier=notify) + + assert result == 1 + assert sent == [('https://client.example/logout', 'fresh.logout.token')] + assert await redis.lpop(queue_key) is None + + +@pytest.mark.asyncio +async def test_logout_endpoint_disabled_is_local_404_without_legacy_pre_auth( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """OIDC 关闭时直接返回裸 404,不执行数据库或 Legacy PreAuth。""" + + config = SimpleNamespace(oidc_enabled=False) + monkeypatch.setattr(controller, 'OidcConfig', config) + response = await controller.logout(_request(), object()) + + assert response.status_code == _HTTP_NOT_FOUND + + +@pytest.mark.asyncio +async def test_logout_endpoint_prepares_confirmation_without_clearing_sso_cookie( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """初次退出只建立确认,不清理 SSO Cookie 或撤销登录。""" + + config = SimpleNamespace( + oidc_enabled=True, + oidc_sso_cookie_name='__Host-ruoyi-sso', + oidc_sso_cookie_secure=True, + oidc_sso_cookie_domain=None, + oidc_sso_cookie_samesite='lax', + ) + monkeypatch.setattr(controller, 'OidcConfig', config) + + async def fake_logout(*_args: object, **_kwargs: object) -> LogoutResult: + return LogoutResult('https://portal.example/logged-out', 'state-1', True) + + monkeypatch.setattr(controller.LogoutService, 'execute_logout', fake_logout) + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + monkeypatch.setattr(controller.LogoutConfirmationService, 'issue', AsyncMock(return_value=('t' * 43, 'n' * 43))) + execute = AsyncMock() + monkeypatch.setattr(controller.LogoutService, 'execute_logout', execute) + response = await controller.logout(_request(), db) + + assert response.status_code == _HTTP_OK + assert controller.LogoutConfirmationService.COOKIE_NAME in response.headers['set-cookie'] + assert '__Host-ruoyi-sso=' not in response.headers['set-cookie'] + assert '确认退出'.encode() in response.body + execute.assert_not_awaited() + assert OidcUtil.append_state('https://portal.example/logged-out?state=old&next=1', 'state-1') == ( + 'https://portal.example/logged-out?next=1&state=state-1' + ) + + +@pytest.mark.asyncio +async def test_logout_endpoint_unexpected_service_error_returns_503_without_clearing_cookie( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Logout Service 异常时返回 503,避免清理仍有效的认证 Cookie。""" + + config = SimpleNamespace( + oidc_enabled=True, + oidc_sso_cookie_name='__Host-ruoyi-sso', + oidc_sso_cookie_secure=True, + oidc_sso_cookie_domain=None, + oidc_sso_cookie_samesite='lax', + ) + monkeypatch.setattr(controller, 'OidcConfig', config) + + async def fail(*_args: object, **_kwargs: object) -> LogoutResult: + raise RuntimeError('unexpected failure') + + monkeypatch.setattr(controller.LogoutConfirmationService, 'issue', fail) + request = _request(headers=[(b'cookie', b'__Host-ruoyi-sso=valid-cookie')]) + response = await controller.logout(request, SimpleNamespace()) + + assert response.status_code == _HTTP_SERVICE_UNAVAILABLE + assert 'set-cookie' not in response.headers + + +@pytest.mark.asyncio +async def test_logout_endpoint_post_form_forwards_logout_parameters( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """POST application/x-www-form-urlencoded 正确解析并转发 Logout 参数。""" + + config = SimpleNamespace( + oidc_enabled=True, + oidc_sso_cookie_name='__Host-ruoyi-sso', + oidc_sso_cookie_secure=True, + oidc_sso_cookie_domain=None, + oidc_sso_cookie_samesite='lax', + ) + monkeypatch.setattr(controller, 'OidcConfig', config) + monkeypatch.setattr('module_identity.dependencies.OidcConfig', config) + captured: dict[str, object] = {} + + async def prepare(_redis: object, parameters: dict[str, str], _cookie: str | None) -> tuple[str, str]: + captured.update(parameters) + return 't' * 43, 'n' * 43 + + monkeypatch.setattr(controller.LogoutConfirmationService, 'issue', prepare) + body = b'id_token_hint=hint-value&post_logout_redirect_uri=https%3A%2F%2Fportal.example%2Fdone&state=state-1' + request = _request( + method='POST', + headers=[ + (b'content-type', b'application/x-www-form-urlencoded'), + (b'content-length', str(len(body)).encode()), + ], + body=body, + ) + + response = await controller.logout(request, SimpleNamespace()) + + assert response.status_code == _HTTP_OK + assert captured['id_token_hint'] == 'hint-value' + assert captured['post_logout_redirect_uri'] == 'https://portal.example/done' + assert captured['state'] == 'state-1' + + +def _request( + *, + method: str = 'GET', + headers: list[tuple[bytes, bytes]] | None = None, + body: bytes = b'', +) -> Request: + """构造带最小 Redis 状态的 Starlette Request。""" + + app = SimpleNamespace(state=SimpleNamespace(redis=object())) + + async def receive() -> dict[str, object]: + return {'type': 'http.request', 'body': body, 'more_body': False} + + return Request( + { + 'type': 'http', + 'method': method, + 'path': '/oauth2/logout', + 'headers': headers or [], + 'query_string': b'', + 'app': app, + }, + receive, + ) + + +def _async_return(value: object) -> object: + """构造固定值异步桩。""" + + async def return_value(*_args: object, **_kwargs: object) -> object: + return value + + return return_value + + +async def _async_noop(*_args: object, **_kwargs: object) -> None: + """构造无副作用异步桩。""" + + +@pytest.mark.asyncio +async def test_expired_id_token_is_accepted_only_as_logout_hint(monkeypatch: pytest.MonkeyPatch, now: datetime) -> None: + private_key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + payload = jwt.decode(_id_token(private_key, now=now), options={'verify_signature': False}) + payload.update(iat=int(now.timestamp()) - 1200, exp=int(now.timestamp()) - 600) + token = encode_id_token(payload, private_key, 'key-1') + monkeypatch.setattr( + 'module_identity.service.session_service.OAuthClientDao.get_by_client_id', + _async_return(SimpleNamespace(client_pk=20, client_id='portal', status='0')), + ) + monkeypatch.setattr( + 'module_identity.service.session_service.OidcKeyDao.get_verifying', + _async_return(_key_record(private_key, now)), + ) + claims, _ = await LogoutService._validate_id_token_hint(object(), token, now) + assert claims['sid'] == 'sid-1' + with pytest.raises(JwtProfileError): + decode_id_token(token, verification_key=private_key.public_key(), issuer=_ISSUER, audience='portal') + + +@pytest.mark.asyncio +async def test_confirmed_logout_does_not_revoke_another_accounts_hint_session( + monkeypatch: pytest.MonkeyPatch, now: datetime +) -> None: + client = SimpleNamespace(client_pk=20, client_id='portal', status='0') + other = SimpleNamespace(sid='other-sid', subject_id='other-subject', status='active') + current = SimpleNamespace(sid='current-sid', subject_id='current-subject', status='active') + monkeypatch.setattr( + LogoutService, '_validate_id_token_hint', _async_return(({'sid': other.sid, 'sub': other.subject_id}, client)) + ) + monkeypatch.setattr('module_identity.service.session_service.SsoSessionDao.get_by_sid', _async_return(other)) + monkeypatch.setattr( + 'module_identity.service.session_service.SsoSessionService.validate_logout_cookie', _async_return(current) + ) + monkeypatch.setattr(LogoutService, '_lock_session_refresh_tokens', _async_return([])) + revoked = AsyncMock() + monkeypatch.setattr(LogoutService, '_revoke_session_state', revoked) + monkeypatch.setattr(LogoutService, '_register_backchannel', _async_noop) + monkeypatch.setattr( + 'module_identity.service.session_service.SsoSessionDao.client_ids_for_sid', AsyncMock(return_value=[]) + ) + result = await LogoutService._logout( + object(), + object(), + id_token_hint='signed-other-account', + cookie='current-cookie', + confirmed=True, + coordinator=AfterCommitCoordinator(), + now=now, + ) + assert revoked.await_args.args[2] == 'current-sid' + assert result.redirect_uri is None diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_oauth_management_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_oauth_management_service.py new file mode 100644 index 000000000..03e4edd81 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_oauth_management_service.py @@ -0,0 +1,576 @@ +from datetime import datetime, timedelta, timezone + +import pytest +import pytest_asyncio +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from config.database import Base +from config.env import OidcConfig +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.do.oauth_client_do import ( + SysOAuthClient, + SysOAuthClientSecret, + SysOAuthClientUri, +) +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken, SysSsoSession +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.oauth_client_vo import ( + ClientCreateModel, + ClientPageQueryModel, + ClientStatusModel, + ClientUpdateModel, + ClientUriModel, +) +from module_identity.security.client_auth import verify_client_secret +from module_identity.service.oauth_management_service import ( + OAuthClientManagementError, + OAuthClientManagementService, + OAuthManagementBaseService, +) + +_EXPECTED_PAGE_TOTAL = 3 + + +@pytest_asyncio.fixture +async def management_session() -> AsyncSession: + """创建只包含 OAuth 管理表的真实 SQLite 异步会话。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + tables = [ + SysOAuthClient.__table__, + SysOAuthResource.__table__, + SysOAuthScope.__table__, + SysOAuthClientSecret.__table__, + SysOAuthClientUri.__table__, + SysOAuthClientScope.__table__, + SysOAuthClientResource.__table__, + SysOAuthGrant.__table__, + SysSsoSession.__table__, + SysOAuthRefreshToken.__table__, + SysOAuthAuditLog.__table__, + ] + async with engine.begin() as connection: + await connection.run_sync(lambda sync_connection: Base.metadata.create_all(sync_connection, tables=tables)) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as session: + yield session + await engine.dispose() + + +async def _seed_definitions(session: AsyncSession) -> None: + """写入用于绑定校验的启用 Resource 和 Scope。""" + resource = SysOAuthResource( + resource_id='resource-a', + resource_name='Resource A', + audience='urn:test:a', + allowed_claims=['sub', 'email'], + status='0', + create_by='tester', + update_by='tester', + ) + session.add(resource) + await session.flush() + session.add_all( + [ + SysOAuthScope( + scope_code='openid', + scope_name='OpenID', + scope_type='identity', + resource_pk=None, + claims=['sub'], + status='0', + create_by='tester', + update_by='tester', + ), + SysOAuthScope( + scope_code='resource.read', + scope_name='Read', + scope_type='resource', + resource_pk=resource.resource_pk, + claims=['email'], + status='0', + create_by='tester', + update_by='tester', + ), + ] + ) + await session.flush() + await session.commit() + + +def _confidential_payload(**changes: object) -> ClientCreateModel: + """构造一个需要授权码、PKCE 和回调地址的机密 Client。""" + values: dict[str, object] = { + 'client_name': '管理测试 Client', + 'client_type': 'confidential', + 'token_endpoint_auth_method': 'client_secret_basic', + 'grant_types': ['authorization_code', 'refresh_token'], + 'response_types': ['code'], + 'require_pkce': True, + 'scope_codes': ['openid', 'resource.read'], + 'pre_authorized_scope_codes': ['openid'], + 'resource_ids': ['resource-a'], + 'redirect_uris': ['https://client.example/callback'], + 'post_logout_redirect_uris': ['https://client.example/logout'], + 'cors_origins': ['https://client.example'], + } + values.update(changes) + return ClientCreateModel.model_validate(values) + + +@pytest.mark.asyncio +async def test_after_commit_failure_does_not_rollback_committed_management_change() -> None: + """提交后的运行时回调失败不得回滚已提交的管理事务。""" + + class _Db: + commits = 0 + rollbacks = 0 + + async def commit(self) -> None: + self.commits += 1 + + async def rollback(self) -> None: + self.rollbacks += 1 + + db = _Db() + + async def operation() -> str: + return 'done' + + async def after_commit() -> None: + raise RuntimeError('runtime refresh failed') + + with pytest.raises(RuntimeError, match='runtime refresh failed'): + await OAuthManagementBaseService._transaction(db, operation, after_commit) + assert db.commits == 1 + assert db.rollbacks == 0 + + +@pytest.mark.asyncio +async def test_create_detail_bindings_secret_and_cors_snapshot(management_session: AsyncSession) -> None: + """创建应原子写入绑定,创建不自动发放 Secret,轮换时只返回一次明文。""" + await _seed_definitions(management_session) + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + detail = await OAuthClientManagementService.create_client( + management_session, _confidential_payload(), actor='admin', now=now + ) + secret = await OAuthClientManagementService.rotate_secret( + management_session, detail.client_id, actor='admin', now=now + ) + assert secret.client_secret.startswith('cs1.') + assert detail.client_id.startswith('cli_') + assert detail.scope_codes == ['openid', 'resource.read'] + assert detail.pre_authorized_scope_codes == ['openid'] + assert detail.resource_ids == ['resource-a'] + assert verify_client_secret( + secret.client_secret, (await management_session.get(SysOAuthClientSecret, secret.secret_id)).secret_hash + ) + assert secret.client_secret not in repr(detail) + assert await OAuthClientManagementService.list_active_cors_origins(management_session) == ( + 'https://client.example', + ) + + +@pytest.mark.asyncio +async def test_update_replaces_bindings_and_increments_policy(management_session: AsyncSession) -> None: + """更新策略应锁定 Client、替换旧绑定并单调增加 policy_version。""" + await _seed_definitions(management_session) + detail = await OAuthClientManagementService.create_client(management_session, _confidential_payload(), actor='a') + update_values = _confidential_payload( + client_name='更新后的 Client', + pre_authorized_scope_codes=['openid', 'resource.read'], + cors_origins=['https://new-client.example'], + ).model_dump() + update_values['client_id'] = detail.client_id + update = ClientUpdateModel.model_validate(update_values) + changed = await OAuthClientManagementService.update_client(management_session, update, actor='b') + assert changed.client_name == '更新后的 Client' + assert changed.policy_version == detail.policy_version + 1 + assert changed.cors_origins == ['https://new-client.example'] + assert ( + (await management_session.execute(select(SysOAuthClientUri).where(SysOAuthClientUri.client_pk == 1))) + .scalars() + .all() + ) + + +@pytest.mark.asyncio +async def test_status_soft_disable_and_secret_revoke_are_safe(management_session: AsyncSession) -> None: + """状态变更、跨 Client 撤销和重复撤销都不泄漏 Secret。""" + await _seed_definitions(management_session) + first = await OAuthClientManagementService.create_client(management_session, _confidential_payload(), actor='a') + second = await OAuthClientManagementService.create_client( + management_session, _confidential_payload(client_name='第二个'), actor='a' + ) + first_secret = await OAuthClientManagementService.rotate_secret(management_session, first.client_id, actor='a') + second_secret = await OAuthClientManagementService.rotate_secret(management_session, second.client_id, actor='a') + with pytest.raises(OAuthClientManagementError, match='不属于当前客户端'): + await OAuthClientManagementService.revoke_secret( + management_session, first.client_id, second_secret.secret_id, actor='admin' + ) + assert await OAuthClientManagementService.revoke_secret( + management_session, first.client_id, first_secret.secret_id, actor='admin' + ) + assert not await OAuthClientManagementService.revoke_secret( + management_session, first.client_id, first_secret.secret_id, actor='admin' + ) + second_row = ( + await management_session.execute(select(SysOAuthClient).where(SysOAuthClient.client_id == second.client_id)) + ).scalar_one() + management_session.add( + SysOAuthGrant( + grant_id='grant-disable', + user_id=7, + subject_id='subject-disable', + client_pk=second_row.client_pk, + granted_scopes=[], + granted_resources=[], + client_policy_version=second.policy_version, + status='active', + ) + ) + management_session.add( + SysSsoSession( + sid='sid-disable', + session_secret_hash='a' * 64, + user_id=7, + subject_id='subject-disable', + auth_version=1, + auth_time=datetime.now(timezone.utc), + last_seen_at=datetime.now(timezone.utc), + idle_expires_at=datetime.now(timezone.utc) + timedelta(hours=1), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + acr='urn:test', + amr=['pwd'], + status='active', + ) + ) + management_session.add( + SysOAuthRefreshToken( + token_id='refresh-disable', + token_hash='b' * 64, + family_id='family-disable', + grant_id='grant-disable', + user_id=7, + subject_id='subject-disable', + auth_version=1, + client_pk=second_row.client_pk, + sid='sid-disable', + scopes=[], + resources=[], + status='active', + issued_at=datetime.now(timezone.utc), + idle_expires_at=datetime.now(timezone.utc) + timedelta(hours=1), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + ) + ) + await management_session.flush() + disabled = await OAuthClientManagementService.change_client_status( + management_session, ClientStatusModel(client_id=second.client_id, status='1'), actor='admin' + ) + assert disabled.status == '1' + assert (await management_session.get(SysOAuthGrant, 'grant-disable')).status == 'revoked' + assert (await management_session.get(SysOAuthRefreshToken, 'refresh-disable')).status == 'revoked' + assert (await management_session.get(SysSsoSession, 'sid-disable')).status == 'active' + disabled_again = await OAuthClientManagementService.change_client_status( + management_session, ClientStatusModel(client_id=second.client_id, status='1'), actor='admin' + ) + assert disabled_again.status == '1' + + +@pytest.mark.asyncio +async def test_client_ttl_limits_and_scope_default_are_fail_closed( + management_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + """Client TTL 受平台上限约束,允许 Scope 不会被误标为默认 Scope。""" + await _seed_definitions(management_session) + monkeypatch.setattr(OidcConfig, 'oidc_max_access_token_ttl_seconds', 60) + monkeypatch.setattr(OidcConfig, 'oidc_access_token_ttl_seconds', 30) + monkeypatch.setattr(OidcConfig, 'oidc_refresh_token_idle_seconds', 100) + monkeypatch.setattr(OidcConfig, 'oidc_refresh_token_absolute_seconds', 200) + with pytest.raises(OAuthClientManagementError, match='访问令牌有效期'): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(access_token_ttl_seconds=61), + actor='admin', + ) + with pytest.raises(OAuthClientManagementError, match='闲置有效期超过平台上限'): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(refresh_token_idle_seconds=150, refresh_token_absolute_seconds=200), + actor='admin', + ) + with pytest.raises(OAuthClientManagementError, match='闲置有效期不能超过绝对有效期'): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(refresh_token_idle_seconds=100, refresh_token_absolute_seconds=50), + actor='admin', + ) + with pytest.raises(OAuthClientManagementError, match='绝对有效期超过平台上限'): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(refresh_token_idle_seconds=100, refresh_token_absolute_seconds=201), + actor='admin', + ) + await OAuthClientManagementService.create_client(management_session, _confidential_payload(), actor='admin') + defaults = ( + (await management_session.execute(select(SysOAuthClientScope).where(SysOAuthClientScope.client_pk == 1))) + .scalars() + .all() + ) + assert defaults and all(row.is_default == 0 for row in defaults) + + +@pytest.mark.asyncio +async def test_client_count_matches_filter_and_ignores_pagination(management_session: AsyncSession) -> None: + """Client 分页的 total 必须是过滤后的全量数量,而不是当前页数量。""" + await _seed_definitions(management_session) + for index in range(3): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(client_name=f'分页计数 Client {index}'), + actor='admin', + ) + query = ClientPageQueryModel(client_name='分页计数', page_num=2, page_size=2) + rows = await OAuthClientManagementService.list_clients(management_session, query) + total = await OAuthClientManagementService.count_clients(management_session, query) + assert len(rows) == 1 + assert total == _EXPECTED_PAGE_TOTAL + + +@pytest.mark.asyncio +async def test_secret_rotation_bounds_old_active_secret(management_session: AsyncSession) -> None: + """轮换应将旧 Active Secret 置为 retiring 并设置明确退役窗口。""" + await _seed_definitions(management_session) + detail = await OAuthClientManagementService.create_client( + management_session, _confidential_payload(), actor='admin' + ) + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + first = await OAuthClientManagementService.rotate_secret( + management_session, detail.client_id, actor='admin', now=now, retirement_seconds=120 + ) + future = now + timedelta(days=2) + second = await OAuthClientManagementService.rotate_secret( + management_session, + detail.client_id, + actor='admin', + now=now, + not_before=future, + retirement_seconds=120, + ) + old = await management_session.get(SysOAuthClientSecret, first.secret_id) + assert old is not None and old.status == 'retiring' + assert old.expires_at == future + timedelta(seconds=120) + new = await management_session.get(SysOAuthClientSecret, second.secret_id) + assert new is not None + assert new.create_time == now + assert new.not_before == future + assert second.secret_hint == f'...{second.client_secret[-6:]}' + secret_tail = second.client_secret[-6:] + assert second.secret_hint.startswith('...') + assert len(second.secret_hint) <= len('...') + len(secret_tail) + + +@pytest.mark.asyncio +async def test_secret_rotation_respects_original_expiry_and_rejects_a_gap( + management_session: AsyncSession, +) -> None: + """轮换不延长原过期时间,远期过期会被截断,排期断档则原子拒绝。""" + await _seed_definitions(management_session) + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + + early_client = await OAuthClientManagementService.create_client( + management_session, _confidential_payload(client_name='原过期较早'), actor='admin' + ) + early_secret = await OAuthClientManagementService.rotate_secret( + management_session, + early_client.client_id, + actor='admin', + now=now, + expires_at=now + timedelta(seconds=60), + ) + await OAuthClientManagementService.rotate_secret( + management_session, + early_client.client_id, + actor='admin', + now=now, + not_before=now + timedelta(seconds=30), + retirement_seconds=120, + ) + early_old = await management_session.get(SysOAuthClientSecret, early_secret.secret_id) + assert early_old is not None and early_old.status == 'retiring' + assert early_old.expires_at == now + timedelta(seconds=60) + + far_client = await OAuthClientManagementService.create_client( + management_session, _confidential_payload(client_name='原过期较远'), actor='admin' + ) + far_secret = await OAuthClientManagementService.rotate_secret( + management_session, + far_client.client_id, + actor='admin', + now=now, + expires_at=now + timedelta(seconds=1000), + ) + await OAuthClientManagementService.rotate_secret( + management_session, far_client.client_id, actor='admin', now=now, retirement_seconds=120 + ) + far_old = await management_session.get(SysOAuthClientSecret, far_secret.secret_id) + assert far_old is not None and far_old.status == 'retiring' + assert far_old.expires_at == now + timedelta(seconds=120) + + reject_client = await OAuthClientManagementService.create_client( + management_session, _confidential_payload(client_name='排期断档'), actor='admin' + ) + reject_secret = await OAuthClientManagementService.rotate_secret( + management_session, + reject_client.client_id, + actor='admin', + now=now, + expires_at=now + timedelta(seconds=60), + ) + with pytest.raises(OAuthClientManagementError, match='新旧密钥的重叠有效期'): + await OAuthClientManagementService.rotate_secret( + management_session, + reject_client.client_id, + actor='admin', + now=now, + not_before=now + timedelta(seconds=60), + retirement_seconds=120, + ) + reject_old = await management_session.get(SysOAuthClientSecret, reject_secret.secret_id) + assert reject_old is not None and reject_old.status == 'active' + assert reject_old.expires_at == now + timedelta(seconds=60) + reject_rows = ( + ( + await management_session.execute( + select(SysOAuthClientSecret).where(SysOAuthClientSecret.client_pk == reject_old.client_pk) + ) + ) + .scalars() + .all() + ) + assert len(reject_rows) == 1 + + +@pytest.mark.asyncio +async def test_uri_add_and_soft_remove_increment_policy(management_session: AsyncSession) -> None: + """URI 增删应锁定 Client、递增版本且保留停用历史行。""" + await _seed_definitions(management_session) + detail = await OAuthClientManagementService.create_client(management_session, _confidential_payload(), actor='a') + uri_id = await OAuthClientManagementService.add_uri( + management_session, + detail.client_id, + ClientUriModel(uri_type='redirect', uri='https://client.example/second-callback'), + actor='a', + ) + added = await OAuthClientManagementService.detail(management_session, detail.client_id) + assert added.policy_version == detail.policy_version + 1 + assert 'https://client.example/second-callback' in added.redirect_uris + assert await OAuthClientManagementService.remove_uri(management_session, detail.client_id, uri_id, actor='a') + removed = await OAuthClientManagementService.detail(management_session, detail.client_id) + assert removed.policy_version == added.policy_version + 1 + assert 'https://client.example/second-callback' not in removed.redirect_uris + historical = await management_session.get(SysOAuthClientUri, uri_id) + assert historical is not None and historical.status == '1' + + +@pytest.mark.asyncio +async def test_public_client_has_no_secret_and_uri_policy_is_fail_closed(management_session: AsyncSession) -> None: + """Public Client 禁止轮换密钥,危险 URI 和错误绑定均在服务边界拒绝。""" + public = ClientCreateModel.model_validate( + { + 'client_name': '公共 Client', + 'client_type': 'public', + 'token_endpoint_auth_method': 'none', + 'grant_types': ['authorization_code'], + 'response_types': ['code'], + 'scope_codes': [], + 'redirect_uris': ['https://public.example/callback'], + } + ) + detail = await OAuthClientManagementService.create_client(management_session, public, actor='admin') + with pytest.raises(OAuthClientManagementError, match='公开客户端'): + await OAuthClientManagementService.rotate_secret(management_session, detail.client_id, actor='admin') + await _seed_definitions(management_session) + with pytest.raises(OAuthClientManagementError, match='保留参数'): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(redirect_uris=['https://client.example/callback?state=bad']), + actor='admin', + ) + with pytest.raises(OAuthClientManagementError, match='资源权限'): + await OAuthClientManagementService.create_client( + management_session, + _confidential_payload(resource_ids=[]), + actor='admin', + ) + + +@pytest.mark.asyncio +async def test_create_failure_can_be_rolled_back_by_caller(management_session: AsyncSession) -> None: + """公共写入口提交事务,调用方后续回滚不会撤销已提交变更。""" + await _seed_definitions(management_session) + await OAuthClientManagementService.create_client(management_session, _confidential_payload(), actor='admin') + await management_session.rollback() + result = await management_session.execute(select(SysOAuthClient)) + assert len(result.scalars().all()) == 1 + + +@pytest.mark.asyncio +async def test_role_allowlist_is_persisted_and_changes_client_policy(management_session: AsyncSession) -> None: + """仅发布白名单中的角色,修改白名单后旧策略授权失效。""" + await _seed_definitions(management_session) + management_session.add( + SysOAuthScope( + scope_code='roles', + scope_name='Roles', + scope_type='identity', + claims=['roles'], + status='0', + create_by='tester', + update_by='tester', + ) + ) + await management_session.commit() + payload = _confidential_payload( + scope_codes=['openid', 'roles'], + resource_ids=[], + allowed_role_keys=['analyst'], + ) + created = await OAuthClientManagementService.create_client(management_session, payload, actor='admin') + assert created.allowed_role_keys == ['analyst'] + values = payload.model_dump() + values.update(client_id=created.client_id, allowed_role_keys=['reader']) + changed = await OAuthClientManagementService.update_client( + management_session, + ClientUpdateModel.model_validate(values), + actor='admin', + ) + assert changed.policy_version == created.policy_version + 1 + loaded = await OAuthClientManagementService.detail(management_session, created.client_id) + assert loaded.allowed_role_keys == ['reader'] + filters = ( + ( + await management_session.execute( + select(SysOAuthClientScope.claim_filter) + .join(SysOAuthScope, SysOAuthScope.scope_pk == SysOAuthClientScope.scope_pk) + .where(SysOAuthScope.scope_code == 'roles') + ) + ) + .scalars() + .all() + ) + assert filters == [{'claims': ['roles'], 'allowed_role_keys': ['reader']}] + values['allowed_role_keys'] = [] + cleared = await OAuthClientManagementService.update_client( + management_session, + ClientUpdateModel.model_validate(values), + actor='admin', + ) + assert cleared.allowed_role_keys == [] + assert cleared.policy_version == changed.policy_version + 1 diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_oauth_resource_scope_management.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_oauth_resource_scope_management.py new file mode 100644 index 000000000..3b5acb8dd --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_oauth_resource_scope_management.py @@ -0,0 +1,410 @@ +from datetime import datetime, timezone + +import pytest +import pytest_asyncio +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from config.database import Base +from module_identity.entity.do.oauth_audit_do import SysOAuthAuditLog +from module_identity.entity.do.oauth_client_do import ( + SysOAuthClient, + SysOAuthClientSecret, + SysOAuthClientUri, +) +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken, SysSsoSession +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.entity.vo.oauth_client_vo import ClientCreateModel, ClientUpdateModel, ClientUriModel +from module_identity.entity.vo.oauth_resource_vo import ( + ResourceCreateModel, + ResourcePageQueryModel, + ResourceStatusModel, + ResourceUpdateModel, + ScopeModel, + ScopePageQueryModel, + ScopeStatusModel, +) +from module_identity.entity.vo.oidc_key_vo import OidcKeyRotateModel +from module_identity.service.oauth_management_service import ( + OAuthClientManagementError, + OAuthClientManagementService, + OAuthResourceManagementService, +) + +_FIRST_POLICY_CHANGE = 2 +_SECOND_POLICY_CHANGE = 3 +_EXPECTED_PAGE_TOTAL = 3 + + +@pytest_asyncio.fixture +async def resource_session() -> AsyncSession: + """创建只包含 Resource/Scope/Client 管理表的真实 SQLite 会话。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + tables = [ + SysOAuthClient.__table__, + SysOAuthResource.__table__, + SysOAuthScope.__table__, + SysOAuthClientSecret.__table__, + SysOAuthClientUri.__table__, + SysOAuthClientScope.__table__, + SysOAuthClientResource.__table__, + SysOAuthGrant.__table__, + SysSsoSession.__table__, + SysOAuthRefreshToken.__table__, + SysOAuthAuditLog.__table__, + ] + async with engine.begin() as connection: + await connection.run_sync(lambda sync_connection: Base.metadata.create_all(sync_connection, tables=tables)) + factory = async_sessionmaker(engine, expire_on_commit=False) + async with factory() as session: + yield session + await engine.dispose() + + +def _resource_payload(**changes: object) -> ResourceCreateModel: + """构造一个合法 Resource DTO。""" + values: dict[str, object] = { + 'resource_id': 'orders-api', + 'resource_name': '订单 API', + 'audience': 'https://orders-api.example/api', + 'allowed_claims': ['sub', 'scope'], + } + values.update(changes) + return ResourceCreateModel.model_validate(values) + + +@pytest.mark.parametrize('value', ['bad/id', 'bad,code', 'bad%20code', 'bad code', '\\x00bad']) +def test_management_identifiers_are_safe_path_segments(value: str) -> None: + """Resource、Scope 和 Key 标识拒绝路径分隔符、CSV 分隔符和控制字符。""" + with pytest.raises(ValueError): + ResourceCreateModel(resource_id=value, resource_name='x', audience='https://x.example') + with pytest.raises(ValueError): + ScopeModel(scope_code=value, scope_name='x', scope_type='identity') + with pytest.raises(ValueError): + OidcKeyRotateModel(kid=value, publish_at=datetime.now(timezone.utc)) + + +def _client_payload(resource_id: str = 'orders-api', scope_code: str | None = None) -> ClientCreateModel: + """构造授权码机密 Client。""" + return ClientCreateModel.model_validate( + { + 'client_name': 'Resource 管理测试 Client', + 'client_type': 'confidential', + 'token_endpoint_auth_method': 'client_secret_basic', + 'grant_types': ['authorization_code'], + 'response_types': ['code'], + 'scope_codes': [scope_code] if scope_code else [], + 'resource_ids': [resource_id], + 'redirect_uris': ['https://client.example/callback'], + } + ) + + +async def _create_resource(session: AsyncSession, resource_id: str = 'orders-api') -> object: + """创建一个基础启用 Resource。""" + return await OAuthResourceManagementService.create_resource( + session, + _resource_payload(resource_id=resource_id, audience=f'https://{resource_id}.example/api'), + actor='admin', + ) + + +@pytest.mark.asyncio +async def test_resource_crud_introspection_policy_and_invalidation_snapshot(resource_session: AsyncSession) -> None: + """Resource 更新应校验 introspection Client、递增绑定 Client 策略并提供撤销目标接口。""" + await _create_resource(resource_session) + client = await OAuthClientManagementService.create_client(resource_session, _client_payload(), actor='admin') + update_values = _resource_payload( + allowed_claims=['sub', 'email'], introspection_client_id=client.client_id + ).model_dump() + update_values['status'] = '0' + updated = await OAuthResourceManagementService.update_resource( + resource_session, + ResourceUpdateModel.model_validate(update_values), + actor='admin', + now=datetime(2026, 1, 1, tzinfo=timezone.utc), + ) + assert updated.introspection_client_id == client.client_id + assert updated.allowed_claims == ['sub', 'email'] + stored_client = await resource_session.get(SysOAuthClient, 1) + assert stored_client is not None and stored_client.policy_version == _FIRST_POLICY_CHANGE + resource_session.add( + SysOAuthGrant( + grant_id='grant-1', + user_id=7, + subject_id='subject-1', + client_pk=stored_client.client_pk, + granted_scopes=[], + granted_resources=['orders-api'], + client_policy_version=stored_client.policy_version, + status='active', + ) + ) + await resource_session.flush() + targets = await OAuthResourceManagementService.collect_resource_invalidation_targets(resource_session, 'orders-api') + assert targets.client_ids == (client.client_id,) + assert targets.grant_ids == ('grant-1',) + assert targets.refresh_token_ids == () + disabled = await OAuthResourceManagementService.change_resource_status( + resource_session, ResourceStatusModel(resource_id='orders-api', status='1'), actor='admin' + ) + assert disabled.status == '1' + assert (await resource_session.get(SysOAuthClient, 1)).policy_version == _SECOND_POLICY_CHANGE + + +@pytest.mark.asyncio +async def test_resource_audience_and_introspection_validation_is_fail_closed(resource_session: AsyncSession) -> None: + """Resource audience 必须绝对 HTTPS,introspection Client 必须 active confidential。""" + for audience in ( + 'http://orders.example/api', + 'https://user:pass@orders.example/api', + 'https://orders.example/api#fragment', + 'https:///missing-host', + ): + with pytest.raises(OAuthClientManagementError): + await OAuthResourceManagementService.create_resource( + resource_session, _resource_payload(audience=audience), actor='a' + ) + await _create_resource(resource_session) + public = await OAuthClientManagementService.create_client( + resource_session, + ClientCreateModel.model_validate( + { + 'client_name': 'Public', + 'client_type': 'public', + 'token_endpoint_auth_method': 'none', + 'grant_types': ['authorization_code'], + 'response_types': ['code'], + 'redirect_uris': ['https://public.example/callback'], + } + ), + actor='a', + ) + with pytest.raises(OAuthClientManagementError, match='机密客户端'): + await OAuthResourceManagementService.update_resource( + resource_session, + ResourceUpdateModel.model_validate( + _resource_payload(introspection_client_id=public.client_id).model_dump() + ), + actor='a', + ) + + +@pytest.mark.asyncio +async def test_scope_binding_claim_whitelist_and_builtin_openid(resource_session: AsyncSession) -> None: + """Scope 的 Resource 归属、Claim 白名单和内置 openid 保护必须 fail-closed。""" + await _create_resource(resource_session) + scope = await OAuthResourceManagementService.create_scope( + resource_session, + ScopeModel( + scope_code='orders.read', + scope_name='读取订单', + scope_type='resource', + resource_id='orders-api', + claims=['sub', 'scope'], + ), + actor='admin', + ) + client = await OAuthClientManagementService.create_client( + resource_session, _client_payload(scope_code='orders.read'), actor='admin' + ) + changed = await OAuthResourceManagementService.update_scope( + resource_session, + ScopeModel( + scope_code='orders.read', + scope_name='读取订单', + scope_type='resource', + resource_id='orders-api', + claims=['sub'], + ), + actor='admin', + ) + assert changed.claims == ['sub'] + assert (await resource_session.get(SysOAuthClient, 1)).policy_version == _FIRST_POLICY_CHANGE + with pytest.raises(OAuthClientManagementError, match='不允许发布的声明'): + await OAuthResourceManagementService.create_scope( + resource_session, + ScopeModel(scope_code='bad', scope_name='Bad', scope_type='identity', claims=['user_id']), + actor='admin', + ) + await OAuthResourceManagementService.create_scope( + resource_session, + ScopeModel(scope_code='openid', scope_name='OpenID', scope_type='identity', claims=['sub']), + actor='admin', + ) + with pytest.raises(OAuthClientManagementError, match='不能停用'): + await OAuthResourceManagementService.change_scope_status( + resource_session, ScopeStatusModel(scope_code='openid', status='1'), actor='admin' + ) + with pytest.raises(OAuthClientManagementError, match='不能停用或更改归属'): + await OAuthResourceManagementService.update_scope( + resource_session, + ScopeModel( + scope_code='openid', + scope_name='OpenID', + scope_type='resource', + resource_id='orders-api', + claims=['sub'], + ), + actor='admin', + ) + assert scope.scope_code == 'orders.read' and client.client_id + + +@pytest.mark.asyncio +async def test_display_only_client_resource_scope_updates_do_not_bump_policy_version( + resource_session: AsyncSession, +) -> None: + """仅修改展示字段时,不应使绑定凭据的策略版本失效。""" + await _create_resource(resource_session) + await OAuthResourceManagementService.create_scope( + resource_session, + ScopeModel( + scope_code='orders.read', + scope_name='读取订单', + scope_type='resource', + resource_id='orders-api', + claims=['sub'], + ), + actor='admin', + ) + client = await OAuthClientManagementService.create_client( + resource_session, _client_payload(scope_code='orders.read'), actor='admin' + ) + assert client.policy_version == 1 + + client_update = ClientUpdateModel.model_validate( + { + **_client_payload(scope_code='orders.read').model_dump(), + 'client_id': client.client_id, + 'client_name': '新的展示名称', + 'logo_uri': 'https://client.example/logo.svg', + 'policy_uri': 'https://client.example/policy', + 'tos_uri': 'https://client.example/terms', + 'remark': '展示备注', + } + ) + changed_client = await OAuthClientManagementService.update_client(resource_session, client_update, actor='admin') + assert changed_client.policy_version == 1 + + changed_resource = await OAuthResourceManagementService.update_resource( + resource_session, + ResourceUpdateModel( + **_resource_payload(resource_name='新的 Resource 展示名', remark='展示备注').model_dump(), + status='0', + ), + actor='admin', + ) + assert changed_resource.resource_name == '新的 Resource 展示名' + assert (await resource_session.get(SysOAuthClient, 1)).policy_version == 1 + + changed_scope = await OAuthResourceManagementService.update_scope( + resource_session, + ScopeModel( + scope_code='orders.read', + scope_name='新的 Scope 展示名', + scope_type='resource', + resource_id='orders-api', + claims=['sub'], + remark='展示备注', + ), + actor='admin', + ) + assert changed_scope.scope_name == '新的 Scope 展示名' + assert (await resource_session.get(SysOAuthClient, 1)).policy_version == 1 + + +@pytest.mark.asyncio +async def test_scope_requires_active_resource(resource_session: AsyncSession) -> None: + """Resource Scope 不得绑定未知或停用 Resource。""" + await _create_resource(resource_session) + await OAuthResourceManagementService.change_resource_status( + resource_session, ResourceStatusModel(resource_id='orders-api', status='1'), actor='admin' + ) + with pytest.raises(OAuthClientManagementError, match='已停用'): + await OAuthResourceManagementService.create_scope( + resource_session, + ScopeModel(scope_code='orders.read', scope_name='读取', scope_type='resource', resource_id='orders-api'), + actor='admin', + ) + + +@pytest.mark.asyncio +async def test_resource_scope_count_matches_filter_and_ignores_pagination(resource_session: AsyncSession) -> None: + """Resource/Scope 分页的 total 必须是过滤后的全量数量。""" + await _create_resource(resource_session) + for index in range(2): + await _create_resource(resource_session, resource_id=f'orders-api-{index}') + resource_query = ResourcePageQueryModel(resource_name='订单', page_num=2, page_size=1) + resource_rows = await OAuthResourceManagementService.list_resources(resource_session, resource_query) + resource_total = await OAuthResourceManagementService.count_resources(resource_session, resource_query) + assert len(resource_rows) == 1 + assert resource_total == _EXPECTED_PAGE_TOTAL + + for index in range(3): + await OAuthResourceManagementService.create_scope( + resource_session, + ScopeModel( + scope_code=f'orders.read.{index}', + scope_name=f'读取订单 {index}', + scope_type='resource', + resource_id='orders-api', + claims=['sub'], + ), + actor='admin', + ) + scope_query = ScopePageQueryModel(scope_name='读取订单', page_num=2, page_size=1) + scope_rows = await OAuthResourceManagementService.list_scopes(resource_session, scope_query) + scope_total = await OAuthResourceManagementService.count_scopes(resource_session, scope_query) + assert len(scope_rows) == 1 + assert scope_total == _EXPECTED_PAGE_TOTAL + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + 'uri', + [ + 'http://localhost/backchannel', + 'https://127.0.0.1/backchannel', + 'https://[::1]/backchannel', + 'https://10.0.0.1/backchannel', + 'https://169.254.1.1/backchannel', + 'https://[ff02::1]/backchannel', + 'https://0.0.0.0/backchannel', + ], +) +async def test_backchannel_registration_has_first_layer_ssrf_boundary(resource_session: AsyncSession, uri: str) -> None: + """Backchannel 注册期拒绝非 HTTPS 和特殊 IP,运行时仍会再次解析校验。""" + await _create_resource(resource_session) + client = await OAuthClientManagementService.create_client(resource_session, _client_payload(), actor='admin') + with pytest.raises(OAuthClientManagementError): + await OAuthClientManagementService.add_uri( + resource_session, + client.client_id, + ClientUriModel(uri_type='backchannel_logout', uri=uri), + actor='admin', + ) + + +@pytest.mark.asyncio +async def test_backchannel_registration_rechecks_dns_target( + resource_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + """Backchannel 注册必须执行与发送端一致的 DNS 公网重解析。""" + await _create_resource(resource_session) + client = await OAuthClientManagementService.create_client(resource_session, _client_payload(), actor='admin') + monkeypatch.setattr( + 'module_identity.security.uri_validator.socket.getaddrinfo', + lambda *_args, **_kwargs: [(2, 1, 6, '', ('169.254.169.254', 443))], + ) + with pytest.raises(OAuthClientManagementError): + await OAuthClientManagementService.add_uri( + resource_session, + client.client_id, + ClientUriModel(uri_type='backchannel_logout', uri='https://client.example/logout'), + actor='admin', + ) diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_review_completion.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_review_completion.py new file mode 100644 index 000000000..0ca3fd46f --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_review_completion.py @@ -0,0 +1,459 @@ +import base64 +from datetime import timedelta +from http import HTTPStatus +from types import SimpleNamespace +from unittest.mock import AsyncMock +from urllib.parse import parse_qs, quote_plus, urlencode, urlsplit + +import pytest +from fastapi import APIRouter, FastAPI +from fastapi.routing import APIRoute +from httpx import ASGITransport, AsyncClient +from sqlalchemy import delete, func, select +from sqlalchemy.ext.asyncio import AsyncSession + +from common.constant import OidcAuditEvent +from exceptions.exception import OAuthProtocolException +from exceptions.handle import handle_exception +from module_identity.controller.authorization_controller import authorization_controller +from module_identity.controller.oauth_session_controller import oauth_grant_controller +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysSsoSessionClient +from module_identity.entity.vo.oauth_session_vo import ( + AccessPolicyPageQueryModel, + GrantPageQueryModel, + SessionPageQueryModel, +) +from module_identity.entity.vo.protocol_vo import AuthorizeRequest +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import AuthorizationCodeService, AuthorizationService +from module_identity.service.consent_service import ConsentService +from module_identity.service.infrastructure_service import AfterCommitCoordinator, OidcRateLimiter +from module_identity.service.oauth_management_service import OAuthClientManagementService +from module_identity.service.oauth_session_management_service import OAuthSessionManagementService +from module_identity.service.session_service import LogoutService, SsoSessionService +from module_identity.service.token_protocol_service import IntrospectionService, RevocationService +from module_identity.service.token_service import TokenService +from tests.module_identity.services.test_authentication_regressions import ( + _PEPPER, + _VERIFIER, + _authorize_request, + _introspect, + _user_token, +) +from tests.module_identity.services.test_authentication_regressions import ( + auth_flow as _auth_flow_fixture, +) + +# 复用真实认证夹具,避免不同回归场景使用不一致的授权配置。 +auth_flow = _auth_flow_fixture + + +def _basic(client_id: str, secret: str) -> str: + """生成符合协议编码要求的机密客户端认证头。""" + + return 'Basic ' + base64.b64encode(f'{quote_plus(client_id)}:{quote_plus(secret)}'.encode()).decode() + + +async def _confidential_token(flow: SimpleNamespace) -> tuple: + """通过机密客户端的真实授权码流程签发离线令牌。""" + + flow.app.client_type = 'confidential' + flow.app.token_endpoint_auth_method = 'client_secret_basic' + secret = await OAuthClientManagementService._new_secret(flow.db, flow.app, 'admin', flow.now, not_before=flow.now) + await flow.db.commit() + authorization = _basic(flow.app.client_id, secret.client_secret) + request = AuthorizeRequest(**_authorize_request(flow, offline=True)) + context = await AuthorizationService.validate_request(flow.db, request) + consent = await ConsentService.submit_consent( + flow.db, + context, + True, + context.scopes, + True, + user_id=flow.user.user_id, + subject_id=flow.subject.subject_id, + ) + code = await AuthorizationCodeService.issue( + flow.redis, + { + 'clientPk': flow.app.client_pk, + 'redirectUri': request.redirect_uri, + 'userId': flow.user.user_id, + 'subjectId': flow.subject.subject_id, + 'authVersion': flow.subject.auth_version, + 'sid': flow.session.sid, + 'grantId': consent.grant.grant_id, + 'scopes': list(consent.scopes), + 'resources': [flow.resource.audience], + 'nonce': request.nonce, + 'codeChallenge': request.code_challenge, + 'codeChallengeMethod': 'S256', + 'authTime': flow.now.isoformat(), + }, + pepper=_PEPPER, + ) + token = await TokenService.issue_token_request( + flow.db, + flow.redis, + { + 'grant_type': 'authorization_code', + 'code': code, + 'redirect_uri': request.redirect_uri, + 'code_verifier': _VERIFIER, + }, + authorization=authorization, + signing_key=flow.signer, + kid='regression-key', + ) + return token, secret, authorization + + +@pytest.mark.asyncio +@pytest.mark.parametrize('stored_status', ['active', 'expired']) +async def test_revoke_user_terminates_expired_offline_sessions(auth_flow: SimpleNamespace, stored_status: str) -> None: + """自然过期不影响离线续期,但管理员按用户撤销必须终止该能力。""" + + flow = auth_flow + token, _ = await _user_token(flow, offline=True) + flow.session.idle_expires_at = flow.now - timedelta(seconds=1) + flow.session.status = stored_status + await flow.db.commit() + assert (await _introspect(flow, token.refresh_token))['active'] is True + count = await OAuthSessionManagementService.revoke_user( + flow.db, flow.redis, flow.user.user_id, 'admin', '终止该用户的全部外部会话' + ) + assert count == 1 + assert flow.session.status == 'revoked' + assert await _introspect(flow, token.access_token) == {'active': False} + assert await _introspect(flow, token.refresh_token) == {'active': False} + with pytest.raises(OAuthProtocolException) as error: + await TokenService.issue_token_request( + flow.db, + flow.redis, + { + 'grant_type': 'refresh_token', + 'client_id': 'business-app', + 'refresh_token': token.refresh_token, + }, + client_id='business-app', + signing_key=flow.signer, + kid='regression-key', + ) + assert error.value.error == 'invalid_grant' + + +@pytest.mark.asyncio +async def test_secret_rotation_and_revocation_preserve_existing_authorization(auth_flow: SimpleNamespace) -> None: + """凭据轮换只改变客户端认证,旧授权可通过新凭据继续续期。""" + + flow = auth_flow + token, old_secret, old_header = await _confidential_token(flow) + policy_version = flow.app.policy_version + secret = await OAuthClientManagementService.rotate_secret(flow.db, 'business-app', 'admin', now=flow.now) + assert flow.app.policy_version == policy_version + assert (await _introspect(flow, token.access_token))['active'] is True + assert (await TokenService.authenticate_client(flow.db, authorization=old_header))[1].client_id == 'business-app' + new_header = _basic('business-app', secret.client_secret) + refreshed = await TokenService.issue_token_request( + flow.db, + flow.redis, + { + 'grant_type': 'refresh_token', + 'refresh_token': token.refresh_token, + }, + authorization=old_header, + signing_key=flow.signer, + kid='regression-key', + ) + await OAuthClientManagementService.revoke_secret(flow.db, 'business-app', old_secret.secret_id, 'admin') + assert flow.app.policy_version == policy_version + assert (await _introspect(flow, refreshed.access_token))['active'] is True + with pytest.raises(OAuthProtocolException) as error: + await TokenService.authenticate_client(flow.db, authorization=old_header) + assert error.value.error == 'invalid_client' + renewed = await TokenService.issue_token_request( + flow.db, + flow.redis, + { + 'grant_type': 'refresh_token', + 'refresh_token': refreshed.refresh_token, + }, + authorization=new_header, + signing_key=flow.signer, + kid='regression-key', + ) + assert (await _introspect(flow, renewed.access_token))['active'] is True + + +@pytest.mark.asyncio +@pytest.mark.parametrize('kind', ['access_token', 'refresh_token']) +@pytest.mark.parametrize('hint', [None, 'access_token', 'refresh_token', 'unknown_hint']) +async def test_token_hint_does_not_override_real_token_type( + auth_flow: SimpleNamespace, kind: str, hint: str | None +) -> None: + """内省和撤销均忽略不匹配或未知的可选令牌类型提示。""" + + flow = auth_flow + token, _ = await _user_token(flow, offline=True) + value = getattr(token, kind) + result = await IntrospectionService.introspect( + flow.db, + flow.redis, + value, + flow.caller, + verification_key=flow.signer.public_key(), + now=flow.now, + token_type_hint=hint, + ) + assert result['active'] is True + assert await RevocationService.revoke_request( + flow.db, + flow.redis, + value, + client_id='business-app', + verification_key=flow.signer.public_key(), + token_type_hint=hint, + ) + assert await _introspect(flow, value) == {'active': False} + + +@pytest.mark.asyncio +@pytest.mark.parametrize('expiry', ['idle', 'absolute']) +async def test_effective_expiry_matches_list_count_detail_without_mutating_state( + auth_flow: SimpleNamespace, expiry: str +) -> None: + """列表筛选、总数和详情采用相同有效状态,读取不改变离线凭据语义。""" + + flow = auth_flow + _, grant = await _user_token(flow) + setattr(flow.session, f'{expiry}_expires_at', flow.now - timedelta(seconds=1)) + grant.expires_at = flow.now - timedelta(seconds=1) + await flow.db.commit() + sessions, count = await OAuthSessionManagementService.list_sessions( + flow.db, SessionPageQueryModel(status='expired') + ) + grants, grant_count = await OAuthSessionManagementService.list_grants( + flow.db, GrantPageQueryModel(status='expired') + ) + assert count == grant_count == 1 + assert sessions[0].status == grants[0].status == 'expired' + assert sessions[0].client_ids == ['business-app'] + assert (await OAuthSessionManagementService.get_session(flow.db, flow.session.sid)).status == 'expired' + assert (await OAuthSessionManagementService.get_grant(flow.db, grant.grant_id)).status == 'expired' + assert await OAuthSessionManagementService.list_sessions(flow.db, SessionPageQueryModel(status='active')) == ([], 0) + assert await OAuthSessionManagementService.list_grants(flow.db, GrantPageQueryModel(status='active')) == ([], 0) + assert flow.session.status == grant.status == 'active' + flow.session.status = grant.status = 'revoked' + await flow.db.commit() + assert await OAuthSessionManagementService.list_sessions(flow.db, SessionPageQueryModel(status='expired')) == ( + [], + 0, + ) + assert await OAuthSessionManagementService.list_grants(flow.db, GrantPageQueryModel(status='expired')) == ([], 0) + + +@pytest.mark.asyncio +async def test_online_authorization_records_exact_session_participation( + auth_flow: SimpleNamespace, monkeypatch: pytest.MonkeyPatch +) -> None: + """没有刷新令牌的授权也记录应用,其他设备的参与应用不混入当前登出。""" + + flow = auth_flow + token, _ = await _user_token(flow) + await _user_token(flow) + assert token.refresh_token is None + assert await flow.db.scalar(select(func.count()).select_from(SysSsoSessionClient)) == 1 + coordinator = AfterCommitCoordinator() + _, other = await SsoSessionService.create( + flow.db, + flow.redis, + flow.user.user_id, + flow.subject.subject_id, + flow.subject.auth_version, + 'urn:ruoyi:acr:pwd', + ('pwd',), + pepper=_PEPPER, + now=flow.now, + coordinator=coordinator, + ) + await SsoSessionDao.record_client(flow.db, other.sid, flow.machine.client_pk, flow.now) + await coordinator.commit(flow.db) + assert await SsoSessionDao.client_ids_for_sid(flow.db, flow.session.sid) == ['business-app'] + assert await SsoSessionDao.client_ids_for_sid(flow.db, other.sid) == ['machine'] + register = AsyncMock() + monkeypatch.setattr(LogoutService, '_register_backchannel', register) + coordinator = AfterCommitCoordinator() + result = await LogoutService._logout( + flow.db, + flow.redis, + cookie=flow.cookie, + confirmed=True, + coordinator=coordinator, + now=flow.now, + ) + assert result.session_revoked is True + assert register.call_args.args[3] == {flow.app.client_pk} + await coordinator.commit(flow.db) + assert other.status == 'active' + + +@pytest.mark.asyncio +async def test_legacy_refresh_binding_is_an_exact_session_fallback(auth_flow: SimpleNamespace) -> None: + """升级前的离线凭据仍可恢复精确会话参与者,不依赖同用户授权时间推断。""" + + flow = auth_flow + await _user_token(flow, offline=True) + await flow.db.execute(delete(SysSsoSessionClient)) + await flow.db.commit() + assert await SsoSessionDao.client_ids_for_sid(flow.db, flow.session.sid) == ['business-app'] + + +@pytest.mark.asyncio +async def test_failed_token_transaction_rolls_back_session_participation( + auth_flow: SimpleNamespace, monkeypatch: pytest.MonkeyPatch +) -> None: + """令牌签发事务失败时不能残留成功登录关联记录。""" + + record = AuditService.record + + async def fail_token_audit(db: AsyncSession, event_type: str, *args: object, **kwargs: object) -> object: + if event_type == OidcAuditEvent.TOKEN_ISSUED: + raise RuntimeError('audit unavailable') + return await record(db, event_type, *args, **kwargs) + + monkeypatch.setattr(AuditService, 'record', fail_token_audit) + with pytest.raises(RuntimeError, match='audit unavailable'): + await _user_token(auth_flow) + assert await auth_flow.db.scalar(select(func.count()).select_from(SysSsoSessionClient)) == 0 + + +@pytest.mark.asyncio +async def test_access_policy_without_grant_is_queryable_and_unblocking_does_not_create_grant( + auth_flow: SimpleNamespace, +) -> None: + """首次授权前的禁止策略独立可见,解除后也不凭空恢复授权。""" + + flow = auth_flow + await OAuthSessionManagementService.set_access( + flow.db, flow.user.user_id, 'business-app', True, 'admin', '预先禁止' + ) + query = AccessPolicyPageQueryModel( + user_id=flow.user.user_id, client_id='business-app', access_status='blocked', page_size=1 + ) + rows, count = await OAuthSessionManagementService.list_access_policies(flow.db, query) + assert count == 1 and len(rows) == 1 + assert rows[0].user_name == 'alice' and rows[0].client_name == 'Business app' + assert rows[0].reason == '预先禁止' + assert await flow.db.scalar(select(func.count()).select_from(SysOAuthGrant)) == 0 + query.page_num = 2 + assert await OAuthSessionManagementService.list_access_policies(flow.db, query) == ([], 1) + await OAuthSessionManagementService.set_access( + flow.db, flow.user.user_id, 'business-app', False, 'admin', '允许重新授权' + ) + query.page_num = 1 + assert await OAuthSessionManagementService.list_access_policies(flow.db, query) == ([], 0) + query.access_status = 'allowed' + assert (await OAuthSessionManagementService.list_access_policies(flow.db, query))[1] == 1 + assert await flow.db.scalar(select(func.count()).select_from(SysOAuthGrant)) == 0 + + +def _http_app(flow: SimpleNamespace, router: APIRouter) -> FastAPI: + """挂载真实路由并将数据库依赖定向到隔离测试库。""" + + app = FastAPI() + app.state.redis = flow.redis + app.include_router(router) + handle_exception(app) + for route in router.routes: + if isinstance(route, APIRoute): + for dependency in route.dependant.dependencies: + if dependency.name == 'query_db': + app.dependency_overrides[dependency.call] = lambda: flow.db + else: + app.dependency_overrides[dependency.call] = lambda: None + return app + + +@pytest.mark.asyncio +@pytest.mark.parametrize('method', ['GET', 'POST']) +async def test_authorize_accepts_query_response_mode_on_both_methods( + auth_flow: SimpleNamespace, monkeypatch: pytest.MonkeyPatch, method: str +) -> None: + """标准授权方法均进入真实交互流程,未登录时保持 prompt=none 的标准错误。""" + + monkeypatch.setattr(OidcRateLimiter, 'enforce', AsyncMock()) + values = {**_authorize_request(auth_flow), 'response_mode': 'query', 'prompt': 'none', 'state': 'test-state'} + app = _http_app(auth_flow, authorization_controller) + async with AsyncClient(transport=ASGITransport(app=app), base_url='https://auth.example.com') as client: + response = await client.request( + method, '/oauth2/authorize', **({'params': values} if method == 'GET' else {'data': values}) + ) + assert response.status_code == HTTPStatus.SEE_OTHER + location = urlsplit(response.headers['location']) + assert location.netloc == 'app.example.com' + assert parse_qs(location.query)['error'] == ['login_required'] + assert parse_qs(location.query)['state'] == ['test-state'] + assert response.headers['cache-control'] == 'no-store' + + +@pytest.mark.asyncio +@pytest.mark.parametrize('case', ['duplicate', 'query_and_body', 'json', 'bad_escape', 'oversize']) +async def test_authorize_post_rejects_ambiguous_or_invalid_forms(auth_flow: SimpleNamespace, case: str) -> None: + """支持 POST 不放宽重复参数、编码、媒体类型和体积限制。""" + + body = urlencode(_authorize_request(auth_flow)) + url = '/oauth2/authorize' + media_type = 'application/x-www-form-urlencoded' + expected = 400 + if case == 'duplicate': + body += '&client_id=other' + elif case == 'query_and_body': + url += '?client_id=other' + elif case == 'json': + body, media_type, expected = '{}', 'application/json', 415 + elif case == 'bad_escape': + body += '&state=%XX' + else: + body += '&state=' + 'x' * 70000 + expected = 413 + async with AsyncClient( + transport=ASGITransport(app=_http_app(auth_flow, authorization_controller)), base_url='https://auth.example.com' + ) as client: + response = await client.post(url, content=body, headers={'content-type': media_type}) + assert response.status_code == expected + assert response.json()['error'] == 'invalid_request' + assert 'location' not in response.headers + + +@pytest.mark.asyncio +async def test_unsupported_response_mode_is_only_redirected_to_verified_client(auth_flow: SimpleNamespace) -> None: + """不支持的响应模式只能向已登记回调返回错误。""" + + values = {**_authorize_request(auth_flow), 'response_mode': 'fragment'} + with pytest.raises(OAuthProtocolException) as error: + await AuthorizationService._parse_request(auth_flow.db, values) + assert error.value.error == 'unsupported_response_mode' and error.value.can_redirect + values['redirect_uri'] = 'https://attacker.example/callback' + with pytest.raises(OAuthProtocolException) as error: + await AuthorizationService._parse_request(auth_flow.db, values) + assert not error.value.can_redirect + + +@pytest.mark.asyncio +async def test_access_policy_list_route_returns_camel_case_without_grants(auth_flow: SimpleNamespace) -> None: + """独立策略列表路由返回规范字段,不被授权详情路径截获。""" + + flow = auth_flow + await OAuthSessionManagementService.set_access( + flow.db, flow.user.user_id, 'business-app', True, 'admin', '预先禁止' + ) + app = _http_app(flow, oauth_grant_controller) + async with AsyncClient(transport=ASGITransport(app=app), base_url='http://testserver') as client: + response = await client.get( + '/system/oauth/grant/access/list', params={'accessStatus': 'blocked', 'pageSize': 10} + ) + assert response.status_code == HTTPStatus.OK + payload = response.json() + assert payload['total'] == 1 and payload['rows'][0]['accessStatus'] == 'blocked' + assert payload['rows'][0]['clientId'] == 'business-app' diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_review_regressions.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_review_regressions.py new file mode 100644 index 000000000..f97894faa --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_review_regressions.py @@ -0,0 +1,235 @@ +import json +from datetime import datetime, timedelta, timezone +from unittest.mock import AsyncMock + +import pytest +from sqlalchemy import func, select, text +from sqlalchemy.ext.asyncio import AsyncSession + +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_admin.entity.do.user_do import SysUser +from module_identity.controller.oidc_key_controller import list_oidc_keys +from module_identity.dao.oauth_token_dao import OAuthTokenDao +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_grant_do import SysOAuthGrant, SysOAuthRefreshToken, SysSsoSession +from module_identity.entity.do.oauth_resource_do import SysOAuthClientScope, SysOAuthScope +from module_identity.entity.vo.oauth_client_vo import ClientViewModel +from module_identity.security.opaque_token import parse_opaque_token +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.identity_service import ClaimService +from module_identity.service.session_service import LogoutService +from module_identity.service.token_service import RefreshTokenReuseDetected, TokenResult, TokenService +from utils.oidc_util import OidcUtil +from utils.response_util import ResponseUtil + +NOW = datetime(2026, 9, 12, 10, 0, tzinfo=timezone.utc) +PEPPER = 'regression-refresh-pepper-' + 'x' * 32 + + +async def seed_refresh(db: AsyncSession, monkeypatch: pytest.MonkeyPatch, *, expired: bool = False) -> str: + # 显式确认外键校验已启用,避免SQLite默认配置掩盖原有问题 + await db.execute(text('PRAGMA foreign_keys=ON')) + assert await db.scalar(text('PRAGMA foreign_keys')) == 1 + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', PEPPER) + monkeypatch.setattr(TokenService, '_issue_access_token', AsyncMock(return_value=('signed-access', 600))) + user = SysUser(user_id=9901, user_name='refresh-regression', nick_name='Refresh', status='0', del_flag='0') + client = SysOAuthClient( + client_pk=9901, + client_id='refresh-regression', + client_name='Refresh regression', + client_type='public', + token_endpoint_auth_method='none', + grant_types=['authorization_code', 'refresh_token'], + response_types=['code'], + policy_version=1, + status='0', + ) + db.add_all([user, client]) + await db.flush() + subject = SysIdentitySubject(user_id=user.user_id, subject_id='subject-9901', auth_version=1) + grant = SysOAuthGrant( + grant_id='grant-9901', + user_id=user.user_id, + subject_id=subject.subject_id, + client_pk=client.client_pk, + granted_scopes=['openid', 'offline_access'], + granted_resources=[], + client_policy_version=1, + consented_at=NOW, + status='active', + ) + session = SysSsoSession( + sid='99010000-0000-4000-8000-000000000001', + session_secret_hash='a' * 64, + user_id=user.user_id, + subject_id=subject.subject_id, + auth_version=1, + auth_time=NOW - timedelta(days=1), + last_seen_at=NOW - timedelta(hours=1), + idle_expires_at=NOW + timedelta(minutes=-30 if expired else 30), + absolute_expires_at=NOW + timedelta(hours=-1 if expired else 8), + acr='pwd', + amr=['pwd'], + status='expired' if expired else 'active', + ) + scopes = [ + SysOAuthScope( + scope_pk=9901, + scope_code='openid', + scope_name='OpenID', + scope_type='identity', + claims=['sub'], + create_by='test', + update_by='test', + ), + SysOAuthScope( + scope_pk=9902, + scope_code='offline_access', + scope_name='Offline', + scope_type='identity', + claims=[], + create_by='test', + update_by='test', + ), + ] + db.add_all([subject, grant, session, *scopes]) + await db.flush() + db.add_all([SysOAuthClientScope(client_pk=9901, scope_pk=s.scope_pk) for s in scopes]) + await db.flush() + token = await TokenService._create_refresh_token( + db, + client, + grant, + user, + subject, + session, + ['openid', 'offline_access'], + [], + NOW, + token_pepper=PEPPER, + ) + await db.commit() + return token + + +async def refresh(db: AsyncSession, token: str) -> TokenResult: + return await TokenService.refresh_token( + db, + {'grant_type': 'refresh_token', 'client_id': 'refresh-regression', 'refresh_token': token}, + OAuthClientPrincipal('refresh-regression', 'public', 'none'), + token_pepper=PEPPER, + now=NOW, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize('expired', [False, True]) +async def test_rotation_with_real_foreign_keys_and_replay_revocation( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch, expired: bool +) -> None: + old = await seed_refresh(data_session, monkeypatch, expired=expired) + result = await refresh(data_session, old) + await data_session.commit() + old_row = await OAuthTokenDao.get_by_token_id(data_session, parse_opaque_token(old, 'rt1').token_id) + new_id = parse_opaque_token(result.refresh_token, 'rt1').token_id + successor = await OAuthTokenDao.get_by_token_id(data_session, new_id) + assert old_row.status == 'used' + assert old_row.replaced_by_token_id == new_id + assert successor.parent_token_id == old_row.token_id + assert successor.absolute_expires_at == old_row.absolute_expires_at + with pytest.raises(RefreshTokenReuseDetected): + await refresh(data_session, old) + await data_session.commit() + assert successor.status == 'revoked' + + +@pytest.mark.asyncio +async def test_failed_signing_rolls_back_both_rotation_rows( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + old = await seed_refresh(data_session, monkeypatch) + monkeypatch.setattr(TokenService, '_issue_access_token', AsyncMock(side_effect=RuntimeError('signer unavailable'))) + with pytest.raises(RuntimeError, match='signer unavailable'): + await refresh(data_session, old) + await data_session.rollback() + assert await data_session.scalar(select(func.count()).select_from(SysOAuthRefreshToken)) == 1 + old_row = await OAuthTokenDao.get_by_token_id(data_session, parse_opaque_token(old, 'rt1').token_id) + assert old_row.status == 'active' + assert old_row.replaced_by_token_id is None + + +@pytest.mark.asyncio +async def test_explicit_revocation_still_blocks_offline_refresh( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + old = await seed_refresh(data_session, monkeypatch, expired=True) + assert await SsoSessionDao.revoke( + data_session, '99010000-0000-4000-8000-000000000001', reason='rp_initiated_logout', now=NOW + ) + await data_session.commit() + with pytest.raises(OAuthProtocolException): + await refresh(data_session, old) + + +@pytest.mark.asyncio +async def test_real_key_list_readiness_and_client_response_are_rfc3339(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(OidcConfig, 'oidc_enabled', False) + monkeypatch.setattr( + 'module_identity.controller.oidc_key_controller.OidcKeyManagementService.list_page', + AsyncMock(return_value=([], 0)), + ) + response = await list_oidc_keys(object(), status=None, page_num=1, page_size=10) + assert json.loads(response.body)['readinessCheckedAt'].endswith('Z') + model = ClientViewModel( + client_id='public', + client_name='Public', + client_type='public', + token_endpoint_auth_method='none', + redirect_uris=['https://client.example/callback'], + update_time=NOW, + ) + assert json.loads(ResponseUtil.success(rows=[model]).body)['rows'][0]['updateTime'] == '2026-09-12T10:00:00.000Z' + + +def test_role_values_require_an_explicit_client_allowlist() -> None: + user = {'subject_id': 'sub-1'} + for policy, expected in [ + ({'roles': True}, []), + ({'roles': {'claims': ['roles'], 'allowed_role_keys': ['analyst']}}, ['analyst']), + ]: + claims = ClaimService.build_claims( + user, + ['roles'], + policy, + ['roles'], + roles=['admin', 'analyst', 'finance-internal'], + ) + assert claims['roles'] == expected + + +@pytest.mark.asyncio +async def test_confirmed_logout_revokes_offline_grant_with_expired_owned_cookie( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + old = await seed_refresh(data_session, monkeypatch, expired=True) + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + sid = '99010000-0000-4000-8000-000000000001' + session = await SsoSessionDao.get_by_sid(data_session, sid) + secret = 's' * 43 + session.session_secret_hash = OidcUtil.session_secret_digest(secret, PEPPER) + await data_session.commit() + result = await LogoutService.execute_logout( + data_session, + AsyncMock(), + confirmed=True, + cookie=f'ss1.{sid}.{secret}', + now=NOW, + ) + assert result.session_revoked is True + assert (await SsoSessionDao.get_by_sid(data_session, sid)).status == 'revoked' + assert ( + await OAuthTokenDao.get_by_token_id(data_session, parse_opaque_token(old, 'rt1').token_id) + ).status == 'revoked' diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_revocation_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_revocation_service.py new file mode 100644 index 000000000..539080a20 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_revocation_service.py @@ -0,0 +1,252 @@ +from datetime import datetime, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from config.env import OidcConfig +from module_identity.security.opaque_token import generate_refresh_token, token_digest +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.token_protocol_service import RevocationService +from tests.module_identity.support.redis_fakes import FakeRedis + +_PEPPER = 'revocation-test-pepper-' + 'x' * 32 + + +def _config() -> SimpleNamespace: + """构造撤销所需的最小配置。""" + + return SimpleNamespace( + oidc_issuer='https://auth.example.com', + oidc_allowed_clock_skew_seconds=60, + oidc_token_hash_pepper=_PEPPER, + ) + + +@pytest.fixture(autouse=True) +def _configure_oidc(monkeypatch: pytest.MonkeyPatch) -> None: + """将全局 OIDC 配置固定为本文件测试所需的协议值。""" + + monkeypatch.setattr(OidcConfig, 'oidc_issuer', 'https://auth.example.com') + monkeypatch.setattr(OidcConfig, 'oidc_allowed_clock_skew_seconds', 60) + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', _PEPPER) + + +def _client(client_type: str = 'public') -> SimpleNamespace: + """构造启用的撤销 Client。""" + + return SimpleNamespace( + client_pk=10, + client_id='client-a', + client_type=client_type, + status='0', + token_endpoint_auth_method='client_secret_basic' if client_type == 'confidential' else 'none', + ) + + +@pytest.mark.asyncio +async def test_public_client_can_revoke_only_its_own_refresh_family(monkeypatch: pytest.MonkeyPatch) -> None: + """Public Client 只能撤销自己签发的 Refresh Token Family。""" + + monkeypatch.setattr('module_identity.service.token_protocol_service.AuditService.record', AsyncMock()) + client = _client() + token = generate_refresh_token('token-0001') + row = SimpleNamespace( + client_pk=10, + family_id='family-1', + status='active', + token_hash=token_digest(token, _PEPPER), + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _client_lookup(client) + ) + get_by_token_id = AsyncMock(return_value=row) + revoke_family = AsyncMock() + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthTokenDao.get_by_token_id', get_by_token_id) + monkeypatch.setattr('module_identity.service.token_protocol_service.OAuthTokenDao.revoke_family', revoke_family) + db = object() + + await RevocationService.revoke( + db, + FakeRedis(), + token, + OAuthClientPrincipal('client-a', 'public', 'none'), + ) + + get_by_token_id.assert_awaited_once_with(db, 'token-0001', for_update=True) + revoke_family.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ('client_type', 'db_method', 'principal_type', 'principal_method', 'expected'), + [ + ('public', 'none', 'public', 'none', True), + ('confidential', 'client_secret_basic', 'confidential', 'client_secret_basic', True), + ('public', 'client_secret_basic', 'public', 'none', False), + ('public', 'none', 'public', 'client_secret_basic', False), + ('confidential', 'none', 'confidential', 'client_secret_basic', False), + ('confidential', 'client_secret_basic', 'confidential', 'client_secret_post', False), + ], +) +async def test_revocation_auth_method_pair_is_exact( + monkeypatch: pytest.MonkeyPatch, + client_type: str, + db_method: str, + principal_type: str, + principal_method: str, + expected: bool, +) -> None: + """Revocation 对 Public/none 与 Confidential/Basic 组合逐项 fail closed。""" + + client = _client(client_type) + client.token_endpoint_auth_method = db_method + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _client_lookup(client) + ) + + result = await RevocationService._resolve_caller( + object(), OAuthClientPrincipal('client-a', principal_type, principal_method) + ) + + assert (result is not None) is expected + + +@pytest.mark.asyncio +async def test_unknown_refresh_is_idempotent_success(monkeypatch: pytest.MonkeyPatch) -> None: + """未知或错误 Secret 不得暴露存在性,且 RFC7009 仍返回成功。""" + + client = _client() + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _client_lookup(client) + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthTokenDao.get_by_token_id', AsyncMock(return_value=None) + ) + + result = await RevocationService.revoke( + object(), + FakeRedis(), + generate_refresh_token('token-0002'), + OAuthClientPrincipal('client-a', 'public', 'none'), + ) + + assert result is None + + +@pytest.mark.asyncio +async def test_access_revocation_is_written_only_after_commit(monkeypatch: pytest.MonkeyPatch) -> None: + """Access JTI 撤销必须在数据库提交成功后写入 Redis。""" + + client = _client('confidential') + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _client_lookup(client) + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'subject', + 'client_id': 'client-a', + 'aud': ['https://api.example'], + 'scope': 'api.read', + 'iat': 100, + 'exp': 500, + 'jti': 'access-jti', + }, + ) + redis = FakeRedis() + coordinator = AfterCommitCoordinator() + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + now = datetime.fromtimestamp(200, timezone.utc) + + await RevocationService.revoke( + db, + redis, + 'signed-access', + OAuthClientPrincipal('client-a', 'confidential', 'client_secret_basic'), + coordinator=coordinator, + now=now, + ) + assert await redis.get('oidc:revoked_jti:access-jti') is None + await coordinator.commit(db) + assert await redis.get('oidc:revoked_jti:access-jti') == '1' + + +@pytest.mark.asyncio +async def test_public_client_can_revoke_its_own_access_after_commit(monkeypatch: pytest.MonkeyPatch) -> None: + """Public Client 可撤销自己通过 PKCE 获得的 Access JWT。""" + + client = _client('public') + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _client_lookup(client) + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'subject', + 'client_id': 'client-a', + 'aud': ['https://api.example'], + 'scope': 'api.read', + 'gty': 'authorization_code', + 'iat': 100, + 'exp': 500, + 'jti': 'public-access-jti', + }, + ) + redis = FakeRedis() + coordinator = AfterCommitCoordinator() + db = SimpleNamespace(commit=AsyncMock(), rollback=AsyncMock()) + + await RevocationService.revoke( + db, + redis, + 'signed-public-access', + OAuthClientPrincipal('client-a', 'public', 'none'), + coordinator=coordinator, + now=datetime.fromtimestamp(200, timezone.utc), + ) + await coordinator.commit(db) + + assert await redis.get('oidc:revoked_jti:public-access-jti') == '1' + + +@pytest.mark.asyncio +async def test_public_client_cannot_revoke_other_client_access(monkeypatch: pytest.MonkeyPatch) -> None: + """Public Client 不能撤销其他业务 Client 的 Access JWT。""" + + client = _client('public') + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.OAuthClientDao.get_by_client_id', _client_lookup(client) + ) + monkeypatch.setattr( + 'module_identity.service.token_protocol_service.decode_access_token', + lambda *_args, **_kwargs: { + 'iss': 'https://auth.example.com', + 'sub': 'subject', + 'client_id': 'client-b', + 'aud': ['https://api.example'], + 'scope': 'api.read', + 'gty': 'authorization_code', + 'iat': 100, + 'exp': 500, + 'jti': 'foreign-access-jti', + }, + ) + redis = FakeRedis() + await RevocationService.revoke( + object(), redis, 'foreign-access', OAuthClientPrincipal('client-a', 'public', 'none') + ) + + assert not redis.values + + +def _client_lookup(client: SimpleNamespace) -> object: + """返回指定 Client 的异步 DAO 查询桩。""" + + async def lookup(*_args: object, **_kwargs: object) -> SimpleNamespace: + return client + + return lookup diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_sso_session_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_sso_session_service.py new file mode 100644 index 000000000..0b0f62bdf --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_sso_session_service.py @@ -0,0 +1,622 @@ +import asyncio +import builtins +import json +from datetime import datetime, timedelta, timezone +from uuid import UUID, uuid4 + +import pytest +from fastapi import Response +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from config.env import OidcConfig +from module_admin.entity.do.user_do import SysUser +from module_identity.dao.sso_session_dao import SsoSessionDao +from module_identity.entity.do.identity_subject_do import SysIdentitySubject +from module_identity.entity.do.oauth_grant_do import SysSsoSession +from module_identity.redis_keys import OidcRedisKey +from module_identity.service.infrastructure_service import AfterCommitCoordinator +from module_identity.service.session_service import SsoSessionError, SsoSessionService +from utils.oidc_util import OidcUtil + +_PEPPER = 's' * 32 +_NOW = datetime(2026, 1, 1, tzinfo=timezone.utc) +_COOKIE_SECRET_TEXT_LENGTH = 43 +_SUBJECT_ID = 'aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa' + + +class SessionFakeRedis: + """覆盖 Session 服务所需字符串、Set 和 publish 语义的 FakeRedis。""" + + def __init__(self) -> None: + self.values: dict[str, tuple[str, float | None]] = {} + self.sets: dict[str, set[str]] = {} + self.set_expiries: dict[str, float] = {} + self.expire_seconds: dict[str, int] = {} + self.published: list[tuple[str, str]] = [] + self.lock = asyncio.Lock() + + async def set(self, key: str, value: str, ex: int | None = None, **kwargs: object) -> bool: + async with self.lock: + self.values[key] = (value, None if ex is None else asyncio.get_running_loop().time() + ex) + return True + + async def get(self, key: str) -> str | None: + async with self.lock: + value = self.values.get(key) + if value is not None and value[1] is not None and value[1] <= asyncio.get_running_loop().time(): + self.values.pop(key, None) + return None + return value[0] if value else None + + async def delete(self, *keys: str) -> int: + async with self.lock: + return sum(self.values.pop(key, None) is not None for key in keys) + + async def sadd(self, key: str, *members: str) -> int: + self._expire_set_if_needed(key) + current = self.sets.setdefault(key, set()) + before = len(current) + current.update(members) + return len(current) - before + + async def srem(self, key: str, *members: str) -> int: + self._expire_set_if_needed(key) + current = self.sets.setdefault(key, set()) + removed = sum(member in current for member in members) + current.difference_update(members) + return removed + + async def smembers(self, key: str) -> builtins.set[str]: + self._expire_set_if_needed(key) + return builtins.set(self.sets.get(key, builtins.set())) + + async def expire(self, key: str, seconds: int) -> bool: + """设置 Set 的剩余 TTL。""" + if key not in self.sets: + return False + self.set_expiries[key] = asyncio.get_running_loop().time() + seconds + self.expire_seconds[key] = seconds + return True + + async def ttl(self, key: str) -> int: + """返回 Set 的 Redis 风格剩余 TTL。""" + self._expire_set_if_needed(key) + expiry = self.set_expiries.get(key) + if key not in self.sets: + return -2 + if expiry is None: + return -1 + return max(0, int(expiry - asyncio.get_running_loop().time())) + + def _expire_set_if_needed(self, key: str) -> None: + expiry = self.set_expiries.get(key) + if expiry is not None and expiry <= asyncio.get_running_loop().time(): + self.sets.pop(key, None) + self.set_expiries.pop(key, None) + + async def publish(self, channel: str, message: str) -> int: + self.published.append((channel, message)) + return 1 + + +@pytest.fixture(autouse=True) +def _configure_oidc(monkeypatch: pytest.MonkeyPatch) -> None: + """将全局 OIDC 配置固定为 Session 测试所需的基线。""" + + values = { + 'oidc_enabled': True, + 'oidc_sso_absolute_seconds': 8 * 60 * 60, + 'oidc_sso_remember_absolute_seconds': 7 * 24 * 60 * 60, + 'oidc_sso_idle_seconds': 1800, + 'oidc_sso_cookie_name': '__Host-ruoyi-sso', + 'oidc_sso_cookie_secure': True, + 'oidc_sso_cookie_samesite': 'lax', + 'oidc_sso_cookie_domain': '', + } + for name, value in values.items(): + monkeypatch.setattr(OidcConfig, name, value) + + +def _config() -> OidcConfig: + """返回当前测试使用的全局 OIDC 配置。""" + + return OidcConfig + + +def _config_with_long_remember() -> OidcConfig: + """创建超过协议上限的 remember-me 配置。""" + config = _config() + config.oidc_sso_remember_absolute_seconds = 8 * 24 * 60 * 60 + return config + + +async def _seed_identity(db: AsyncSession, *, auth_version: int = 2) -> None: + """写入真实用户和 Subject 事实。""" + db.add(SysUser(user_id=1, user_name='session-user', nick_name='Session User', status='0', del_flag='0')) + db.add(SysIdentitySubject(user_id=1, subject_id=_SUBJECT_ID, auth_version=auth_version, create_by='test')) + await db.flush() + + +async def _create(db: AsyncSession, redis: SessionFakeRedis, *, remember: bool = False) -> tuple[str, SysSsoSession]: + """创建测试 Session。""" + queue = AfterCommitCoordinator() + result = await SsoSessionService.create( + db, + redis, + 1, + _SUBJECT_ID, + 2, + 'pwd', + ['pwd', 'captcha'], + remember_me=remember, + pepper=_PEPPER, + now=_NOW, + ip_address='127.0.0.1', + user_agent='pytest-agent', + coordinator=queue, + ) + await queue.commit(db) + return result + + +@pytest.mark.asyncio +async def test_create_validate_and_remember_absolute_ttl(data_session: AsyncSession) -> None: + """创建持久化完整 Session,普通和 remember 绝对 TTL 均受配置约束。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + assert cookie.startswith('ss1.') + sid, secret = OidcUtil.parse_sso_cookie(cookie) + UUID(sid) + assert len(secret) == _COOKIE_SECRET_TEXT_LENGTH + assert secret not in row.session_secret_hash + assert cookie not in json.dumps(redis.values) + user_set_key = 'oidc:user_sessions:1' + assert redis.expire_seconds[user_set_key] == 8 * 60 * 60 + session_payload = json.loads(redis.values[OidcRedisKey.sso_session(row.sid)][0]) + assert 'session_secret_hash' not in session_payload + queue = AfterCommitCoordinator() + validated = await SsoSessionService.validate( + data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue + ) + await queue.commit(data_session) + assert validated.subject_id == _SUBJECT_ID + remember_cookie, remember_row = await _create(data_session, redis, remember=True) + assert remember_cookie != cookie + assert remember_row.absolute_expires_at == _NOW + timedelta(days=7) + assert redis.expire_seconds[user_set_key] == 7 * 24 * 60 * 60 + + queue = AfterCommitCoordinator() + capped = await SsoSessionService.create( + data_session, + redis, + 1, + _SUBJECT_ID, + 2, + 'pwd', + ['pwd'], + remember_me=True, + pepper=_PEPPER, + now=_NOW, + coordinator=queue, + ) + await queue.commit(data_session) + assert capped[1].absolute_expires_at == _NOW + timedelta(days=7) + + +@pytest.mark.asyncio +async def test_user_session_set_ttl_never_shortens(data_session: AsyncSession) -> None: + """先创建 remember Session 时,普通 Session 不得缩短用户 Set TTL。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + await _create(data_session, redis, remember=True) + await _create(data_session, redis, remember=False) + assert redis.expire_seconds['oidc:user_sessions:1'] == 7 * 24 * 60 * 60 + + +@pytest.mark.asyncio +async def test_stale_cookie_index_does_not_reject_valid_database_session(data_session: AsyncSession) -> None: + """错误 Redis Cookie 映射视为陈旧,DB 校验成功后覆盖为正确 sid。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + sid, secret = OidcUtil.parse_sso_cookie(cookie) + digest = OidcUtil.session_secret_digest(secret, _PEPPER) + await redis.set(OidcRedisKey.sso_cookie(digest), str(uuid4()), ex=3600) + coordinator = AfterCommitCoordinator() + validated = await SsoSessionService.validate( + data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=coordinator + ) + await coordinator.commit(data_session) + assert validated.sid == sid == row.sid + assert await redis.get(OidcRedisKey.sso_cookie(digest)) == sid + + +@pytest.mark.asyncio +async def test_coordinator_runs_all_callbacks_after_first_failure(data_session: AsyncSession) -> None: + """首个提交后回调失败不阻断后续回调,数据库提交仍视为成功。""" + coordinator = AfterCommitCoordinator() + called: list[str] = [] + + async def fail() -> None: + called.append('first') + raise RuntimeError('副作用失败') + + async def succeed() -> None: + called.append('second') + + await coordinator.register(fail) + await coordinator.register(succeed) + await coordinator.commit(data_session) + assert called == ['first', 'second'] + assert len(coordinator.callback_errors) == 1 + + +@pytest.mark.asyncio +async def test_wrong_secret_and_security_version_fail_closed(data_session: AsyncSession) -> None: + """错误 Secret、Subject auth_version 变化都会撤销并清缓存。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + sid, _ = OidcUtil.parse_sso_cookie(cookie) + wrong_secret = 'a' * 43 + wrong = f'ss1.{sid}.{wrong_secret}' + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate(data_session, redis, wrong, pepper=_PEPPER, now=_NOW, coordinator=queue) + await queue.commit(data_session) + assert row.status == 'active' + + cookie, row = await _create(data_session, redis) + subject = await data_session.scalar(select(SysIdentitySubject).where(SysIdentitySubject.user_id == 1)) + subject.auth_version = 3 + await data_session.flush() + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + await queue.commit(data_session) + assert row.status == 'revoked' + + +@pytest.mark.asyncio +async def test_touch_never_crosses_absolute_and_rotation_invalidates_old(data_session: AsyncSession) -> None: + """滑动 touch 受 absolute 上界约束,Cookie 轮换立即淘汰旧 Secret。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + row.absolute_expires_at = _NOW + timedelta(seconds=100) + await data_session.flush() + queue = AfterCommitCoordinator() + touched = await SsoSessionService.touch( + data_session, + redis, + cookie, + pepper=_PEPPER, + now=_NOW + timedelta(seconds=90), + coordinator=queue, + ) + await queue.commit(data_session) + assert touched.idle_expires_at == _NOW + timedelta(seconds=100) + queue = AfterCommitCoordinator() + rotated = await SsoSessionService.rotate_cookie( + data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue + ) + await queue.commit(data_session) + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + await queue.commit(data_session) + queue = AfterCommitCoordinator() + assert ( + await SsoSessionService.validate(data_session, redis, rotated, pepper=_PEPPER, now=_NOW, coordinator=queue) + ).sid == row.sid + await queue.commit(data_session) + + +@pytest.mark.asyncio +async def test_expired_session_is_marked_expired(data_session: AsyncSession) -> None: + """超过 absolute 后拒绝 Cookie 并把数据库事实标记为 expired。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate( + data_session, + redis, + cookie, + pepper=_PEPPER, + now=_NOW + timedelta(seconds=8 * 60 * 60 + 1), + coordinator=queue, + ) + await queue.commit(data_session) + assert row.status == 'expired' + + +@pytest.mark.asyncio +async def test_db_revocation_invalidates_cached_hit_and_after_commit_order(data_session: AsyncSession) -> None: + """跨服务 DB 撤销不会被缓存命中绕过,撤销清理可延迟到 after-commit。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + await SsoSessionDao.revoke(data_session, row.sid, reason='external', now=_NOW) + await data_session.commit() + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + await queue.commit(data_session) + + cookie, row = await _create(data_session, redis) + coordinator = AfterCommitCoordinator() + + assert await SsoSessionService.revoke( + data_session, redis, row.sid, reason='admin', now=_NOW, coordinator=coordinator + ) + assert row.status == 'revoked' + assert redis.values + await coordinator.commit(data_session) + assert not redis.values + assert json.loads(redis.published[-1][1])['sid'] == row.sid + + +@pytest.mark.asyncio +async def test_revoke_db_failure_does_not_clear_cache( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + """数据库撤销失败时不能提前删除 Redis 状态。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + _, row = await _create(data_session, redis) + + async def fail(*args: object, **kwargs: object) -> bool: + raise RuntimeError('database unavailable') + + monkeypatch.setattr(SsoSessionDao, 'revoke', fail) + queue = AfterCommitCoordinator() + with pytest.raises(RuntimeError): + await SsoSessionService.revoke(data_session, redis, row.sid, now=_NOW, coordinator=queue) + assert redis.values + + +@pytest.mark.asyncio +async def test_revoke_rollback_keeps_cache_until_after_commit(data_session: AsyncSession) -> None: + """数据库回滚时不执行 after-commit 清理,避免缓存先于事实源失效。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + _, row = await _create(data_session, redis) + await data_session.commit() + sid = row.sid + coordinator = AfterCommitCoordinator() + + assert await SsoSessionService.revoke(data_session, redis, sid, now=_NOW, coordinator=coordinator) + await coordinator.rollback(data_session) + assert redis.values + assert coordinator.pending_count == 0 + restored = await SsoSessionDao.get_by_sid(data_session, sid) + assert restored is not None and restored.status == 'active' + + +@pytest.mark.asyncio +async def test_commit_failure_does_not_run_registered_cleanup( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + """数据库 commit 失败时不执行缓存删除或撤销事件。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + _, row = await _create(data_session, redis) + queue = AfterCommitCoordinator() + assert await SsoSessionService.revoke(data_session, redis, row.sid, now=_NOW, coordinator=queue) + + async def fail_commit() -> None: + raise RuntimeError('commit failed') + + monkeypatch.setattr(data_session, 'commit', fail_commit) + with pytest.raises(RuntimeError): + await queue.commit(data_session) + assert redis.values + assert not redis.published + + +@pytest.mark.asyncio +async def test_inactive_validation_only_cleans_cache_without_event(data_session: AsyncSession) -> None: + """已撤销 Session 的验证只清理陈旧缓存,不重复发布撤销事件。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, row = await _create(data_session, redis) + await SsoSessionDao.revoke(data_session, row.sid, reason='external', now=_NOW) + await data_session.commit() + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + await queue.commit(data_session) + assert not redis.published + + +@pytest.mark.asyncio +async def test_create_requires_uuid_and_positive_integer_ttls(data_session: AsyncSession) -> None: + """Subject 必须为 RFC4122 UUID,TTL 拒绝零值和 bool。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + queue = AfterCommitCoordinator() + config = _config() + config.oidc_sso_idle_seconds = True + with pytest.raises(SsoSessionError): + await SsoSessionService.create( + data_session, + redis, + 1, + _SUBJECT_ID, + 2, + 'pwd', + ['pwd'], + pepper=_PEPPER, + now=_NOW, + coordinator=queue, + ) + config.oidc_sso_idle_seconds = 0 + with pytest.raises(SsoSessionError): + await SsoSessionService.create( + data_session, + redis, + 1, + _SUBJECT_ID, + 2, + 'pwd', + ['pwd'], + pepper=_PEPPER, + now=_NOW, + coordinator=queue, + ) + + +@pytest.mark.asyncio +async def test_create_rotate_touch_cache_writes_wait_for_commit(data_session: AsyncSession) -> None: + """创建、轮换和 touch 在事务回滚时都不提前写入或删除缓存。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + queue = AfterCommitCoordinator() + await SsoSessionService.create( + data_session, + redis, + 1, + _SUBJECT_ID, + 2, + 'pwd', + ['pwd'], + pepper=_PEPPER, + now=_NOW, + coordinator=queue, + ) + assert not redis.values + await data_session.rollback() + assert not redis.values + + await _seed_identity(data_session) + cookie, _ = await _create(data_session, redis) + queue = AfterCommitCoordinator() + await SsoSessionService.rotate_cookie(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + assert redis.values + await data_session.rollback() + assert redis.values + + +@pytest.mark.asyncio +async def test_touch_cache_write_waits_for_commit(data_session: AsyncSession) -> None: + """touch 的缓存刷新在事务回滚时不执行。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, _ = await _create(data_session, redis) + queue = AfterCommitCoordinator() + await SsoSessionService.touch(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + await data_session.rollback() + assert redis.values + + +@pytest.mark.asyncio +async def test_after_commit_uses_snapshot_and_publishes_when_cache_clear_fails( + data_session: AsyncSession, monkeypatch: pytest.MonkeyPatch +) -> None: + """提交后回调不读取 ORM 后续突变,缓存清理失败仍发布失效事件。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + coordinator = AfterCommitCoordinator() + cookie, row = await SsoSessionService.create( + data_session, + redis, + 1, + _SUBJECT_ID, + 2, + 'pwd', + ['pwd'], + pepper=_PEPPER, + now=_NOW, + coordinator=coordinator, + ) + row.status = 'revoked' + await coordinator.commit(data_session) + sid, _ = OidcUtil.parse_sso_cookie(cookie) + payload = json.loads(redis.values[OidcRedisKey.sso_session(sid)][0]) + assert payload['status'] == 'active' + + snapshot = SsoSessionService._snapshot(row) + + async def fail_clear(*args: object, **kwargs: object) -> None: + raise RuntimeError('redis clear failed') + + monkeypatch.setattr(SsoSessionService, '_clear_cache', fail_clear) + await SsoSessionService._best_effort_cleanup(redis, snapshot, reason='logout') + assert redis.published + + +@pytest.mark.asyncio +async def test_user_status_and_cookie_response_security(data_session: AsyncSession) -> None: + """停用用户 fail-closed,响应 Cookie 强制 __Host 安全属性。""" + await _seed_identity(data_session) + redis = SessionFakeRedis() + cookie, _ = await _create(data_session, redis) + user = await data_session.scalar(select(SysUser).where(SysUser.user_id == 1)) + user.status = '1' + await data_session.flush() + queue = AfterCommitCoordinator() + with pytest.raises(SsoSessionError): + await SsoSessionService.validate(data_session, redis, cookie, pepper=_PEPPER, now=_NOW, coordinator=queue) + await queue.commit(data_session) + response = Response() + OidcUtil.parse_sso_cookie(cookie) + response.set_cookie(value=cookie, **SsoSessionService.cookie_parameters()) + header = response.headers['set-cookie'] + assert 'Secure' in header and 'HttpOnly' in header and 'Path=/' in header and 'SameSite=lax' in header + cleared = Response() + cleared.delete_cookie(**SsoSessionService.cookie_parameters()) + clear_header = cleared.headers['set-cookie'] + assert '__Host-ruoyi-sso=' in clear_header + assert 'Secure' in clear_header and 'HttpOnly' in clear_header and 'Path=/' in clear_header + + +def test_cookie_parser_rejects_noncanonical_sid_and_invalid_cookie_policy() -> None: + """Cookie sid 必须是服务端生成的规范 UUID,清理也不得绕过安全策略。""" + secret = 'a' * _COOKIE_SECRET_TEXT_LENGTH + sid = 'aaaaaaaa-aaaa-4aaa-8aaa-aaaaaaaaaaaa' + with pytest.raises(ValueError): + OidcUtil.parse_sso_cookie(f'ss1.{sid.upper()}.{secret}') + with pytest.raises(ValueError): + OidcUtil.parse_sso_cookie(f'ss1.{sid.replace("-", "")}.{secret}') + + insecure = _config() + insecure.oidc_sso_cookie_secure = False + with pytest.raises(SsoSessionError): + SsoSessionService.cookie_parameters() + + +@pytest.mark.asyncio +@pytest.mark.parametrize('operation', ['validate', 'validate_logout_cookie']) +@pytest.mark.parametrize( + ('cookie', 'pepper', 'message'), + [ + ('malformed-cookie', _PEPPER, '会话 Cookie 格式无效'), + (f'ss1.{_SUBJECT_ID}.{"a" * _COOKIE_SECRET_TEXT_LENGTH}', 'short', '会话摘要密钥至少需要 32 字节'), + ], +) +async def test_session_entry_points_preserve_domain_errors_for_invalid_credentials( + data_session: AsyncSession, + monkeypatch: pytest.MonkeyPatch, + operation: str, + cookie: str, + pepper: str, + message: str, +) -> None: + """直接调用工具后,格式错误和摘要配置错误仍按会话异常处理。""" + monkeypatch.setattr(OidcConfig, 'oidc_token_hash_pepper', pepper) + with pytest.raises(SsoSessionError, match=message): + if operation == 'validate_logout_cookie': + await SsoSessionService.validate_logout_cookie(data_session, cookie) + else: + await SsoSessionService.validate( + data_session, + SessionFakeRedis(), + cookie, + pepper=pepper, + now=_NOW, + coordinator=AfterCommitCoordinator(), + ) diff --git a/ruoyi-fastapi-backend/tests/module_identity/services/test_token_service.py b/ruoyi-fastapi-backend/tests/module_identity/services/test_token_service.py new file mode 100644 index 000000000..261501417 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/services/test_token_service.py @@ -0,0 +1,741 @@ +import asyncio +import base64 +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from typing import Any +from unittest.mock import AsyncMock + +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine + +from config.env import OidcConfig +from exceptions.exception import OAuthProtocolException +from module_admin.entity.do.dept_do import SysDept +from module_admin.entity.do.role_do import SysRole +from module_admin.entity.do.user_do import SysUser, SysUserRole +from module_identity.dao.oauth_access_policy_dao import OAuthAccessPolicyDao +from module_identity.entity.do.oauth_client_do import SysOAuthClient +from module_identity.entity.do.oauth_resource_do import ( + SysOAuthClientResource, + SysOAuthClientScope, + SysOAuthResource, + SysOAuthScope, +) +from module_identity.security.client_auth import hash_client_secret +from module_identity.security.jwt_profile import decode_access_token, decode_id_token +from module_identity.security.opaque_token import generate_refresh_token, token_digest +from module_identity.security.pkce import generate_code_challenge +from module_identity.security.principal import OAuthClientPrincipal +from module_identity.service.audit_service import AuditService +from module_identity.service.authorization_service import AuthorizationCodeReuseError, AuthorizationCodeService +from module_identity.service.token_service import RefreshTokenReuseDetected, TokenResult, TokenService +from tests.module_identity.support.redis_fakes import FakeRedis +from utils.oidc_util import OidcUtil + +_PEPPER = 'token-service-test-pepper-' + 'x' * 32 +_VERIFIER = 'v' * 64 +_CHALLENGE = generate_code_challenge(_VERIFIER) +_ACCESS_TTL = 120 +_AUTH_VERSION = 4 +_DEPT_ID = 7301 +_DEFAULT_TTL = 600 +_LONG_TTL = 1200 +_RESOURCE_TTL = 900 +_MAX_TTL = 1800 + + +@pytest.mark.asyncio +async def test_confidential_auth_marks_only_matching_secret_used(monkeypatch: pytest.MonkeyPatch) -> None: + """成功 Basic 认证在同一事务标记匹配 Secret,错误 Secret 不更新。""" + secret = 'cs1.' + 's' * 32 + client = _client(client_type='confidential', token_endpoint_auth_method='client_secret_basic') + secret_row = SimpleNamespace(secret_id='secret-1', secret_hash=hash_client_secret(secret)) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.get_by_client_id', AsyncMock(return_value=client) + ) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.list_secrets', AsyncMock(return_value=[secret_row]) + ) + mark_used = AsyncMock(return_value=True) + monkeypatch.setattr('module_identity.service.token_service.OAuthClientDao.mark_secret_used', mark_used) + db = object() + header = 'Basic ' + base64.b64encode(f'{client.client_id}:{secret}'.encode()).decode() + + await TokenService.authenticate_client(db, authorization=header) + + mark_used.assert_awaited_once_with(db, 'secret-1') + mark_used.reset_mock() + bad_header = 'Basic ' + base64.b64encode(f'{client.client_id}:wrong-secret'.encode()).decode() + with pytest.raises(OAuthProtocolException): + await TokenService.authenticate_client(db, authorization=bad_header) + mark_used.assert_not_awaited() + + +def _config() -> SimpleNamespace: + """构造最小 JWT/TTL 配置快照。""" + return SimpleNamespace( + oidc_issuer='https://auth.example.com', + oidc_access_token_ttl_seconds=600, + oidc_max_access_token_ttl_seconds=1800, + oidc_id_token_ttl_seconds=300, + oidc_refresh_token_idle_seconds=3600, + oidc_refresh_token_absolute_seconds=7200, + oidc_token_hash_pepper=_PEPPER, + ) + + +@pytest.fixture(autouse=True) +def _configure_oidc(monkeypatch: pytest.MonkeyPatch) -> None: + """将全局 OIDC 配置固定为本文件测试所需的协议值。""" + + values = { + 'oidc_issuer': 'https://auth.example.com', + 'oidc_access_token_ttl_seconds': _DEFAULT_TTL, + 'oidc_max_access_token_ttl_seconds': _MAX_TTL, + 'oidc_id_token_ttl_seconds': 300, + 'oidc_refresh_token_idle_seconds': 3600, + 'oidc_refresh_token_absolute_seconds': 7200, + 'oidc_token_hash_pepper': _PEPPER, + } + for name, value in values.items(): + monkeypatch.setattr(OidcConfig, name, value) + + +def _client(**overrides: object) -> SimpleNamespace: + """构造 Token 流程所需的 Client 标量快照。""" + value: dict[str, object] = { + 'client_pk': 1001, + 'client_id': 'portal-client', + 'client_type': 'public', + 'grant_types': ['authorization_code', 'refresh_token'], + 'access_token_ttl_seconds': None, + 'refresh_token_idle_seconds': None, + 'refresh_token_absolute_seconds': None, + 'policy_version': 3, + 'status': '0', + } + value.update(overrides) + return SimpleNamespace(**value) + + +def _code_payload(**overrides: object) -> dict[str, object]: + """构造已由 AuthorizationCodeService 验证的 Code 绑定。""" + value: dict[str, object] = { + 'clientPk': 1001, + 'redirectUri': 'https://portal.example/callback', + 'userId': 2001, + 'subjectId': 'subject-2001', + 'authVersion': 4, + 'sid': 'sid-2001', + 'grantId': 'grant-2001', + 'scopes': ['openid', 'profile'], + 'resources': [], + 'nonce': 'nonce-2001', + 'codeChallenge': _CHALLENGE, + 'codeChallengeMethod': 'S256', + 'authTime': '2026-08-24T04:00:00+00:00', + } + value.update(overrides) + return value + + +@pytest.mark.asyncio +@pytest.mark.parametrize('remembered', [True, False]) +async def test_authorization_code_binds_client_redirect_and_pkce( + monkeypatch: pytest.MonkeyPatch, remembered: bool +) -> None: + """验证 Code 兑换严格绑定 Client、Redirect URI 和 S256 verifier。""" + monkeypatch.setattr(OAuthAccessPolicyDao, 'lock_client', AsyncMock()) + monkeypatch.setattr(AuditService, 'record', AsyncMock()) + monkeypatch.setattr(AuditService, 'record_independent', AsyncMock()) + client = _client() + payload = _code_payload() + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.get_by_client_id', + lambda db, client_id, active_only=True: _async_value(client), + ) + + async def identity(*args: object, **kwargs: object) -> tuple[object, object]: + return SimpleNamespace(user_id=2001, status='0', del_flag='0'), SimpleNamespace( + subject_id='subject-2001', auth_version=4 + ) + + async def session(*args: object, **kwargs: object) -> SimpleNamespace: + return SimpleNamespace( + sid='sid-2001', + user_id=2001, + subject_id='subject-2001', + auth_version=4, + status='active', + auth_time=datetime.now(timezone.utc), + acr='urn:test', + amr=['pwd'], + ) + + monkeypatch.setattr('module_identity.service.token_service.SsoSessionDao.record_client', AsyncMock()) + monkeypatch.setattr(TokenService, '_require_user_identity', identity) + monkeypatch.setattr(TokenService, '_require_session', session) + monkeypatch.setattr( + TokenService, + '_require_grant', + lambda *args, **kwargs: _async_value( + SimpleNamespace(grant_id='grant-2001', remembered_scopes=['openid'] if remembered else []) + ), + ) + monkeypatch.setattr( + TokenService, + '_validate_client_scope_resource', + lambda *args, **kwargs: _async_value((['openid', 'profile'], [], None)), + ) + monkeypatch.setattr( + TokenService, + '_issue_access_token', + lambda *args, **kwargs: _async_value(('access', 600)), + ) + monkeypatch.setattr(TokenService, '_issue_id_token', lambda *args, **kwargs: _async_value('id')) + + redis = FakeRedis() + code = await AuthorizationCodeService.issue(redis, payload, pepper=_PEPPER) + result = await TokenService.authorization_code( + object(), + redis, + { + 'grant_type': 'authorization_code', + 'client_id': 'portal-client', + 'code': code, + 'redirect_uri': 'https://portal.example/callback', + 'code_verifier': _VERIFIER, + }, + OAuthClientPrincipal('portal-client', 'public', 'none'), + ) + assert isinstance(result, TokenResult) + assert result.access_token == 'access' and result.id_token == 'id' + assert result.refresh_token is None + assert await redis.get(f'oidc:authorization_code:{code.split(".")[1]}') is None + + with pytest.raises(OAuthProtocolException) as raised: + await TokenService.authorization_code( + object(), + redis, + { + 'grant_type': 'authorization_code', + 'client_id': 'portal-client', + 'code': code, + 'redirect_uri': 'https://portal.example/callback', + 'code_verifier': 'wrong-' + _VERIFIER, + }, + OAuthClientPrincipal('portal-client', 'public', 'none'), + token_pepper=_PEPPER, + ) + assert raised.value.error == 'invalid_grant' + + with pytest.raises(OAuthProtocolException) as raised: + await TokenService.authorization_code( + object(), + redis, + { + 'grant_type': 'authorization_code', + 'client_id': 'portal-client', + 'code': code, + 'redirect_uri': 'https://portal.example/other', + 'code_verifier': _VERIFIER, + }, + OAuthClientPrincipal('portal-client', 'public', 'none'), + token_pepper=_PEPPER, + ) + assert raised.value.error == 'invalid_grant' + + +@pytest.mark.asyncio +async def test_authorization_code_reuse_revokes_bound_grant(monkeypatch: pytest.MonkeyPatch) -> None: + """授权码重用必须撤销首次兑换关联的 Grant 和 Refresh Token。""" + client = _client() + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.get_by_client_id', + AsyncMock(return_value=client), + ) + monkeypatch.setattr( + AuthorizationCodeService, + 'consume', + AsyncMock(side_effect=AuthorizationCodeReuseError()), + ) + monkeypatch.setattr( + AuthorizationCodeService, + 'consumed_payload', + AsyncMock(return_value=_code_payload()), + ) + revoke = AsyncMock(return_value=True) + monkeypatch.setattr('module_identity.service.token_service.OAuthGrantDao.revoke', revoke) + monkeypatch.setattr(AuditService, 'record_independent', AsyncMock()) + db = SimpleNamespace(execute=AsyncMock()) + + with pytest.raises(AuthorizationCodeReuseError): + await TokenService.authorization_code( + db, + FakeRedis(), + { + 'grant_type': 'authorization_code', + 'client_id': 'portal-client', + 'code': 'ac1.code-id.code-secret', + 'redirect_uri': 'https://portal.example/callback', + 'code_verifier': _VERIFIER, + }, + OAuthClientPrincipal('portal-client', 'public', 'none'), + ) + + revoke.assert_awaited_once_with(db, 'grant-2001', reason='authorization_code_reuse') + + +@pytest.mark.asyncio +async def test_client_credentials_is_machine_domain_isolated(monkeypatch: pytest.MonkeyPatch) -> None: + """验证 Client Credentials 仅签发机器 sub,不产生用户 Session Claims。""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + client = _client( + client_type='confidential', + grant_types=['client_credentials'], + ) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.get_by_client_id', + lambda db, client_id, active_only=True: _async_value(client), + ) + monkeypatch.setattr( + TokenService, + '_validate_client_scope_resource', + lambda *args, **kwargs: _async_value((['portal.service.read'], ['https://api.example'], None)), + ) + monkeypatch.setattr(TokenService, '_resolve_signer', lambda *args, **kwargs: _async_value((key, 'kid-1'))) + result = await TokenService.client_credentials( + object(), + {'grant_type': 'client_credentials', 'scope': 'portal.service.read', 'resource': 'https://api.example'}, + OAuthClientPrincipal('portal-client', 'confidential', 'client_secret_basic'), + ) + claims = decode_access_token( + result.access_token, + verification_key=key.public_key(), + issuer='https://auth.example.com', + audience='https://api.example', + ) + assert claims['sub'] == 'client:portal-client' + assert claims['gty'] == 'client_credentials' + assert 'sid' not in claims + assert 'ver' not in claims + assert result.refresh_token is None + assert result.id_token is None + + +@pytest.mark.asyncio +async def test_client_credentials_uses_real_async_dao_and_resource_scope(monkeypatch: pytest.MonkeyPatch) -> None: + """验证 Client Credentials 通过真实 AsyncSession 查询 Client/Scope/Resource 绑定。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + tables = [ + SysOAuthClient.__table__, + SysOAuthResource.__table__, + SysOAuthScope.__table__, + SysOAuthClientScope.__table__, + SysOAuthClientResource.__table__, + ] + async with engine.begin() as connection: + await connection.run_sync(lambda sync: SysOAuthClient.metadata.create_all(sync, tables=tables)) + session_factory = async_sessionmaker(engine, expire_on_commit=False) + async with session_factory() as db: + db.add_all( + [ + SysOAuthClient( + client_pk=9101, + client_id='real-machine', + client_name='Real Machine', + client_type='confidential', + token_endpoint_auth_method='client_secret_basic', + grant_types=['client_credentials'], + response_types=[], + status='0', + access_token_ttl_seconds=120, + ), + SysOAuthResource( + resource_pk=9102, + resource_id='real-api', + resource_name='Real API', + audience='https://real.api.example', + allowed_claims=[], + create_by='test', + update_by='test', + status='0', + ), + SysOAuthScope( + scope_pk=9103, + scope_code='real.service.read', + scope_name='Real Service Read', + scope_type='resource', + resource_pk=9102, + claims=[], + create_by='test', + update_by='test', + status='0', + ), + SysOAuthClientScope(client_pk=9101, scope_pk=9103), + SysOAuthClientResource(client_pk=9101, resource_pk=9102), + ] + ) + await db.flush() + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + monkeypatch.setattr(TokenService, '_resolve_signer', lambda *args, **kwargs: _async_value((key, 'kid-real'))) + result = await TokenService.client_credentials( + db, + {'grant_type': 'client_credentials', 'scope': 'real.service.read', 'resource': 'https://real.api.example'}, + OAuthClientPrincipal('real-machine', 'confidential', 'client_secret_basic'), + ) + assert result.scope == 'real.service.read' + assert result.refresh_token is None + assert result.id_token is None + await engine.dispose() + + +@pytest.mark.asyncio +async def test_user_access_and_id_profiles_have_fixed_claims_and_ttl(monkeypatch: pytest.MonkeyPatch) -> None: + """验证用户 Access/ID Profile、audience、at_hash 和 TTL 上限。""" + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + config = _config() + client = _client(access_token_ttl_seconds=_ACCESS_TTL) + user = SimpleNamespace(user_id=2001, user_name='user', nick_name='User', status='0', del_flag='0') + subject = SimpleNamespace(subject_id='subject-2001', auth_version=4) + now = datetime.now(timezone.utc) + session = SimpleNamespace( + sid='sid-2001', + auth_time=now - timedelta(minutes=1), + acr='urn:test:pwd', + amr=['pwd'], + ) + definition = SimpleNamespace(scope_pk=1, scope_code='openid', claims=['sub'], status='0') + binding = SimpleNamespace(scope_pk=1, claim_filter=None) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.list_scope_bindings', + lambda *args, **kwargs: _async_value([binding]), + ) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.list_scope_definitions', + lambda *args, **kwargs: _async_value([definition]), + ) + monkeypatch.setattr( + 'module_identity.service.token_service.ClaimService.load_roles_and_department', + lambda *args, **kwargs: _async_value(([], None)), + ) + monkeypatch.setattr(TokenService, '_resolve_signer', lambda *args, **kwargs: _async_value((key, 'kid-1'))) + access, expires = await TokenService._issue_access_token( + object(), + client, + user, + subject, + session, + ['openid'], + [], + None, + grant_type='authorization_code', + signing_key=key, + kid='kid-1', + now=now, + ) + claims = decode_access_token( + access, + verification_key=key.public_key(), + issuer=config.oidc_issuer, + audience=f'{config.oidc_issuer}/oauth2/userinfo', + ) + assert expires == _ACCESS_TTL + assert claims['aud'] == [f'{config.oidc_issuer}/oauth2/userinfo'] + assert claims['ver'] == _AUTH_VERSION + assert claims['sid'] == 'sid-2001' + assert 'user_id' not in claims + identity = await TokenService._issue_id_token( + object(), + client, + user, + subject, + session, + ['openid'], + 'nonce-2001', + access, + signing_key=key, + kid='kid-1', + now=now, + ) + id_claims = decode_id_token( + identity, + verification_key=key.public_key(), + issuer=config.oidc_issuer, + audience=client.client_id, + nonce='nonce-2001', + ) + assert id_claims['nonce'] == 'nonce-2001' + assert id_claims['at_hash'] == OidcUtil.access_token_hash(access) + + +@pytest.mark.asyncio +async def test_user_claims_load_real_roles_and_department_rows(monkeypatch: pytest.MonkeyPatch) -> None: + """验证 Claims 从真实角色、用户角色和部门行加载,并仍受三层白名单约束。""" + engine = create_async_engine('sqlite+aiosqlite:///:memory:') + tables = [SysUser.__table__, SysDept.__table__, SysRole.__table__, SysUserRole.__table__] + async with engine.begin() as connection: + await connection.run_sync(lambda sync: SysUser.metadata.create_all(sync, tables=tables)) + session_factory = async_sessionmaker(engine, expire_on_commit=False) + async with session_factory() as db: + user = SysUser( + user_id=7201, + user_name='claims-user', + nick_name='Claims User', + dept_id=_DEPT_ID, + status='0', + del_flag='0', + ) + dept = SysDept(dept_id=_DEPT_ID, dept_name='Security', status='0', del_flag='0') + role = SysRole( + role_id=7401, + role_name='Administrator', + role_key='admin', + role_sort=1, + status='0', + del_flag='0', + ) + db.add_all([user, dept, role, SysUserRole(user_id=7201, role_id=7401)]) + await db.flush() + client = _client(client_pk=7501) + definitions = [ + SimpleNamespace(scope_pk=1, scope_code='roles', claims=['roles'], status='0'), + SimpleNamespace(scope_pk=2, scope_code='dept', claims=['dept_id', 'dept_name'], status='0'), + ] + bindings = [ + SimpleNamespace( + scope_pk=1, claim_filter={'claims': ['roles', 'dept_id', 'dept_name'], 'allowed_role_keys': ['admin']} + ), + SimpleNamespace( + scope_pk=2, claim_filter={'claims': ['roles', 'dept_id', 'dept_name'], 'allowed_role_keys': ['admin']} + ), + ] + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.list_scope_bindings', + lambda *args, **kwargs: _async_value(bindings), + ) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.list_scope_definitions', + lambda *args, **kwargs: _async_value(definitions), + ) + claims = await TokenService._user_claims( + db, + client, + user, + SimpleNamespace(subject_id='subject-7201'), + ['roles', 'dept'], + None, + ) + assert claims['roles'] == ['admin'] + assert claims['dept_id'] == _DEPT_ID + assert claims['dept_name'] == 'Security' + await engine.dispose() + + +def test_access_ttl_uses_overrides_and_global_cap() -> None: + """验证默认、Client/Resource 覆盖和全局上限的最严格策略。""" + assert TokenService._access_ttl(_client(), None) == _DEFAULT_TTL + assert TokenService._access_ttl(_client(access_token_ttl_seconds=_LONG_TTL), None) == _LONG_TTL + resource = SimpleNamespace(access_token_ttl_seconds=_RESOURCE_TTL) + assert TokenService._access_ttl(_client(access_token_ttl_seconds=_LONG_TTL), resource) == _RESOURCE_TTL + assert TokenService._access_ttl(_client(access_token_ttl_seconds=3000), None) == _MAX_TTL + + +@pytest.mark.asyncio +async def test_public_token_entrypoints_require_authenticated_principal() -> None: + """验证公开 Token 流程拒绝 ORM 行和裸 Client ID。""" + with pytest.raises(OAuthProtocolException) as raised: + await TokenService.client_credentials( + object(), + {'grant_type': 'client_credentials'}, + 'portal-client', + ) + assert raised.value.error == 'invalid_client' + + +@pytest.mark.asyncio +async def test_signer_loads_private_key_from_same_active_record(monkeypatch: pytest.MonkeyPatch) -> None: + """验证签名私钥和返回 kid 来自同一轮 Active Key 查询。""" + record = SimpleNamespace(kid='kid-stable') + private_key = object() + monkeypatch.setattr( + 'module_identity.service.token_service.OidcKeyDao.get_active', + lambda *args, **kwargs: _async_value(record), + ) + calls: list[object] = [] + + async def load(value: object, **kwargs: object) -> object: + calls.append(value) + return private_key + + monkeypatch.setattr('module_identity.service.token_service.KeyService.load_private_key_async', load) + key, kid = await TokenService._resolve_signer(object(), None, None, datetime.now(timezone.utc)) + assert key is private_key + assert kid == 'kid-stable' + assert calls == [record] + + +@pytest.mark.asyncio +async def test_refresh_scope_expansion_is_rejected_before_rotation(monkeypatch: pytest.MonkeyPatch) -> None: + """验证 Refresh 请求不能扩大原 Token Scope。""" + monkeypatch.setattr(OAuthAccessPolicyDao, 'lock_client', AsyncMock()) + client = _client(grant_types=['authorization_code', 'refresh_token']) + row = SimpleNamespace( + token_id='token-1', + token_hash='', + family_id='family-1', + client_pk=1001, + status='active', + user_id=2001, + subject_id='subject-2001', + auth_version=4, + sid='sid-2001', + scopes=['openid'], + resources=[], + grant_id='grant-1', + idle_expires_at=datetime.now(timezone.utc) + timedelta(hours=1), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + ) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.get_by_client_id', + lambda db, client_id, active_only=True: _async_value(client), + ) + token = generate_refresh_token() + row.token_hash = token_digest(token, _PEPPER) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthTokenDao.get_by_token_id', lambda *a, **k: _async_value(row) + ) + monkeypatch.setattr( + TokenService, + '_require_user_identity', + lambda *args, **kwargs: _async_value( + (SimpleNamespace(user_id=2001), SimpleNamespace(subject_id='subject-2001', auth_version=4)) + ), + ) + monkeypatch.setattr( + TokenService, + '_require_session', + lambda *args, **kwargs: _async_value(SimpleNamespace(sid='sid-2001')), + ) + monkeypatch.setattr( + TokenService, + '_require_grant', + lambda *args, **kwargs: _async_value(SimpleNamespace(grant_id='grant-1')), + ) + with pytest.raises(OAuthProtocolException) as raised: + await TokenService.refresh_token( + object(), + { + 'grant_type': 'refresh_token', + 'client_id': 'portal-client', + 'refresh_token': token, + 'scope': 'openid profile', + }, + OAuthClientPrincipal('portal-client', 'public', 'none'), + token_pepper=_PEPPER, + ) + assert raised.value.error == 'invalid_grant' + + +@pytest.mark.asyncio +async def test_refresh_concurrent_rotation_has_one_success_and_reuse_revoke(monkeypatch: pytest.MonkeyPatch) -> None: + """验证并发轮换只有一个成功,失败请求触发 Family 重放处理。""" + monkeypatch.setattr(OAuthAccessPolicyDao, 'lock_client', AsyncMock()) + monkeypatch.setattr(AuditService, 'record', AsyncMock()) + client = _client(grant_types=['authorization_code', 'refresh_token']) + row = SimpleNamespace( + token_id='token-1', + token_hash='', + family_id='family-1', + client_pk=1001, + status='active', + user_id=2001, + subject_id='subject-2001', + auth_version=4, + sid='sid-2001', + scopes=['openid'], + resources=[], + grant_id='grant-1', + idle_expires_at=datetime.now(timezone.utc) + timedelta(hours=1), + absolute_expires_at=datetime.now(timezone.utc) + timedelta(hours=2), + ) + token = generate_refresh_token() + row.token_hash = token_digest(token, _PEPPER) + state = {'rotated': False, 'reuse': 0, 'mark_now': None} + lock = asyncio.Lock() + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthClientDao.get_by_client_id', + lambda db, client_id, active_only=True: _async_value(client), + ) + monkeypatch.setattr( + 'module_identity.service.token_service.OAuthTokenDao.get_by_token_id', lambda *a, **k: _async_value(row) + ) + monkeypatch.setattr( + TokenService, + '_require_user_identity', + lambda *args, **kwargs: _async_value( + (SimpleNamespace(user_id=2001), SimpleNamespace(subject_id='subject-2001', auth_version=4)) + ), + ) + monkeypatch.setattr( + TokenService, + '_require_session', + lambda *args, **kwargs: _async_value(SimpleNamespace(sid='sid-2001')), + ) + monkeypatch.setattr( + TokenService, + '_require_grant', + lambda *args, **kwargs: _async_value(SimpleNamespace(grant_id='grant-1')), + ) + monkeypatch.setattr( + TokenService, + '_validate_client_scope_resource', + lambda *args, **kwargs: _async_value((['openid'], [], None)), + ) + + async def mark_used(*args: object, **kwargs: object) -> bool: + state['mark_now'] = kwargs.get('now') + async with lock: + if state['rotated']: + return False + state['rotated'] = True + return True + + async def reuse(*args: object, **kwargs: object) -> int: + state['reuse'] += 1 + return 1 + + monkeypatch.setattr('module_identity.service.token_service.OAuthTokenDao.mark_used', mark_used) + monkeypatch.setattr('module_identity.service.token_service.OAuthTokenDao.refresh_token_family_reuse', reuse) + monkeypatch.setattr(TokenService, '_create_refresh_token', lambda *a, **k: _async_value(token)) + monkeypatch.setattr(TokenService, '_issue_access_token', lambda *a, **k: _async_value(('access', 600))) + results = await asyncio.gather( + TokenService.refresh_token( + object(), + {'grant_type': 'refresh_token', 'client_id': 'portal-client', 'refresh_token': token}, + OAuthClientPrincipal('portal-client', 'public', 'none'), + token_pepper=_PEPPER, + ), + TokenService.refresh_token( + object(), + {'grant_type': 'refresh_token', 'client_id': 'portal-client', 'refresh_token': token}, + OAuthClientPrincipal('portal-client', 'public', 'none'), + token_pepper=_PEPPER, + ), + return_exceptions=True, + ) + assert sum(isinstance(item, TokenResult) for item in results) == 1 + assert sum(isinstance(item, RefreshTokenReuseDetected) for item in results) == 1 + assert state['reuse'] == 1 + assert isinstance(state['mark_now'], datetime) + assert state['mark_now'].tzinfo is timezone.utc + + +def _async_value(value: object) -> Any: + """返回异步测试桩。""" + + async def result() -> object: + return value + + return result() diff --git a/ruoyi-fastapi-backend/tests/module_identity/support/__init__.py b/ruoyi-fastapi-backend/tests/module_identity/support/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/ruoyi-fastapi-backend/tests/module_identity/support/redis_fakes.py b/ruoyi-fastapi-backend/tests/module_identity/support/redis_fakes.py new file mode 100644 index 000000000..862b0ccc4 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/module_identity/support/redis_fakes.py @@ -0,0 +1,139 @@ +import asyncio +import json +import time +from typing import Any + + +class FakeRedis: + """实现本轮服务所用 SET NX、TTL、GET/DEL 和 Lua 等效语义的 FakeRedis。""" + + def __init__(self) -> None: + self.values: dict[str, tuple[str, float | None]] = {} + self.lists: dict[str, list[str]] = {} + self.lock = asyncio.Lock() + self.set_calls: list[tuple[str, dict[str, Any]]] = [] + self.eval_calls: list[tuple[str, tuple[Any, ...]]] = [] + + def _purge(self, key: str) -> None: + item = self.values.get(key) + if item is not None and item[1] is not None and item[1] <= time.monotonic(): + self.values.pop(key, None) + + async def set(self, key: str, value: str, ex: int | None = None, nx: bool = False, **kwargs: Any) -> bool: + async with self.lock: + self._purge(key) + self.set_calls.append((key, {'ex': ex, 'nx': nx, **kwargs})) + if nx and key in self.values: + return False + self.values[key] = (value, time.monotonic() + ex if ex is not None else None) + return True + + async def get(self, key: str) -> str | None: + async with self.lock: + self._purge(key) + value = self.values.get(key) + return value[0] if value else None + + async def exists(self, key: str) -> int: + """按 Redis EXISTS 语义返回键数量。""" + async with self.lock: + self._purge(key) + return int(key in self.values) + + async def delete(self, key: str) -> int: + async with self.lock: + self._purge(key) + return int(self.values.pop(key, None) is not None) + + async def rpush(self, key: str, value: str) -> int: + """追加队列元素。""" + async with self.lock: + self.lists.setdefault(key, []).append(value) + return len(self.lists[key]) + + async def lpop(self, key: str) -> str | None: + """原子弹出队首元素。""" + async with self.lock: + values = self.lists.get(key, []) + return values.pop(0) if values else None + + async def ltrim(self, key: str, start: int, end: int) -> bool: + """按 Redis 负索引语义裁剪列表。""" + async with self.lock: + values = self.lists.get(key, []) + size = len(values) + first = start if start >= 0 else max(0, size + start) + last = end if end >= 0 else size + end + self.lists[key] = values[first : last + 1] if first <= last else [] + return True + + async def ttl(self, key: str) -> int: + async with self.lock: + self._purge(key) + value = self.values.get(key) + if value is None or value[1] is None: + return -1 + return max(0, int(value[1] - time.monotonic())) + + async def eval(self, script: str, numkeys: int, *args: Any) -> Any: # noqa: PLR0911, PLR0912 + async with self.lock: + self.eval_calls.append((script, args)) + key = args[0] + self._purge(key) + if 'INCR' in script and 'EXPIRE' in script: + current = int(self.values.get(key, ('0', None))[0]) + 1 + ttl = int(args[1]) + expiry = time.monotonic() + ttl + self.values[key] = (str(current), expiry) + return [current, ttl] + if 'codeHash' in script and 'KEYS[3]' in script: + value = self.values.get(key) + tombstone_key = args[1] + payload_key = args[2] + if value is None: + self._purge(tombstone_key) + consumed = self.values.get(payload_key) + if consumed is not None: + try: + consumed_payload = json.loads(consumed[0]) + except json.JSONDecodeError: + consumed_payload = None + if isinstance(consumed_payload, dict) and consumed_payload.get('codeHash') != args[3]: + return -3 + return -4 if tombstone_key in self.values else None + try: + payload = json.loads(value[0]) + except json.JSONDecodeError: + return -2 + if str(payload.get('version')) != str(args[4]) or payload.get('codeHash') != args[3]: + return -3 + self.values.pop(key, None) + self.values[tombstone_key] = (str(args[5]), time.monotonic() + int(args[6])) + self.values[payload_key] = (value[0], time.monotonic() + int(args[6])) + return value[0] + if 'codeHash' in script: + value = self.values.get(key) + if value is None: + return None + try: + payload = json.loads(value[0]) + except json.JSONDecodeError: + return -2 + if str(payload.get('version')) != str(args[2]) or payload.get('codeHash') != args[1]: + return -3 + self.values.pop(key, None) + return value[0] + value = self.values.get(key) + if value is None: + return -1 + current = json.loads(value[0]) + expected = json.loads(args[2]) + if value[1] is None or value[1] <= time.monotonic(): + self.values.pop(key, None) + return -5 + if str(current.get('version')) != str(args[1]): + return -2 + if current.get('status') not in expected: + return -3 + self.values[key] = (args[3], value[1]) + return 1 diff --git a/ruoyi-fastapi-backend/tests/server/test_oidc_cors_refresh.py b/ruoyi-fastapi-backend/tests/server/test_oidc_cors_refresh.py new file mode 100644 index 000000000..c9bd54925 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/server/test_oidc_cors_refresh.py @@ -0,0 +1,46 @@ +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest + +from config.env import OidcConfig +from module_identity.service.runtime_service import OidcRuntimeService + + +@pytest.mark.asyncio +async def test_registered_cors_changes_reach_other_workers_and_fail_closed(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(OidcConfig, 'oidc_enabled', True) + state = {'version': '0', 'origins': ('https://old.example',)} + + async def get(_key: str) -> str: + return state['version'] + + async def incr(_key: str) -> None: + state['version'] = str(int(state['version']) + 1) + + async def load() -> tuple[str, ...]: + return state['origins'] + + redis = SimpleNamespace(get=get, incr=incr) + worker_a = SimpleNamespace(state=SimpleNamespace(redis=redis)) + worker_b = SimpleNamespace(state=SimpleNamespace(redis=redis)) + monkeypatch.setattr(OidcRuntimeService, 'load_cors_origins', load) + for app in (worker_a, worker_b): + await OidcRuntimeService.ensure_cors_snapshot(app) + state['origins'] = ('https://new.example',) + await OidcRuntimeService.cors_snapshot_callback(worker_a)() + await OidcRuntimeService.ensure_cors_snapshot(worker_b) + assert ( + worker_a.state.oidc_registered_cors_origins == worker_b.state.oidc_registered_cors_origins == state['origins'] + ) + # 有时限的缓存兜底无需依赖变更通知 + state['origins'] = () + worker_b.state.oidc_cors_loaded_at = 0 + await OidcRuntimeService.ensure_cors_snapshot(worker_b) + assert worker_b.state.oidc_registered_cors_origins == () + worker_a.state.oidc_cors_loaded_at = 0 + monkeypatch.setattr( + OidcRuntimeService, 'load_cors_origins', AsyncMock(side_effect=RuntimeError('database unavailable')) + ) + await OidcRuntimeService.ensure_cors_snapshot(worker_a) + assert worker_a.state.oidc_registered_cors_origins == () diff --git a/ruoyi-fastapi-backend/tests/server/test_oidc_runtime.py b/ruoyi-fastapi-backend/tests/server/test_oidc_runtime.py new file mode 100644 index 000000000..0d9b0447c --- /dev/null +++ b/ruoyi-fastapi-backend/tests/server/test_oidc_runtime.py @@ -0,0 +1,177 @@ +import asyncio +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import pytest + +from module_identity.service.key_service import KeyServiceError +from module_identity.service.runtime_service import OidcRuntimeService +from server import _start_background_tasks + + +@pytest.mark.asyncio +async def test_oidc_runtime_validation_is_a_noop_when_disabled() -> None: + """关闭认证中心时不得创建额外数据库会话。""" + with ( + patch('module_identity.service.runtime_service.OidcConfig.oidc_enabled', False), + patch('module_identity.service.runtime_service.DataSourceRegistry.session') as session, + patch( + 'module_identity.service.runtime_service.KeyService.get_signing_key', + new_callable=AsyncMock, + ) as get_signing_key, + ): + await OidcRuntimeService.validate_runtime() + + session.assert_not_called() + get_signing_key.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_cors_snapshot_callback_fails_closed_without_raising() -> None: + """管理提交后的 CORS 刷新失败时运行时快照必须进入拒绝态。""" + app = SimpleNamespace(state=SimpleNamespace(oidc_registered_cors_origins=('https://old.example',))) + with patch.object(OidcRuntimeService, 'refresh_cors_snapshot', new=AsyncMock(side_effect=RuntimeError('db down'))): + await OidcRuntimeService.cors_snapshot_callback(app)() + assert app.state.oidc_registered_cors_origins == () + + +@pytest.mark.asyncio +async def test_oidc_runtime_validation_loads_database_active_key() -> None: + """启用认证中心时每个 worker 都验证数据库事实源中的 active 私钥。""" + db = object() + + @asynccontextmanager + async def session() -> AsyncIterator[object]: + yield db + + with ( + patch('module_identity.service.runtime_service.OidcConfig.oidc_enabled', True), + patch('module_identity.service.runtime_service.DataSourceRegistry.session', side_effect=session), + patch( + 'module_identity.service.runtime_service.KeyService.get_signing_key', + new_callable=AsyncMock, + ) as get_signing_key, + ): + await OidcRuntimeService.validate_runtime() + + get_signing_key.assert_awaited_once() + assert get_signing_key.await_args.args == (db,) + + +@pytest.mark.asyncio +async def test_oidc_runtime_validation_fails_closed_without_usable_key() -> None: + """active 元数据或私钥不可用时后台继续启动,但 OIDC 标记为未就绪。""" + + @asynccontextmanager + async def session() -> AsyncIterator[object]: + yield object() + + with ( + patch('module_identity.service.runtime_service.OidcConfig.oidc_enabled', True), + patch('module_identity.service.runtime_service.DataSourceRegistry.session', side_effect=session), + patch( + 'module_identity.service.runtime_service.KeyService.get_signing_key', + new=AsyncMock(side_effect=KeyServiceError('private key does not match public JWK')), + ), + ): + readiness = await OidcRuntimeService.validate_runtime() + + assert readiness.enabled is True + assert readiness.ready is False + assert readiness.reason == 'signing_key_unavailable' + + +@pytest.mark.asyncio +async def test_background_loops_wait_in_both_leader_and_follower() -> None: + """Leader 和 Follower 都启动循环,由每轮当前租约决定是否执行业务。""" + app = SimpleNamespace(state=SimpleNamespace(redis=object(), application_leader=True)) + fake_key_task = object() + fake_retry_task = object() + created: list[object] = [] + + def create_task(coro: object) -> object: + created.append(coro) + coro.close() + return (fake_key_task, fake_retry_task)[len(created) - 1] + + with ( + patch('module_identity.service.runtime_service.OidcConfig.oidc_enabled', True), + patch('module_identity.service.runtime_service.SchedulerManager.is_application_leader', return_value=True), + patch('module_identity.service.runtime_service.asyncio.create_task', side_effect=create_task), + ): + await OidcRuntimeService.start_background_tasks(app) + assert app.state.oidc_key_lifecycle_task is fake_key_task + assert app.state.oidc_backchannel_retry_task is fake_retry_task + + app.state.application_leader = False + created.clear() + with ( + patch('module_identity.service.runtime_service.OidcConfig.oidc_enabled', True), + patch('module_identity.service.runtime_service.SchedulerManager.is_application_leader', return_value=False), + patch('module_identity.service.runtime_service.asyncio.create_task', side_effect=create_task), + ): + await OidcRuntimeService.start_background_tasks(app) + assert app.state.oidc_key_lifecycle_task is fake_key_task + assert app.state.oidc_backchannel_retry_task is fake_retry_task + + +@pytest.mark.asyncio +@pytest.mark.parametrize('loop', ['key', 'retry']) +async def test_background_loops_resume_after_lease_reacquisition(loop: str) -> None: + db = SimpleNamespace(commit=AsyncMock()) + + @asynccontextmanager + async def session() -> AsyncIterator[SimpleNamespace]: + yield db + + with ( + patch( + 'module_identity.service.runtime_service.SchedulerManager.is_application_leader', + side_effect=[False, True, False, True], + ), + patch( + 'module_identity.service.runtime_service.asyncio.sleep', + new=AsyncMock(side_effect=[None, None, None, None, asyncio.CancelledError()]), + ), + patch('module_identity.service.runtime_service.DataSourceRegistry.session', side_effect=session), + patch( + 'module_identity.service.runtime_service.KeyService.activate_due', new=AsyncMock(return_value=0) + ) as activate, + patch('module_identity.service.runtime_service.KeyService.retire_due', new=AsyncMock(return_value=0)), + patch('module_identity.service.runtime_service.OAuthAuditDao.archive_before', new=AsyncMock(return_value=0)), + patch( + 'module_identity.service.runtime_service.LogoutService.consume_backchannel_retry', + new=AsyncMock(return_value=0), + ) as consume, + pytest.raises(asyncio.CancelledError), + ): + if loop == 'key': + await OidcRuntimeService.key_lifecycle_loop(SimpleNamespace(state=SimpleNamespace(redis=object()))) + else: + await OidcRuntimeService.backchannel_retry_loop(object()) + expected_leader_periods = 2 + assert (activate if loop == 'key' else consume).await_count == expected_leader_periods + + +@pytest.mark.asyncio +async def test_server_background_task_setup_delegates_oidc_runtime() -> None: + """应用入口只负责装配通用任务和认证中心运行时。""" + app = SimpleNamespace(state=SimpleNamespace(redis=object())) + log_task = object() + + def create_task(coro: object) -> object: + coro.close() + return log_task + + with ( + patch('server.SchedulerManager.init_system_scheduler', new_callable=AsyncMock) as init_scheduler, + patch('server.LogAggregatorService.consume_stream', new_callable=AsyncMock), + patch('server.asyncio.create_task', side_effect=create_task), + patch.object(OidcRuntimeService, 'start_background_tasks', new_callable=AsyncMock) as start_oidc, + ): + await _start_background_tasks(app) + + init_scheduler.assert_awaited_once_with(app.state.redis) + start_oidc.assert_awaited_once_with(app) diff --git a/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py b/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py index 7813e06ca..dad9a7164 100644 --- a/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py +++ b/ruoyi-fastapi-backend/tests/server/test_plugin_runtime.py @@ -56,6 +56,8 @@ async def test_initialize_application_runtime_delegates_plugin_steps() -> None: patch('server.RedisUtil.check_redis_connection', new_callable=AsyncMock) as check_redis_connection, patch('server.RedisUtil.init_sys_dict', new_callable=AsyncMock) as init_sys_dict, patch('server.RedisUtil.init_sys_config', new_callable=AsyncMock) as init_sys_config, + patch('server.OidcRuntimeService.refresh_cors_snapshot', new_callable=AsyncMock), + patch('server.OidcRuntimeService.validate_runtime', new_callable=AsyncMock), patch('server._start_background_tasks', new_callable=AsyncMock) as start_background_tasks, ): await _initialize_application_runtime(fake_app, application_leader=True) @@ -107,6 +109,8 @@ async def run_as_plugin_writer( patch('server.RedisUtil.check_redis_connection', new_callable=AsyncMock) as check_redis_connection, patch('server.RedisUtil.init_sys_dict', new_callable=AsyncMock), patch('server.RedisUtil.init_sys_config', new_callable=AsyncMock), + patch('server.OidcRuntimeService.refresh_cors_snapshot', new_callable=AsyncMock), + patch('server.OidcRuntimeService.validate_runtime', new_callable=AsyncMock), patch('server._start_background_tasks', new_callable=AsyncMock), ): await _initialize_application_runtime(fake_app, application_leader=False) diff --git a/ruoyi-fastapi-backend/tests/utils/test_oidc_util.py b/ruoyi-fastapi-backend/tests/utils/test_oidc_util.py new file mode 100644 index 000000000..35b5c9d22 --- /dev/null +++ b/ruoyi-fastapi-backend/tests/utils/test_oidc_util.py @@ -0,0 +1,164 @@ +import json +import subprocess +import sys +from pathlib import Path +from types import SimpleNamespace + +import pytest +from cryptography.exceptions import InvalidTag +from cryptography.hazmat.primitives.asymmetric import rsa + +from utils.oidc_util import OidcUtil + +_BACKEND_ROOT = Path(__file__).resolve().parents[2] +_V1_CIPHERTEXT = 'v1.AAECAwQFBgcICQoLh36B9Tc6XrrXlTTTKJydxYN2UsldJqzgWxIn_pyIfSbMrPlmPXHJz6_NAJpbzQ==' +_V2_CIPHERTEXT = ( + 'v2.c2FsdAABAgMEBQYHCAkKCwwNDg8AAQIDBAUGBwgJCgvFPMVj6MHgGRqLTZd7zlLGG0F_YMv1J22s0Nx84W-t468Xts0APDbdZpj8hlwN' +) + + +def test_query_replacement_preserves_unrelated_duplicates_blank_values_and_fragment() -> None: + location = OidcUtil.replace_query_parameters( + 'https://rp.example/cb?keep=one&keep=two&empty=&state=old&error=old#part', + [('code', 'a+/='), ('state', '中文 +&')], + {'code', 'state', 'error'}, + ) + assert location == ( + 'https://rp.example/cb?keep=one&keep=two&empty=&code=a%2B%2F%3D&state=%E4%B8%AD%E6%96%87+%2B%26#part' + ) + + +def test_logout_state_distinguishes_missing_and_empty_state() -> None: + uri = 'https://rp.example/cb?state=old&keep=1&state=older' + assert OidcUtil.append_state(uri, None) == uri + assert OidcUtil.append_state(uri, '') == 'https://rp.example/cb?keep=1&state=' + + +def test_interaction_csrf_is_only_in_fragment() -> None: + assert OidcUtil.interaction_url('https://id.example/login?lang=zh#old', 'i-1', 's+/=') == ( + 'https://id.example/login?lang=zh&interaction=i-1#csrf=s+/=' + ) + + +@pytest.mark.parametrize(('validate_all_first', 'message'), [(False, '参数重复'), (True, '参数无效')]) +def test_batch_validation_preserves_each_callers_error_precedence(validate_all_first: bool, message: str) -> None: + with pytest.raises(ValueError, match=message): + OidcUtil.split_batch('one,one,bad/path', 'ids', max_size=100, validate_all_first=validate_all_first) + + +@pytest.mark.parametrize('value', ['one,,two', 'one,%2Ftwo', 'one,t wo', 'one,t\\wo']) +def test_batch_rejects_ambiguous_path_identifiers(value: str) -> None: + with pytest.raises(ValueError, match='参数无效'): + OidcUtil.split_batch(value, 'ids', max_size=100) + + +def test_json_keeps_unicode_and_key_order_without_accepting_duplicate_input_keys() -> None: + assert OidcUtil.serialize_json({'z': [1], 'a': '中文'}, error_message='载荷无效') == '{"a":"中文","z":[1]}' + with pytest.raises(ValueError, match='重复字段'): + json.loads('{"nested":{"scope":"openid","scope":"admin"}}', object_pairs_hook=OidcUtil.json_object_pairs) + + +@pytest.mark.parametrize('value', [float('nan'), float('inf'), float('-inf'), object()]) +def test_cache_json_rejects_nonfinite_numbers_and_nonserializable_objects(value: object) -> None: + with pytest.raises(ValueError, match='载荷无效'): + OidcUtil.serialize_json({'value': value}, error_message='载荷无效') + + +def test_digest_vectors_preserve_existing_wire_and_browser_binding_formats() -> None: + assert OidcUtil.sha256_digest('abc') == 'ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad' + assert OidcUtil.hmac_sha256('Hi There', b'\x0b' * 20) == ( + 'b0344c61d8db38535ca8afceaf0bf12b881dc200c9833da726e9376c2e32cff7' + ) + assert OidcUtil.access_token_hash('SlAV32hkKG') == 'rXH7QWVTZnXYCou_6Vdpfg' + # 现有退出凭据使用字面量反斜杠加零,不能在重构中替换为 NUL。 + assert OidcUtil.logout_confirmation_digest('ss1.cookie', 'browser-nonce', b'p' * 32) == ( + '55d146629427f15316caa1355ccfbe3d6446b2b6b4e26787750270f890e21d2d' + ) + + +@pytest.mark.parametrize('ciphertext', [_V1_CIPHERTEXT, _V2_CIPHERTEXT], ids=['legacy-v1', 'salted-v2']) +def test_decrypt_accepts_persisted_private_key_formats(ciphertext: str) -> None: + assert OidcUtil.decrypt_signing_private_key(ciphertext, b'm' * 32) == b'private-key-regression-fixture' + + +def test_encrypt_preserves_v2_envelope_and_rejects_wrong_decryption_key() -> None: + ciphertext = OidcUtil.encrypt_signing_private_key( + b'private-key-regression-fixture', b'm' * 32, salt=bytes(range(16)), nonce=bytes(range(12)) + ) + assert ciphertext == _V2_CIPHERTEXT + with pytest.raises(InvalidTag): + OidcUtil.decrypt_signing_private_key(ciphertext, b'x' * 32) + + +@pytest.fixture(scope='module') +def public_jwk() -> dict[str, str]: + key = rsa.generate_private_key(public_exponent=65537, key_size=2048) + return OidcUtil.rsa_public_jwk(key.public_key(), 'k1') + + +def test_public_jwk_strips_private_fields_without_mutating_input(public_jwk: dict[str, str]) -> None: + record = SimpleNamespace(kid='k1', public_jwk={**public_jwk, 'd': 'private-material'}) + actual = OidcUtil.normalize_public_jwk(record, min_rsa_bits=2048, min_exponent=3, max_exponent=2**32) + assert actual == public_jwk + assert record.public_jwk['d'] == 'private-material' + record.kid = 'another-key' + with pytest.raises(ValueError, match='不匹配'): + OidcUtil.normalize_public_jwk(record, min_rsa_bits=2048, min_exponent=3, max_exponent=2**32) + + +@pytest.mark.parametrize(('field', 'value'), [('n', 'AQ'), ('e', 'Ag'), ('e', 'AQAB=')]) +def test_public_jwk_rejects_weak_or_noncanonical_parameters(public_jwk: dict[str, str], field: str, value: str) -> None: + with pytest.raises(ValueError): + OidcUtil.normalize_public_jwk( + {**public_jwk, field: value}, min_rsa_bits=2048, min_exponent=3, max_exponent=2**32 + ) + + +@pytest.mark.parametrize( + 'uri', + [ + 'http://rp.example/logout', + 'https://user:pass@rp.example/logout', + 'https://rp.example/logout?q=x', + 'https://rp.example/logout#part', + 'https://127.0.0.1/logout', + 'https://[::1]/logout', + 'https://rp.example:bad/logout', + ], +) +def test_backchannel_parser_rejects_unsafe_structures_and_literal_private_ips(uri: str) -> None: + assert OidcUtil.parse_backchannel_uri(uri) is None + + +def test_uri_helpers_preserve_different_registration_policies() -> None: + localhost = 'http://localhost:5173/callback' + assert OidcUtil.validate_registered_uri('redirect', localhost) == localhost + assert OidcUtil.is_safe_post_logout_uri(localhost) + with pytest.raises(ValueError, match='HTTPS'): + OidcUtil.validate_resource_audience(localhost, max_length=500) + assert OidcUtil.parse_backchannel_uri('https://rp.example:8443/logout') == ('rp.example', 8443) + + +@pytest.mark.parametrize( + 'modules', + [ + ['utils.oidc_util', 'exceptions.exception', 'utils.time_util'], + ['exceptions.exception', 'utils.time_util', 'utils.oidc_util'], + ['utils.time_util', 'utils.oidc_util', 'exceptions.exception'], + ], +) +def test_utilities_and_exceptions_import_without_cycles_or_loading_application_state(modules: list[str]) -> None: + script = ( + 'import importlib, sys\n' + f'for name in {modules!r}: importlib.import_module(name)\n' + 'assert not {"config.env", "config.database", "redis", "sqlalchemy", "fastapi"}.intersection(sys.modules)\n' + ) + result = subprocess.run( + [sys.executable, '-X', 'utf8', '-c', script], + cwd=_BACKEND_ROOT, + capture_output=True, + text=True, + timeout=30, + check=False, + ) + assert result.returncode == 0, result.stderr diff --git a/ruoyi-fastapi-backend/tests/utils/test_time_util.py b/ruoyi-fastapi-backend/tests/utils/test_time_util.py index 452acacf8..d423eaf37 100644 --- a/ruoyi-fastapi-backend/tests/utils/test_time_util.py +++ b/ruoyi-fastapi-backend/tests/utils/test_time_util.py @@ -25,6 +25,14 @@ def test_datetime_is_normalized_to_utc_and_business_timezone() -> None: assert TimezoneUtil.to_business_time(TimezoneUtil.to_utc(shanghai_time), 'Asia/Shanghai') == shanghai_time +def test_optional_utc_preserves_none_and_rejects_naive_datetime() -> None: + assert TimezoneUtil.to_optional_utc(None) is None + value = datetime.fromisoformat('2026-08-28T10:30:00+08:00') + assert TimezoneUtil.to_optional_utc(value) == datetime(2026, 8, 28, 2, 30, tzinfo=timezone.utc) + with pytest.raises(ValueError, match='必须携带时区信息'): + TimezoneUtil.to_optional_utc(datetime(2026, 8, 28, 10, 30)) + + def test_to_utc_milliseconds_truncates_sub_millisecond_precision() -> None: value = datetime.fromisoformat('2026-08-28T10:30:00.123999+08:00') diff --git a/ruoyi-fastapi-backend/utils/oidc_util.py b/ruoyi-fastapi-backend/utils/oidc_util.py new file mode 100644 index 000000000..af67be352 --- /dev/null +++ b/ruoyi-fastapi-backend/utils/oidc_util.py @@ -0,0 +1,1365 @@ +import base64 +import binascii +import hashlib +import hmac +import ipaddress +import json +import math +import re +import secrets +from collections.abc import Collection, Iterable, Mapping, Sequence +from datetime import datetime +from typing import Any +from urllib.parse import parse_qsl, unquote_plus, urlencode, urlsplit, urlunsplit +from uuid import RFC_4122, UUID + +from cryptography.hazmat.primitives.asymmetric.rsa import RSAPublicKey, RSAPublicNumbers +from cryptography.hazmat.primitives.ciphers.aead import AESGCM +from cryptography.hazmat.primitives.hashes import SHA256 +from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC + + +class OidcUtil: + """ + 统一认证公共工具类。 + """ + + _PROTOCOL_DESCRIPTION_PATTERN = re.compile(r'[\x20-\x21\x23-\x5B\x5D-\x7E]+') + + _ERROR_MESSAGES = { + 'invalid_request': '认证请求无效', + 'invalid_client': '客户端认证失败', + 'invalid_grant': '授权凭据无效或已过期', + 'unauthorized_client': '客户端无权使用当前授权方式', + 'unsupported_grant_type': '不支持当前授权类型', + 'invalid_scope': '请求的权限范围无效', + 'invalid_target': '请求的资源受众无效', + 'unsupported_response_type': '不支持当前授权响应类型', + 'unsupported_response_mode': '不支持当前授权响应模式', + 'access_denied': '授权请求被拒绝', + 'server_error': '认证服务处理失败', + 'temporarily_unavailable': '认证服务暂不可用', + 'invalid_token': '访问令牌无效', + 'insufficient_scope': '访问令牌的权限不足', + 'interaction_required': '需要重新完成认证交互', + 'login_required': '需要登录后继续', + 'consent_required': '需要用户确认授权', + 'account_selection_required': '需要选择登录账号', + 'not_found': '请求的认证资源不存在', + } + + _DEFAULT_PROTOCOL_DESCRIPTIONS = { + 'invalid_request': 'Invalid request', + 'invalid_client': 'Client authentication failed', + 'invalid_grant': 'The authorization grant is invalid or expired', + 'unauthorized_client': 'Client is not authorized', + 'unsupported_grant_type': 'Grant type is not supported', + 'invalid_scope': 'Requested scope is invalid', + 'invalid_target': 'Requested resource is invalid', + 'unsupported_response_type': 'Response type is not supported', + 'unsupported_response_mode': 'Response mode is not supported', + 'access_denied': 'Access is denied', + 'server_error': 'Authorization service is unavailable', + 'temporarily_unavailable': 'Authorization service is temporarily unavailable', + 'invalid_token': 'Invalid access token', + 'insufficient_scope': 'Access token scope is insufficient', + 'interaction_required': 'Interaction is required', + 'login_required': 'Login is required', + 'consent_required': 'Consent is required', + 'account_selection_required': 'Account selection is required', + 'not_found': 'Resource is unavailable', + } + + _DESCRIPTION_MESSAGES = { + 'A current login is required': '需要有效的当前登录状态', + 'A new authorization decision is required': '需要重新确认授权', + 'A resource audience is required for resource scope': '申请资源权限时必须指定资源受众', + 'A transaction coordinator is required': '缺少事务协调器', + 'Access to this application is blocked': '当前用户已被禁止访问此应用', + 'Application is unavailable; restart login': '应用已不可用,请返回应用重新登录', + 'Application permissions have changed; restart login': '应用权限已变更,请返回应用重新登录', + 'Authentication audit service is unavailable': '认证审计服务不可用', + 'Authentication service is unavailable': '认证服务不可用', + 'Authorization audit service is unavailable': '授权审计服务暂不可用', + 'Authorization code could not be completed': '授权码兑换未能完成', + 'Authorization code could not be created': '授权码创建失败', + 'Authorization code is invalid': '授权码无效', + 'Authorization code is invalid or expired': '授权码无效或已过期', + 'Authorization code is not allowed for this client': '当前客户端未启用授权码模式', + 'Authorization code state is invalid': '授权码状态无效', + 'Authorization completion failed': '完成授权失败', + 'Authorization parameters must be in the form body': '授权参数必须全部放在表单请求体中', + 'Authorization server is not ready': '认证中心尚未就绪', + 'Authorization service is unavailable': '授权服务暂不可用', + 'CSRF validation failed': 'CSRF 校验失败,请重新发起认证', + 'Client PKCE policy is weaker than provider policy': '客户端 PKCE 配置不符合认证中心的安全要求', + 'Client authentication failed': '客户端认证失败', + 'Client is inactive': '客户端已停用或不可用', + 'Client is not registered': '客户端未注册或已停用', + 'Consent is required': '需要用户确认授权', + 'Current login could not be verified': '无法验证当前登录状态', + 'Duplicate authorization parameter': '授权请求包含重复参数', + 'Interaction TTL is invalid': '认证交互有效期无效', + 'Interaction cache callback failed': '认证交互缓存更新回调执行失败', + 'Interaction could not be created': '认证交互创建失败', + 'Interaction has already completed': '认证交互已完成,请勿重复提交', + 'Interaction is invalid or expired': '认证交互无效或已过期', + 'Interaction is missing or expired': '认证交互不存在或已过期', + 'Interaction is no longer active': '认证交互已失效,请重新发起认证', + 'Interaction is not complete': '认证交互尚未完成', + 'Interaction state has changed': '认证交互状态已变更,请刷新后重试', + 'Interaction state is invalid': '认证交互状态无效', + 'Interaction transition failed': '认证交互状态更新失败', + 'Invalid authorization form': '授权请求表单无效', + 'Invalid authorization request': '授权请求无效', + 'Invalid interaction request': '认证交互请求无效', + 'Invalid request': '请求参数无效', + 'Invalid token request': '令牌请求无效', + 'Invalid validated redirect URI': '已验证的回调地址无效', + 'Login is required': '需要登录后继续', + 'OIDC provider is disabled': '统一认证服务未启用', + 'Only PKCE S256 is supported': '仅支持 PKCE S256 校验方式', + 'Only authorization code is supported': '仅支持授权码响应类型 code', + 'OpenID scope requires the sub claim': 'openid 权限必须包含用户主体声明 sub', + 'PKCE S256 code_challenge must contain 43 characters': 'PKCE S256 的 code_challenge 必须为 43 个字符', + 'Requested resource is not allowed for this client': '当前客户端无权访问所请求的资源', + 'Requested scope does not belong to the selected resource': '请求的权限范围不属于所选资源', + 'Requested scope is not allowed for this client': '当前客户端不允许申请所请求的权限范围', + 'Requested scope is not authorized': '请求的权限范围尚未获准', + 'Required scope cannot be removed': '必需的权限范围不能取消', + 'Resource is unavailable': '资源不可用', + 'Resource policy has changed': '资源访问策略已变更,请重新授权', + 'Response mode is not supported': '不支持当前授权响应模式', + 'Scope policy has changed': '权限范围策略已变更,请重新授权', + 'Submitted scope is not allowed': '提交的权限范围不在允许列表内', + 'The authorization grant is invalid or expired': '授权凭据无效或已过期', + 'The authorization grant is no longer valid': '原授权已失效,请重新授权', + 'Token endpoint is unavailable': '令牌服务暂不可用', + 'Token issuance is unavailable': '令牌签发服务暂不可用', + 'Token policy is unavailable': '令牌策略不可用', + 'Too many authorization requests': '授权请求过于频繁,请稍后重试', + 'User denied the authorization request': '用户已拒绝授权请求', + 'User identity mapping is unavailable': '用户身份映射暂不可用', + 'User identity version could not be updated': '用户身份安全版本更新失败', + 'Validated redirect URI is invalid': '已校验的回调地址无效', + 'Validated redirect URI is unavailable': '已校验的回调地址不可用', + 'client_id and redirect_uri are required': '必须提供 client_id 和 redirect_uri', + 'max_age must be non-negative': 'max_age 不得为负数', + 'nonce is required for OpenID Connect authorization': 'OpenID Connect 授权请求必须提供 nonce', + 'openid scope is required': '授权请求必须包含 openid 权限', + 'prompt contains an unsupported combination': 'prompt 包含不支持的参数组合', + 'redirect_uri is not registered': '登录回调地址 redirect_uri 未注册', + } + + _PROTOCOL_DESCRIPTIONS = {message: description for description, message in _DESCRIPTION_MESSAGES.items()} + + _SENSITIVE_KEY = re.compile( + r'(?:token|code|secret|cookie|password|passwd|credential|authorization|private[_-]?key|' + r'pkce|verifier|nonce|client[_-]?assertion|assertion|access[_-]?token|refresh[_-]?token|' + r'id[_-]?token|state)', + re.IGNORECASE, + ) + + _SENSITIVE_VALUE = re.compile( + r'(?i)(?:bearer\s+|basic\s+|(?:access|refresh|id)?[_-]?(?:token|secret|code|cookie|password|verifier)\s*[=:])' + ) + + _OPAQUE_CAPABILITY = re.compile( + r'(?i)(? str: + """ + 返回中文异常诊断;未知协议描述按错误码提供安全的中文说明。 + + :param error: OAuth 错误码 + :param description: 原始协议描述或中文诊断 + :return: 中文异常说明 + """ + if description: + translated = cls._DESCRIPTION_MESSAGES.get(description) + if translated is not None: + return translated + if any('\u4e00' <= char <= '\u9fff' for char in description): + return description + return cls._ERROR_MESSAGES.get(error, '认证请求处理失败') + + @classmethod + def protocol_error_description(cls, error: str, description: str | None) -> str | None: + """ + 保留合法 ASCII 协议描述,将中文或不安全字符转换为协议允许的说明。 + + :param error: OAuth 错误码 + :param description: 原始异常说明 + :return: 符合 OAuth 字符约束的描述或 None + """ + if not description: + return None + if cls._PROTOCOL_DESCRIPTION_PATTERN.fullmatch(description): + return description + return cls._PROTOCOL_DESCRIPTIONS.get( + description, cls._DEFAULT_PROTOCOL_DESCRIPTIONS.get(error, 'Authentication request failed') + ) + + # 基础数据、批量参数与认证声明 + + @staticmethod + def read_field(user: Any, name: str, default: Any = None) -> Any: + """ + 从 ORM 用户对象或字典读取字段 + + :param user: SysUser ORM 记录或用户 Claim 字段映射 + :param name: 字段名称 + :param default: 字段缺省值 + :return: 读取的字段值 + """ + + if isinstance(user, Mapping): + return user.get(name, default) + return getattr(user, name, default) + + @staticmethod + def positive_int(value: Any, field: str) -> int: + """ + 校验正整数标量 + + :param value: 待校验的正整数 + :param field: 字段名称 + :return: 校验后的正整数 + """ + + if not isinstance(value, int) or isinstance(value, bool) or value <= 0: + raise ValueError(f'{field} 必须为正整数') + return value + + @staticmethod + def nonnegative_int(value: Any, field: str) -> int: + """ + 校验非负整数标量 + + :param value: 待校验的非负整数 + :param field: 字段名称 + :return: 校验后的非负整数 + """ + + if not isinstance(value, int) or isinstance(value, bool) or value < 0: + raise ValueError(f'{field} 必须为非负整数') + return value + + @staticmethod + def nonempty_string(value: Any, field: str, limit: int) -> str: + """ + 校验非空字符串及长度 + + :param value: 待校验的字符串 + :param field: 字段名称 + :param limit: 字符串的最大长度 + :return: 校验后的字符串 + """ + + if not isinstance(value, str) or not value or len(value) > limit: + raise ValueError(f'{field} 必须为非空字符串') + return value + + @classmethod + def string_list(cls, value: Any, field: str, limit: int) -> list[str]: + """ + 校验字符串列表并复制为 JSON 安全列表 + + :param value: 待校验的字符串列表 + :param field: 字段名称 + :param limit: 字符串列表的最大长度 + :return: 校验后的字符串列表 + """ + + if not isinstance(value, (list, tuple)) or len(value) > limit: + raise ValueError(f'{field} 必须为符合长度限制的字符串列表') + return [cls.nonempty_string(item, field, 500) for item in value] + + @staticmethod + def json_list(value: Any) -> list[Any]: + """ + 将 Scope 或 Resource 字段规范化为列表 + + :param value: JSON 编码的 Scope 或 Resource 列表值 + :return: Scope 或 Resource 元素列表;输入不是 JSON 数组时返回空列表 + """ + + return list(value) if isinstance(value, (list, tuple)) else [] + + @staticmethod + def json_object_pairs(items: list[tuple[str, object]]) -> dict[str, object]: + """ + 构造拒绝重复字段的 JSON 对象 + + :param items: JSON 对象字段 + :return: 唯一字段组成的对象 + :raises ValueError: JSON 对象包含重复字段 + """ + + result: dict[str, object] = {} + for key, value in items: + # 拒绝重复字段,避免不同解析器产生歧义 + if key in result: + raise ValueError('JSON 请求体包含重复字段') + result[key] = value + return result + + @staticmethod + def serialize_json(record: Mapping[str, Any], *, error_message: str) -> str: + """ + 生成稳定紧凑的 JSON,并拒绝不可序列化对象和非有限数字。 + + :param record: 需要缓存的标量载荷 + :param error_message: 调用场景的中文错误说明 + :return: 按键排序的 JSON 字符串 + """ + try: + return json.dumps(record, ensure_ascii=False, separators=(',', ':'), sort_keys=True, allow_nan=False) + except (TypeError, ValueError) as exc: + raise ValueError(error_message) from exc + + @staticmethod + def json_etag(payload: Mapping[str, Any]) -> str: + """ + 为 JWKS 响应计算稳定的强 ETag + + :param payload: 公开响应 JSON 映射 + :return: 带双引号的 SHA-256 ETag + """ + + canonical = json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(',', ':')).encode() + + return '"' + hashlib.sha256(canonical).hexdigest() + '"' + + @staticmethod + def actor_name(value: Any, *, error_message: str = '操作者不能为空') -> str: + """ + 校验操作者名称并按管理字段长度截断,保留原始空白语义。 + + :param value: 待校验的操作者名称 + :param error_message: 名称不可用时的中文说明 + :return: 最多 64 个字符的操作者名称 + """ + if not isinstance(value, str) or not value.strip(): + raise ValueError(error_message) + return value[:64] + + @staticmethod + def split_batch(value: str, field_name: str, *, max_size: int, validate_all_first: bool = False) -> list[str]: + """ + 拆分批量路径标识,拒绝空值、路径字符、空白和重复项。 + + :param value: 逗号分隔的参数 + :param field_name: 错误说明中的字段名称 + :param max_size: 允许的最大项目数量 + :param validate_all_first: 是否先校验全部格式再检查重复,保留调用场景的报错顺序 + :return: 去除首尾空白后的标识列表 + """ + items = value.split(',') if isinstance(value, str) else [] + if not items or len(items) > max_size: + raise ValueError(f'{field_name} 参数无效') + values = [item.strip() for item in items] + + def invalid(item: str) -> bool: + return not item or '%' in item or '/' in item or '\\' in item or any(char.isspace() for char in item) + + if validate_all_first and any(invalid(item) for item in values): + raise ValueError(f'{field_name} 参数无效') + result: list[str] = [] + for item in values: + if invalid(item): + raise ValueError(f'{field_name} 参数无效') + if item in result: + raise ValueError(f'{field_name} 参数重复') + result.append(item) + return result + + @staticmethod + def unique_codes(values: Iterable[str], field_name: str) -> list[str]: + """ + 校验权限或资源编码的边界空白,并保持输入顺序和重复项错误。 + + :param values: 待校验的编码集合 + :param field_name: 中文错误说明中的字段名称 + :return: 保持原序的编码列表 + """ + result: list[str] = [] + for value in values: + if not isinstance(value, str) or not value.strip() or value.strip() != value: + raise ValueError(f'{field_name} 包含无效的标识') + if value in result: + raise ValueError(f'{field_name} 不得包含重复项') + result.append(value) + return result + + @staticmethod + def normalize_user_ids(values: Iterable[int], *, max_size: int) -> tuple[int, ...]: + """ + 校验、去重并排序用户 ID,确保批量锁顺序稳定 + + :param values: 待校验的用户 ID 可迭代集合 + :param max_size: 去重后允许的最大用户数量 + :return: 排序去重后的用户 ID 元组 + """ + + if isinstance(values, (str, bytes)): + raise ValueError('用户编号列表必须全部为整数') + result = tuple(sorted(set(values))) + if not result or len(result) > max_size: + raise ValueError('用户编号列表数量无效') + if any(not isinstance(item, int) or isinstance(item, bool) or item <= 0 for item in result): + raise ValueError('用户编号列表必须全部为正整数') + return result + + @staticmethod + def is_trimmed_identifier(value: str, *, max_length: int) -> bool: + """ + 检查标识长度及首尾空白,保留内部字符。 + + :param value: 待校验的 Session ID + :param max_length: 标识最大长度 + :return: Session ID 是否有效 + """ + + return isinstance(value, str) and 1 <= len(value) <= max_length and value.strip() == value + + @staticmethod + def scope_set(scopes: str | Iterable[str] | None) -> set[str]: + """ + 将空格分隔或集合形式的 Scope 规范化 + + :param scopes: 请求的 Scope 集合 + :return: 规范化 Scope 集合 + """ + + if scopes is None: + return set() + if isinstance(scopes, str): + return {item for item in scopes.split() if item} + return {item for item in scopes if isinstance(item, str) and item} + + @staticmethod + def normalize_audiences(value: Any) -> list[str]: + """ + 将 aud 声明规范化为去重列表 + + :param value: OAuth Token 的 Audience 声明值 + :return: 去重后的 Audience 列表;格式非法时返回空列表 + """ + + values = [value] if isinstance(value, str) else value + if not isinstance(values, list) or not values or any(not isinstance(item, str) or not item for item in values): + return [] + return list(dict.fromkeys(values)) + + @staticmethod + def token_audiences(issuer: str, resources: Sequence[str]) -> list[str]: + """ + 将 aud 声明规范化为去重列表 + + :param issuer: OIDC issuer URL + :param resources: 已校验的 Resource audience 字符串序列 + :return: 包含 userinfo audience 和 Resource audience 的去重列表 + """ + + result = [f'{issuer.rstrip("/")}{OidcUtil._USERINFO_AUDIENCE_SUFFIX}'] + result.extend(resources) + + return list(dict.fromkeys(result)) + + @staticmethod + def claim_numeric_date(value: Any) -> int | float | None: + """ + 将用户更新时间转换为 NumericDate + + :param value: 用户更新时间 + :return: NumericDate 时间戳或 None + """ + + if value is None or isinstance(value, (int, float)): + return value + timestamp = getattr(value, 'timestamp', None) + + return int(timestamp()) if callable(timestamp) else None + + @staticmethod + def numeric_date(value: datetime | None) -> int | None: + """ + 将 datetime 转换为有限的 Unix 时间戳 + + :param value: 可选的 Token 签发时间或过期时间 + :return: 有限的 Unix 时间戳;输入为空或时间戳非有限时返回 None + """ + + if value is None: + return None + timestamp = value.timestamp() + + return int(timestamp) if math.isfinite(timestamp) else None + + @staticmethod + def requires_reauthentication(auth_time: datetime, max_age: int | None, now: datetime | None = None) -> bool: + """ + 纯函数判断 SSO 认证是否超过 max_age + + :param auth_time: SSO 认证时间 + :param max_age: 请求的最大认证年龄 + :param now: 可选当前时间 + :return: 超过 max_age 时为 True + """ + + # 时间工具依赖基础异常,延迟导入避免异常模块初始化时形成循环。 + from utils.time_util import TimezoneUtil # noqa: PLC0415 + + if max_age is None: + return False + if not isinstance(max_age, int) or isinstance(max_age, bool) or max_age < 0: + raise ValueError('max_age 必须为非负整数') + current = TimezoneUtil.to_utc(now) if now is not None else TimezoneUtil.utc_now() + auth = TimezoneUtil.to_utc(auth_time) + + return current.timestamp() - auth.timestamp() > max_age + + # URI、路径标识与交互参数 + + @staticmethod + def validate_path_identifier(value: str, field_name: str) -> str: + """ + 校验可安全放入单一路径段的管理标识 + + :param value: 待校验的资源或权限标识 + :param field_name: 校验失败时展示的字段名称 + :return: 校验通过的管理标识 + """ + + if ( + not isinstance(value, str) + or not value + or not value[0].isalnum() + or any(not (char.isascii() and (char.isalnum() or char in '._:-')) for char in value) + ): + raise ValueError(f'{field_name} 必须符合路径标识的字符要求') + return value + + @classmethod + def is_valid_kid(cls, value: Any) -> bool: + """ + 检查签名密钥标识的字符集和长度。 + + :param value: 待检查的标识 + :return: 是否满足原有安全路径规则 + """ + return isinstance(value, str) and cls._SAFE_KID_PATTERN.fullmatch(value) is not None + + @staticmethod + def validate_registered_uri(uri_type: str, value: str) -> str: + """ + 校验注册 URI + + :param uri_type: 注册地址类型 + :param value: 待校验的完整注册地址 + :return: 校验通过的原始注册地址 + """ + + if '*' in value: + raise ValueError('地址不得包含通配符') + parsed = urlsplit(value) + if parsed.scheme not in {'http', 'https'} or not parsed.netloc: + raise ValueError('地址必须使用 HTTP 或 HTTPS 协议并包含主机名') + if parsed.username is not None or parsed.password is not None: + raise ValueError('地址不得包含用户名或密码') + if parsed.fragment: + raise ValueError('地址不得包含片段标识 fragment') + if uri_type == 'backchannel_logout' and parsed.query: + raise ValueError('后端退出通知地址不得包含查询参数') + try: + hostname = parsed.hostname + _ = parsed.port + except ValueError as exc: + raise ValueError('地址中的主机名或端口无效') from exc + if not hostname: + raise ValueError('地址必须包含主机名') + if parsed.scheme == 'http' and hostname.lower() not in OidcUtil._DEVELOPMENT_HTTP_HOSTS: + raise ValueError('HTTP 地址仅允许用于本机开发环境') + if uri_type == 'cors_origin' and (parsed.path or parsed.query): + raise ValueError('跨域来源只能包含协议、主机名和可选端口') + return value + + @staticmethod + def validate_resource_audience(audience: str, *, max_length: int) -> str: + """ + 校验 Resource audience URI + + :param audience: Resource audience + :param max_length: URI 最大长度 + :return: 校验后的 Resource audience URI + :raises ValueError: audience 不是无用户信息的绝对 HTTPS URI 时抛出 + """ + + if not isinstance(audience, str) or not audience or len(audience) > max_length: + raise ValueError('资源受众必须为非空地址') + parsed = urlsplit(audience) + if parsed.scheme != 'https' or not parsed.netloc or parsed.fragment: + raise ValueError('资源受众必须是无片段标识的绝对 HTTPS 地址') + if parsed.username is not None or parsed.password is not None: + raise ValueError('资源受众地址不得包含用户名或密码') + try: + hostname = parsed.hostname + _ = parsed.port + except ValueError as exc: + raise ValueError('资源受众地址的主机名或端口无效') from exc + if not hostname: + raise ValueError('资源受众地址必须包含主机名') + return audience + + @staticmethod + def is_safe_post_logout_uri(uri: Any) -> bool: + """ + 验证退出后重定向 URI 的安全格式 + + :param uri: 回调 URI + :return: URI 是否满足退出后重定向安全约束 + """ + + if not isinstance(uri, str) or not uri or len(uri) > OidcUtil._MAX_URI_LENGTH or '*' in uri: + return False + try: + parsed = urlsplit(uri) + hostname = parsed.hostname + _ = parsed.port + except ValueError: + return False + return bool( + parsed.scheme in {'http', 'https'} + and hostname + and parsed.netloc + and (parsed.scheme == 'https' or hostname.lower() in OidcUtil._DEVELOPMENT_HTTP_HOSTS) + and not parsed.username + and not parsed.password + and not parsed.fragment + ) + + @staticmethod + def is_public_ip(value: str) -> bool: + """ + 判断地址是否不属于回环、内网、链路本地或保留网段 + + :param value: IPv4 或 IPv6 文本地址 + :return: 仅当地址可作为公网目标时返回 ``True`` + """ + + try: + address = ipaddress.ip_address(value) + except ValueError: + return False + return not ( + address.is_private + or address.is_loopback + or address.is_link_local + or address.is_reserved + or address.is_multicast + or address.is_unspecified + ) + + @staticmethod + def parse_backchannel_uri(uri: object) -> tuple[str, int] | None: + """ + 解析并校验 Back-Channel URI 的非网络安全边界 + + :param uri: 待校验的 URI + :return: 合法时返回主机名和端口,否则返回 ``None`` + """ + + if not isinstance(uri, str) or not uri or len(uri) > OidcUtil._MAX_URI_LENGTH or '*' in uri: + return None + try: + parsed = urlsplit(uri) + hostname = parsed.hostname + username = parsed.username + password = parsed.password + port = parsed.port + except ValueError: + return None + if ( + parsed.scheme != 'https' + or not parsed.netloc + or not hostname + or username + or password + or parsed.query + or parsed.fragment + ): + return None + try: + address = ipaddress.ip_address(hostname) + except ValueError: + address = None + if address is not None and not OidcUtil.is_public_ip(hostname): + return None + return hostname, port or 443 + + @staticmethod + def replace_query_parameters( + uri: str, + parameters: Sequence[tuple[str, str]], + replaced_fields: Collection[str], + *, + fragment: str | None = None, + doseq: bool = False, + ) -> str: + """ + 替换指定查询参数,同时保留其他参数的顺序、重复项和空值。 + + 此方法只负责 URL 编码,调用方必须先完成回调注册和地址安全校验。 + + :param uri: 原始地址 + :param parameters: 按顺序追加的参数 + :param replaced_fields: 必须从原地址移除的参数名称 + :param fragment: 显式指定片段,None 表示保留原片段 + :param doseq: 是否展开参数中的序列值 + :return: 重新编码后的地址 + """ + parsed = urlsplit(uri) + query = [ + (key, value) for key, value in parse_qsl(parsed.query, keep_blank_values=True) if key not in replaced_fields + ] + query.extend(parameters) + return urlunsplit( + ( + parsed.scheme, + parsed.netloc, + parsed.path, + urlencode(query, doseq=doseq), + parsed.fragment if fragment is None else fragment, + ) + ) + + @staticmethod + def append_state(redirect_uri: str, state: str | None) -> str: + """ + 为已验证的退出回调地址附加状态参数 + + :param redirect_uri: 服务端已验证的退出回调地址 + :param state: 原始退出状态参数 + :return: 附加状态参数后的回调地址 + """ + + if state is None: + return redirect_uri + return OidcUtil.replace_query_parameters(redirect_uri, [('state', state)], {'state'}) + + @staticmethod + def interaction_url(base: str, interaction_id: str, csrf_token: str) -> str: + """ + 构建只把原始 CSRF 放入 URL Fragment 的交互地址 + + :param base: 交互地址基址 + :param interaction_id: 交互流程标识 + :param csrf_token: CSRF Token + :return: 交互页面 URL + """ + + parsed = urlsplit(base) + query = dict(parse_qsl(parsed.query, keep_blank_values=True)) + query['interaction'] = interaction_id + + return urlunsplit((parsed.scheme, parsed.netloc, parsed.path, urlencode(query), f'csrf={csrf_token}')) + + @staticmethod + def normalize_prompt(value: str | None) -> str | None: + """ + 校验并规范化认证交互提示,拒绝重复值和非法组合。 + + :param value: 空格分隔的 prompt 参数 + :return: 保持原顺序的规范化参数或 None + """ + if value is None: + return None + prompts = value.split() + if ( + not prompts + or len(set(prompts)) != len(prompts) + or any(item not in {'login', 'consent', 'none'} for item in prompts) + or ('none' in prompts and len(prompts) > 1) + ): + raise ValueError('prompt 只能包含支持的值,且 none 不能与其他值组合') + return ' '.join(prompts) + + @classmethod + def is_s256_challenge(cls, value: Any) -> bool: + """ + 检查无填充 Base64URL 格式的 SHA-256 PKCE 挑战值。 + + :param value: 待检查的挑战值 + :return: 是否恰好包含 43 个允许的 ASCII 字符 + """ + return isinstance(value, str) and cls._S256_CHALLENGE.fullmatch(value) is not None + + # 编码、摘要与凭据 + + @staticmethod + def base64url_encode(value: bytes) -> str: + """ + 将字节编码为无填充的 Base64URL 文本 + + :param value: 待编码字节 + :return: 无填充的 Base64URL 文本 + """ + + return base64.urlsafe_b64encode(value).rstrip(b'=').decode('ascii') + + @staticmethod + def sha256_digest(value: str | bytes) -> str: + """ + 计算原始内容的 SHA-256 十六进制摘要,不规范化或截断输入。 + + :param value: URI、一次性凭据等原始内容 + :return: 小写十六进制摘要 + """ + raw = value.encode('utf-8') if isinstance(value, str) else value + return hashlib.sha256(raw).hexdigest() + + @staticmethod + def sha256_hex_digest(value: str | bytes) -> str: + """ + 校验已生成的 SHA-256 十六进制摘要 + + :param value: 摘要文本或 ASCII 字节 + :return: 小写十六进制摘要 + :raises TypeError: 输入不是 str 或 bytes + :raises ValueError: 摘要长度或字符集不正确 + """ + + if not isinstance(value, (str, bytes)): + raise TypeError('摘要必须为字符串或字节数据') + if isinstance(value, bytes): + try: + value = value.decode('ascii') + except UnicodeDecodeError as exc: + raise ValueError('摘要必须是 SHA-256 十六进制字符串') from exc + if len(value) != OidcUtil._SHA256_HEX_LENGTH: + raise ValueError('摘要必须是 SHA-256 十六进制字符串') + try: + bytes.fromhex(value) + except ValueError as exc: + raise ValueError('摘要必须是 SHA-256 十六进制字符串') from exc + return value.lower() + + @staticmethod + def hmac_sha256(value: str | bytes, key: str | bytes) -> str: + """ + 对已由调用方校验的内容和密钥计算 HMAC-SHA256 摘要。 + + :param value: 原始字符串或字节数据 + :param key: 原始字符串或字节形式的密钥 + :return: 小写十六进制摘要 + """ + raw_value = value.encode('utf-8') if isinstance(value, str) else value + raw_key = key.encode('utf-8') if isinstance(key, str) else key + return hmac.new(raw_key, raw_value, hashlib.sha256).hexdigest() + + @staticmethod + def access_token_hash(access_token: str) -> str: + """ + 计算 OIDC at_hash 值 + + :param access_token: 待计算 OIDC at_hash 的 ASCII JWT Access Token 文本 + :return: SHA-256 前 128 bit 摘要的无填充 Base64URL 字符串 + """ + + digest = hashlib.sha256(access_token.encode('ascii')).digest()[:16] + + return base64.urlsafe_b64encode(digest).rstrip(b'=').decode('ascii') + + @staticmethod + def csrf_digest(token: str, pepper: str) -> str: + """ + 使用独立 Pepper 生成 CSRF HMAC 摘要 + + :param token: 原始 CSRF Token + :param pepper: Token 摘要 Pepper + :return: CSRF 摘要 + """ + + if ( + not isinstance(token, str) + or not isinstance(pepper, str) + or len(pepper.encode()) < OidcUtil._MIN_PEPPER_BYTES + ): + raise ValueError('CSRF 摘要密钥至少需要 32 字节') + return OidcUtil.hmac_sha256(token, pepper) + + @classmethod + def credential_proof( + cls, interaction_id: str, user_id: int, subject_id: str, auth_version: int, *, pepper: str + ) -> str: + """ + 生成绑定交互、用户、主体和安全版本的凭据证明摘要。 + + :param interaction_id: 交互标识 + :param user_id: 本地用户编号 + :param subject_id: 稳定身份主体标识 + :param auth_version: 身份安全版本 + :param pepper: 调用方提供的认证摘要密钥 + :return: 凭据证明摘要 + """ + return cls.hmac_sha256(f'{interaction_id}:{user_id}:{subject_id}:{auth_version}', pepper) + + @classmethod + def logout_confirmation_digest(cls, sso_cookie: str | None, browser_nonce: str, pepper: str | bytes) -> str: + """ + 计算退出确认凭据的浏览器绑定摘要。 + + :param sso_cookie: 首次退出请求携带的会话 Cookie + :param browser_nonce: 绑定浏览器的一次性随机数 + :param pepper: 调用方提供的认证摘要密钥 + :return: 浏览器绑定摘要 + """ + return cls.hmac_sha256((sso_cookie or '') + '\\0' + browser_nonce, pepper) + + @classmethod + def logout_rate_scope(cls, address: Any, pepper: str | bytes) -> str: + """ + 生成退出限流主体摘要,在认证配置尚未就绪时使用原有固定回退值。 + + :param address: 客户端地址或匿名占位值 + :param pepper: 调用方读取的认证摘要密钥 + :return: 限流主体摘要 + """ + raw_pepper = pepper.encode() if isinstance(pepper, str) else pepper + if not isinstance(raw_pepper, bytes) or len(raw_pepper) < cls._MIN_PEPPER_BYTES: + raw_pepper = hashlib.sha256(b'oidc-logout-rate-limit-fallback').digest() + return cls.hmac_sha256(str(address), raw_pepper) + + @classmethod + def redis_key_component(cls, value: str | int, *, name: str = 'Redis Key 组件') -> str: + """ + 校验可作为 Key 路径组件的标识符 + + :param value: 待校验的标识符 + :param name: 错误信息中的字段名称 + :return: 原样返回的安全组件 + :raises ValueError: 组件为空或包含路径/控制字符 + """ + + component = str(value) + if not cls._COMPONENT_PATTERN.fullmatch(component): + raise ValueError(f'{name} 包含不允许的字符') + return component + + @classmethod + def hash_sensitive_identifier(cls, value: str | bytes, pepper: str | bytes) -> str: + """ + 使用独立 Pepper 对敏感标识生成 HMAC-SHA256 摘要 + + :param value: 用户名、IP 等敏感标识 + :param pepper: 独立于 JWT/传输加密密钥的 Pepper,至少 32 bytes + :return: 64 位小写 HMAC-SHA256 摘要 + :raises TypeError: 标识或 Pepper 类型错误 + :raises ValueError: Pepper 为空或长度不足 + """ + + if not isinstance(value, (str, bytes)) or not isinstance(pepper, (str, bytes)): + raise TypeError('敏感标识和摘要密钥必须为字符串或字节数据') + raw_value = value.encode('utf-8') if isinstance(value, str) else value + raw_pepper = pepper.encode('utf-8') if isinstance(pepper, str) else pepper + if len(raw_pepper) < cls._MIN_PEPPER_BYTES: + raise ValueError('敏感标识摘要密钥至少需要 32 字节') + return cls.hmac_sha256(raw_value, raw_pepper) + + @staticmethod + def parse_basic_credentials(authorization: str) -> tuple[str, str]: + """ + 解析 RFC 7617 Basic Header,不接受请求体 Secret 回退 + + :param authorization: Authorization Header + :return: 解码后的 client_id 与 client_secret + :raises ValueError: Header、Base64 或凭据格式不合法 + """ + + if not isinstance(authorization, str) or authorization[:6].lower() != 'basic ': + raise ValueError('客户端认证信息无效') + encoded = authorization[6:].strip() + if not encoded: + raise ValueError('客户端认证信息无效') + try: + decoded = base64.b64decode(encoded, validate=True).decode('utf-8') + except (binascii.Error, UnicodeDecodeError, ValueError): + raise ValueError('客户端认证信息无效') from None + if ':' not in decoded: + raise ValueError('客户端认证信息无效') + client_id, client_secret = decoded.split(':', 1) + # OAuth 表单客户端认证使用 application/x-www-form-urlencoded 约定 + # 该约定同时正确处理 client_id 中的空格与 Secret 中的百分号编码冒号 + client_id, client_secret = unquote_plus(client_id), unquote_plus(client_secret) + if not client_id or not client_secret: + raise ValueError('客户端认证信息无效') + # RFC 7617 user-id 不能包含冒号,但分割后的 Secret 可以包含冒号 + if any(ord(char) < OidcUtil._ASCII_CONTROL_LIMIT for char in client_id): + raise ValueError('客户端认证信息无效') + return client_id, client_secret + + @staticmethod + def generate_client_id() -> str: + """ + 生成带固定前缀的随机客户端标识。 + + :return: cli_ 前缀的客户端标识 + """ + return f'cli_{secrets.token_urlsafe(24)}' + + @staticmethod + def generate_client_secret() -> str: + """ + 生成仅应展示一次的高熵 Client Secret + + :return: 带 cs1 类型前缀的随机 Secret + """ + + return f'cs1.{secrets.token_urlsafe(32)}' + + @staticmethod + def client_secret_hashes(client: Any, explicit: Iterable[str] | None = None) -> list[str]: + """ + 收集 Client 当前可用的 Secret 哈希 + + :param client: Client 映射或对象 + :param explicit: 调用方显式提供的 Secret 哈希集合 + :return: 规范化后的 Secret 哈希列表 + """ + + values = explicit + if values is None: + values = OidcUtil.read_field(client, 'secret_hashes') + if values is None: + values = OidcUtil.read_field(client, 'secrets') + if values is None: + one = OidcUtil.read_field(client, 'secret_hash') + values = [one] if one else [] + result: list[str] = [] + for item in values: + if isinstance(item, str): + result.append(item) + else: + value = OidcUtil.read_field(item, 'secret_hash') + if isinstance(value, str): + result.append(value) + return result + + # 会话数据与审计脱敏 + + @staticmethod + def is_rfc4122_uuid(value: str) -> bool: + """ + 判断字符串是否为规范的 RFC 4122 UUID + + :param value: 待检查的 UUID 字符串 + :return: 是否为规范的 RFC 4122 UUID 字符串 + """ + + try: + parsed = UUID(value) + except (TypeError, ValueError, AttributeError): + return False + return str(parsed) == value and parsed.variant == RFC_4122 + + @staticmethod + def parse_sso_cookie(cookie: str) -> tuple[str, str]: + """ + 拆分并验证 SSO Cookie 中的 Session 标识和随机 Secret + + :param cookie: Cookie 原文 + :return: ``(sid, secret)`` + :raises ValueError: Cookie 格式、UUID 或 Secret 长度不合法 + """ + + if not isinstance(cookie, str): + raise ValueError('会话 Cookie 必须为字符串') + parts = cookie.split('.') + if len(parts) != OidcUtil._COOKIE_PARTS or parts[0] != OidcUtil._COOKIE_PREFIX: + raise ValueError('会话 Cookie 格式无效') + sid, secret = parts[1], parts[2] + if not OidcUtil.is_rfc4122_uuid(sid): + raise ValueError('会话 Cookie 中的 sid 格式不规范') + if len(secret) != OidcUtil._COOKIE_SECRET_TEXT_LENGTH or any( + char not in 'ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789_-' for char in secret + ): + raise ValueError('会话 Cookie 密钥无效') + try: + decoded = base64.urlsafe_b64decode(secret + '===') + except (binascii.Error, ValueError) as exc: + raise ValueError('会话 Cookie 密钥无效') from exc + if len(decoded) != OidcUtil._COOKIE_SECRET_BYTES: + raise ValueError('会话 Cookie 密钥必须为 256 位') + return sid, secret + + @staticmethod + def session_pepper_bytes(pepper: str | bytes) -> bytes: + """ + 验证并转换 Token Pepper 字节串 + + :param pepper: OIDC Token Pepper + :return: Pepper 字节 + :raises ValueError: Pepper 类型或长度不合法 + """ + + if not isinstance(pepper, (str, bytes)): + raise ValueError('会话摘要密钥必须为字符串或字节数据') + raw = pepper.encode() if isinstance(pepper, str) else pepper + if len(raw) < OidcUtil._MIN_PEPPER_BYTES: + raise ValueError('会话摘要密钥至少需要 32 字节') + return raw + + @classmethod + def session_secret_digest(cls, secret: str | bytes, pepper: str | bytes) -> str: + """ + 使用 Pepper 计算 Cookie Secret 的 HMAC 摘要 + + :param secret: Cookie Secret 原文,仅在请求内存中存在 + :param pepper: OIDC Token Pepper + :return: 十六进制 HMAC 摘要 + :raises ValueError: secret 不是 str/bytes,或 pepper 不是 str/bytes、编码后少于 32 字节 + """ + + raw_secret = secret.encode() if isinstance(secret, str) else secret + if not isinstance(raw_secret, bytes): + raise ValueError('会话密钥必须为字节数据或字符串') + return cls.hmac_sha256(raw_secret, cls.session_pepper_bytes(pepper)) + + @classmethod + def session_user_agent_digest(cls, user_agent: str | None, pepper: str | bytes) -> str | None: + """ + 使用 Pepper 计算 User-Agent 的 HMAC 摘要 + + :param user_agent: 请求 User-Agent + :param pepper: OIDC Token Pepper + :return: User-Agent 摘要或 None + :raises ValueError: User-Agent 不是字符串、超过 500 个字符,或 pepper 编码后少于 32 字节 + """ + + if user_agent is None: + return None + if not isinstance(user_agent, str) or len(user_agent) > OidcUtil._MAX_USER_AGENT_LENGTH: + raise ValueError('客户端标识 User-Agent 无效') + return cls.session_secret_digest(user_agent, pepper) + + @staticmethod + def looks_like_jwt(value: str) -> bool: + """ + 判断字符串是否具有 JWT 三段式秘密外观 + + :param value: 待检测的字符串 + :return: 是否符合 JWT 外形 + """ + + parts = value.split('.') + + return len(parts) == OidcUtil._JWT_PART_COUNT and all( + OidcUtil._JWT_PART_MIN_LENGTH <= len(part) <= OidcUtil._JWT_PART_MAX_LENGTH for part in parts + ) + + @classmethod + def sanitize_audit_detail(cls, detail: Any) -> Any: + """ + 递归删除审计详情中的敏感字段 + + :param detail: 待记录的任意可序列化值 + :return: 只包含安全字段的可序列化值 + """ + + if isinstance(detail, Mapping): + return { + str(key): cls.sanitize_audit_detail(value) + for key, value in detail.items() + if not cls._SENSITIVE_KEY.search(str(key)) + } + if isinstance(detail, list): + return [cls.sanitize_audit_detail(value) for value in detail] + if isinstance(detail, tuple): + return [cls.sanitize_audit_detail(value) for value in detail] + if isinstance(detail, str): + if ( + cls._SENSITIVE_VALUE.search(detail) + or cls.looks_like_jwt(detail) + or cls._OPAQUE_CAPABILITY.search(detail) + ): + return '[REDACTED]' + return detail[:1024] + if isinstance(detail, (int, float, bool)) or detail is None: + return detail + return None + + # RSA 公钥与签名私钥 + + @staticmethod + def rsa_public_jwk(public_key: RSAPublicKey, kid: str) -> dict[str, str]: + """ + 将 RSA 公钥编码为仅包含公开参数的 RS256 JWK。 + + :param public_key: RSA 公钥 + :param kid: 已由调用方校验的密钥标识 + :return: 用于签名验证的公开 JWK + """ + numbers = public_key.public_numbers() + return { + 'kty': 'RSA', + 'use': 'sig', + 'kid': kid, + 'alg': 'RS256', + 'n': OidcUtil.base64url_encode(numbers.n.to_bytes((numbers.n.bit_length() + 7) // 8, 'big')), + 'e': OidcUtil.base64url_encode(numbers.e.to_bytes((numbers.e.bit_length() + 7) // 8, 'big')), + } + + @classmethod + def normalize_public_jwk( + cls, record: Any, *, min_rsa_bits: int = 2048, min_exponent: int = 3, max_exponent: int = 2**32 + ) -> dict[str, str]: + """ + 提取并校验公开 RSA JWK + + :param record: 数据库签名密钥记录或包含 public_jwk 字段的映射 + :param min_rsa_bits: 最小 RSA 模数位数 + :param min_exponent: 最小公开指数 + :param max_exponent: 公开指数上限(不包含) + :return: 只包含公开字段的 JWK + :raises ValueError: JWK 结构或算法不合法 + """ + + value = getattr(record, 'public_jwk', record) + if not isinstance(value, Mapping): + raise ValueError('公开 JWK 必须为对象') + if ( + getattr(record, 'key_use', 'sig') != 'sig' + or getattr(record, 'alg', 'RS256') != 'RS256' + or value.get('kty') != 'RSA' + or value.get('use', 'sig') != 'sig' + or value.get('alg', 'RS256') != 'RS256' + ): + raise ValueError('仅支持使用 RS256 签名的 RSA JWK') + record_kid = getattr(record, 'kid', None) + kid = value.get('kid', record_kid) + if not cls.is_valid_kid(kid): + raise ValueError('签名密钥标识 kid 包含不允许的字符') + if record_kid is not None and kid != record_kid: + raise ValueError('公开 JWK 的 kid 与数据库密钥标识不匹配') + result = {field: value.get(field) for field in ('kty', 'use', 'kid', 'alg', 'n', 'e')} + result.update({'kty': 'RSA', 'use': 'sig', 'kid': kid, 'alg': 'RS256'}) + if ( + not isinstance(result.get('n'), str) + or not result['n'] + or not isinstance(result.get('e'), str) + or not result['e'] + ): + raise ValueError('公开 JWK 必须包含 n 和 e') + try: + n_bytes = base64.urlsafe_b64decode(result['n'] + '=' * (-len(result['n']) % 4)) + e_bytes = base64.urlsafe_b64decode(result['e'] + '=' * (-len(result['e']) % 4)) + except (TypeError, ValueError, binascii.Error) as exc: + raise ValueError('公开 JWK 的 n/e 必须使用 Base64URL 编码') from exc + if ( + not n_bytes + or not e_bytes + or n_bytes[0] == 0 + or cls.base64url_encode(n_bytes) != result['n'] + or cls.base64url_encode(e_bytes) != result['e'] + ): + raise ValueError('公开 JWK 的 n/e 编码无效') + modulus = int.from_bytes(n_bytes, 'big') + exponent = int.from_bytes(e_bytes, 'big') + if ( + modulus.bit_length() < min_rsa_bits + or exponent < min_exponent + or exponent >= max_exponent + or exponent % 2 == 0 + ): + raise ValueError('公开 JWK 的 RSA 参数无效') + try: + RSAPublicNumbers(exponent, modulus).public_key() + except ValueError as exc: + raise ValueError('公开 JWK 的 RSA 参数无效') from exc + return result + + @staticmethod + def derive_signing_encryption_key(material: bytes, salt: bytes) -> bytes: + """ + 使用随机盐派生私钥加密密钥 + + :param material: 部署配置中的加密主密钥 + :param salt: 单条密钥记录的随机盐 + :return: AES-GCM 使用的 256 位密钥 + """ + + return PBKDF2HMAC( + algorithm=SHA256(), + length=32, + salt=salt, + iterations=OidcUtil._ENCRYPTION_KDF_ITERATIONS, + ).derive(material) + + @classmethod + def encrypt_signing_private_key(cls, pem: bytes, material: bytes, *, salt: bytes, nonce: bytes) -> str: + """ + 按现有 v2 格式封装使用 AES-GCM 加密的签名私钥。 + + :param pem: 私钥 PEM 字节 + :param material: 调用方校验后的部署加密主密钥 + :param salt: 调用方生成的 16 字节随机盐 + :param nonce: 调用方生成的 12 字节随机数 + :return: 带 v2 前缀的私钥密文 + """ + encrypted = AESGCM(cls.derive_signing_encryption_key(material, salt)).encrypt(nonce, pem, None) + return 'v2.' + base64.urlsafe_b64encode(b'salt' + salt + nonce + encrypted).decode() + + @classmethod + def decrypt_signing_private_key(cls, ciphertext: str, material: bytes) -> bytes: + """ + 解密现有 v1/v2 私钥密文,保留历史数据格式兼容性。 + + :param ciphertext: 带版本前缀的私钥密文 + :param material: 部署加密主密钥 + :return: 已解密的私钥 PEM 字节 + """ + if ciphertext.startswith('v2.'): + envelope = base64.urlsafe_b64decode(ciphertext.removeprefix('v2.')) + salt, nonce, encrypted = envelope[4:20], envelope[20:32], envelope[32:] + key = cls.derive_signing_encryption_key(material, salt) + elif ciphertext.startswith('v1.'): + envelope = base64.urlsafe_b64decode(ciphertext.removeprefix('v1.')) + nonce, encrypted = envelope[:12], envelope[12:] + key = hashlib.sha256(material).digest() + else: + raise ValueError('不支持当前签名私钥密文格式') + return AESGCM(key).decrypt(nonce, encrypted, None) diff --git a/ruoyi-fastapi-backend/utils/time_util.py b/ruoyi-fastapi-backend/utils/time_util.py index 94fee27c9..5ca9d5e9e 100644 --- a/ruoyi-fastapi-backend/utils/time_util.py +++ b/ruoyi-fastapi-backend/utils/time_util.py @@ -157,6 +157,16 @@ def to_utc(cls, value: datetime) -> datetime: """ return cls.ensure_aware(value).astimezone(timezone.utc) + @classmethod + def to_optional_utc(cls, value: datetime | None) -> datetime | None: + """ + 将可选的带时区时刻转换为 UTC,空值原样保留。 + + :param value: 可选的带时区时刻 + :return: UTC 时刻或 None + """ + return cls.to_utc(value) if value is not None else None + @classmethod def to_utc_milliseconds(cls, value: datetime) -> datetime: """ diff --git a/ruoyi-fastapi-frontend/.editorconfig b/ruoyi-fastapi-frontend/.editorconfig new file mode 100644 index 000000000..e26ca0fe7 --- /dev/null +++ b/ruoyi-fastapi-frontend/.editorconfig @@ -0,0 +1,15 @@ +root = true + +[*] +charset = utf-8 +indent_style = space +indent_size = 2 +end_of_line = lf +insert_final_newline = true +trim_trailing_whitespace = true + +[*.md] +trim_trailing_whitespace = false + +[*.{bat,cmd}] +end_of_line = crlf diff --git a/ruoyi-fastapi-frontend/.gitattributes b/ruoyi-fastapi-frontend/.gitattributes new file mode 100644 index 000000000..62781c223 --- /dev/null +++ b/ruoyi-fastapi-frontend/.gitattributes @@ -0,0 +1,6 @@ +# Keep frontend text consistent across Windows, macOS and Linux. +* text=auto eol=lf + +# Windows command scripts use CRLF. +*.bat text eol=crlf +*.cmd text eol=crlf diff --git a/ruoyi-fastapi-frontend/.prettierignore b/ruoyi-fastapi-frontend/.prettierignore new file mode 100644 index 000000000..c5db8e337 --- /dev/null +++ b/ruoyi-fastapi-frontend/.prettierignore @@ -0,0 +1,29 @@ +# Keep the HTML entry DOCTYPE and formatting manually maintained. +/index.html +html/ie.html + +# Dependencies, build output and caches +node_modules/ +dist/ +coverage/ +.cache/ +.vite/ + +# Test output +tests/**/coverage/ +tests/e2e/reports/ +playwright-report/ +test-results/ + +# Package-manager output and generated declarations +package-lock.json +pnpm-lock.yaml +yarn.lock +auto-imports.d.ts +components.d.ts + +# Generated or external assets +**/*.min.js +**/*.min.css +**/*.svg +**/*.map diff --git a/ruoyi-fastapi-frontend/.prettierrc.json b/ruoyi-fastapi-frontend/.prettierrc.json new file mode 100644 index 000000000..8f3fe7b3b --- /dev/null +++ b/ruoyi-fastapi-frontend/.prettierrc.json @@ -0,0 +1,17 @@ +{ + "printWidth": 100, + "tabWidth": 2, + "useTabs": false, + "semi": false, + "singleQuote": true, + "quoteProps": "as-needed", + "trailingComma": "es5", + "bracketSpacing": true, + "bracketSameLine": false, + "arrowParens": "always", + "htmlWhitespaceSensitivity": "css", + "vueIndentScriptAndStyle": false, + "singleAttributePerLine": true, + "endOfLine": "lf", + "proseWrap": "preserve" +} diff --git a/ruoyi-fastapi-frontend/FORMATTING.md b/ruoyi-fastapi-frontend/FORMATTING.md new file mode 100644 index 000000000..bdf9c0884 --- /dev/null +++ b/ruoyi-fastapi-frontend/FORMATTING.md @@ -0,0 +1,122 @@ +# 管理端前端格式化规范 + +本规范适用于 `ruoyi-fastapi-frontend`,覆盖 Vue 单文件组件、JavaScript / TypeScript、CSS / SCSS / Less、HTML、JSON、Markdown 和 YAML,包括 `src`、`plugins`、`vite`、`tests` 以及自行维护的 `public` 脚本。 + +## 参考项目与工具 + +参考 [Element Plus 官方项目的 Prettier 配置](https://github.com/element-plus/element-plus/blob/dev/.prettierrc)(核对日期:2026-09-14)。当前管理端使用 Vue 3 和 Element Plus,采用其 `semi: false`、`singleQuote: true`、`trailingComma: 'es5'` 作为基础风格,适合现有 JavaScript 代码。 + +格式化统一使用 **Prettier 3.6.2**,以 `.prettierrc.json` 为规则来源。该版本与项目已有的本地安装及锁文件一致,并在 `devDependencies` 中固定版本,避免开发者因格式化器版本不同产生差异。 + +## 具体规则 + +| 项目 | 规则 | 说明 | +| --------------------- | ---------------------------------- | --------------------------------------------------------------------- | +| 缩进 | 2 个空格,不使用 Tab | 模板按嵌套层级缩进 | +| 行宽 | `printWidth: 100` | 本项目适配;作为自动换行参考宽度,长字符串等可以超出 | +| JavaScript 字符串 | 优先单引号 | 沿用 Element Plus;必要时由 Prettier 选择更少转义的引号 | +| HTML / Vue 属性、JSON | 双引号 | JavaScript 的单引号设置不改变这些语法的属性引号 | +| JavaScript 分号 | 不写行末分号 | 沿用 Element Plus;Prettier 会保留避免自动分号插入歧义所需的分号 | +| 尾逗号 | `trailingComma: 'es5'` | 沿用 Element Plus;多行对象、数组等保留尾逗号,函数参数不添加尾逗号 | +| 箭头函数 | 参数始终带括号 | `(item) => item.id`,便于增加参数、默认值或类型 | +| 对象字面量 | 大括号内侧加空格 | `{ name: 'demo' }`,属性名仅在语法需要时加引号 | +| Vue / HTML 多属性标签 | 每个属性独占一行 | 本项目适配;减少表单、表格和权限指令挤在同一行的情况 | +| 多行标签的 `>` | 不紧跟最后一个属性 | 使用 `bracketSameLine: false` | +| Vue 的 script / style | 内容不额外缩进一层 | `vueIndentScriptAndStyle: false` | +| HTML 空白 | `htmlWhitespaceSensitivity: 'css'` | 保留 Prettier 默认的空白处理方式,兼顾行内文本空格语义 | +| 文本文件 | UTF-8、LF、末尾换行 | `.editorconfig` 与 `.gitattributes` 配合;Windows 批处理脚本使用 CRLF | +| 行尾空格 | 删除 | Markdown 例外,保留用于硬换行的两个空格 | +| Markdown 段落 | 保留原有换行 | `proseWrap: 'preserve'`,避免中文段落被反复重排 | + +除上述项目适配外,使用 Prettier 默认行为。选项含义见 [Prettier 官方文档](https://prettier.io/docs/options)。 + +Vue 单文件组件建议沿用当前项目的 `