style: auto-format with ruff [skip ci]
This commit is contained in:
parent
76b5156a2c
commit
b7579c539d
@ -8,7 +8,7 @@ import os
|
|||||||
import secrets
|
import secrets
|
||||||
import time
|
import time
|
||||||
import urllib.parse
|
import urllib.parse
|
||||||
from typing import Any, Optional
|
from typing import Any
|
||||||
from urllib.parse import urlencode
|
from urllib.parse import urlencode
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@ -17,14 +17,14 @@ from fastapi.responses import RedirectResponse
|
|||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.exc import IntegrityError
|
from sqlalchemy.exc import IntegrityError
|
||||||
|
|
||||||
from server.utils.auth_utils import AuthUtils
|
|
||||||
from server.utils.common_utils import log_operation
|
|
||||||
from yuxi.repositories.user_repository import UserRepository
|
from yuxi.repositories.user_repository import UserRepository
|
||||||
from yuxi.storage.postgres.models_business import Department, User
|
from yuxi.storage.postgres.models_business import Department, User
|
||||||
from yuxi.utils.datetime_utils import utc_now_naive
|
from yuxi.utils.datetime_utils import utc_now_naive
|
||||||
from yuxi.utils.logging_config import logger
|
from yuxi.utils.logging_config import logger
|
||||||
|
|
||||||
|
from server.utils.auth_utils import AuthUtils
|
||||||
|
from server.utils.common_utils import log_operation
|
||||||
|
|
||||||
# 前端 OIDC 回调路由路径(与 web/src/router/index.js 中的路由保持一致)
|
# 前端 OIDC 回调路由路径(与 web/src/router/index.js 中的路由保持一致)
|
||||||
FRONTEND_CALLBACK_PATH = "/auth/oidc/callback"
|
FRONTEND_CALLBACK_PATH = "/auth/oidc/callback"
|
||||||
# 登录页路径
|
# 登录页路径
|
||||||
@ -106,11 +106,11 @@ class OIDCProviderMetadata:
|
|||||||
"""OIDC Provider 元数据"""
|
"""OIDC Provider 元数据"""
|
||||||
|
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self.authorization_endpoint: Optional[str] = None
|
self.authorization_endpoint: str | None = None
|
||||||
self.token_endpoint: Optional[str] = None
|
self.token_endpoint: str | None = None
|
||||||
self.userinfo_endpoint: Optional[str] = None
|
self.userinfo_endpoint: str | None = None
|
||||||
self.end_session_endpoint: Optional[str] = None
|
self.end_session_endpoint: str | None = None
|
||||||
self.last_error: Optional[str] = None
|
self.last_error: str | None = None
|
||||||
self._loaded = False
|
self._loaded = False
|
||||||
|
|
||||||
async def load(self, issuer_url: str) -> bool:
|
async def load(self, issuer_url: str) -> bool:
|
||||||
@ -150,12 +150,12 @@ class OIDCProviderMetadata:
|
|||||||
class OIDCUtils:
|
class OIDCUtils:
|
||||||
"""OIDC 工具类"""
|
"""OIDC 工具类"""
|
||||||
|
|
||||||
_metadata: Optional[OIDCProviderMetadata] = None
|
_metadata: OIDCProviderMetadata | None = None
|
||||||
_state_store: dict[str, dict[str, Any]] = {}
|
_state_store: dict[str, dict[str, Any]] = {}
|
||||||
_login_code_store: dict[str, dict[str, Any]] = {}
|
_login_code_store: dict[str, dict[str, Any]] = {}
|
||||||
_state_ttl_seconds = 300
|
_state_ttl_seconds = 300
|
||||||
_login_code_ttl_seconds = 60
|
_login_code_ttl_seconds = 60
|
||||||
_last_metadata_error: Optional[str] = None
|
_last_metadata_error: str | None = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _cleanup_expired_state(cls) -> None:
|
def _cleanup_expired_state(cls) -> None:
|
||||||
@ -172,7 +172,7 @@ class OIDCUtils:
|
|||||||
cls._login_code_store.pop(key, None)
|
cls._login_code_store.pop(key, None)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def get_metadata(cls) -> Optional[OIDCProviderMetadata]:
|
async def get_metadata(cls) -> OIDCProviderMetadata | None:
|
||||||
"""获取 OIDC Provider 元数据"""
|
"""获取 OIDC Provider 元数据"""
|
||||||
if not oidc_config.enabled or not oidc_config.is_configured():
|
if not oidc_config.enabled or not oidc_config.is_configured():
|
||||||
cls._last_metadata_error = "OIDC 未启用或基础配置不完整"
|
cls._last_metadata_error = "OIDC 未启用或基础配置不完整"
|
||||||
@ -203,7 +203,7 @@ class OIDCUtils:
|
|||||||
return cls._metadata
|
return cls._metadata
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_last_metadata_error(cls) -> Optional[str]:
|
def get_last_metadata_error(cls) -> str | None:
|
||||||
"""获取最近一次 OIDC 元数据加载错误"""
|
"""获取最近一次 OIDC 元数据加载错误"""
|
||||||
return cls._last_metadata_error
|
return cls._last_metadata_error
|
||||||
|
|
||||||
@ -219,7 +219,7 @@ class OIDCUtils:
|
|||||||
return state
|
return state
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def verify_state(cls, state: str) -> Optional[dict[str, Any]]:
|
def verify_state(cls, state: str) -> dict[str, Any] | None:
|
||||||
"""验证 state 参数"""
|
"""验证 state 参数"""
|
||||||
state_data = cls._state_store.pop(state, None)
|
state_data = cls._state_store.pop(state, None)
|
||||||
if not state_data:
|
if not state_data:
|
||||||
@ -240,7 +240,7 @@ class OIDCUtils:
|
|||||||
return code
|
return code
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def consume_login_code(cls, code: str) -> Optional[dict[str, Any]]:
|
def consume_login_code(cls, code: str) -> dict[str, Any] | None:
|
||||||
"""消费一次性短期登录 code"""
|
"""消费一次性短期登录 code"""
|
||||||
data = cls._login_code_store.pop(code, None)
|
data = cls._login_code_store.pop(code, None)
|
||||||
if not data:
|
if not data:
|
||||||
@ -255,7 +255,7 @@ class OIDCUtils:
|
|||||||
return secrets.token_urlsafe(32)
|
return secrets.token_urlsafe(32)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def build_authorization_url(cls, redirect_path: str = "/") -> Optional[str]:
|
async def build_authorization_url(cls, redirect_path: str = "/") -> str | None:
|
||||||
"""构建授权 URL"""
|
"""构建授权 URL"""
|
||||||
metadata = await cls.get_metadata()
|
metadata = await cls.get_metadata()
|
||||||
if not metadata or not metadata.authorization_endpoint:
|
if not metadata or not metadata.authorization_endpoint:
|
||||||
@ -281,7 +281,7 @@ class OIDCUtils:
|
|||||||
return f"{metadata.authorization_endpoint}?{query_string}"
|
return f"{metadata.authorization_endpoint}?{query_string}"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def exchange_code_for_token(cls, code: str) -> Optional[dict[str, Any]]:
|
async def exchange_code_for_token(cls, code: str) -> dict[str, Any] | None:
|
||||||
"""用授权码交换令牌"""
|
"""用授权码交换令牌"""
|
||||||
metadata = await cls.get_metadata()
|
metadata = await cls.get_metadata()
|
||||||
if not metadata or not metadata.token_endpoint:
|
if not metadata or not metadata.token_endpoint:
|
||||||
@ -313,7 +313,7 @@ class OIDCUtils:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def get_userinfo(cls, access_token: str) -> Optional[dict[str, Any]]:
|
async def get_userinfo(cls, access_token: str) -> dict[str, Any] | None:
|
||||||
"""获取用户信息"""
|
"""获取用户信息"""
|
||||||
metadata = await cls.get_metadata()
|
metadata = await cls.get_metadata()
|
||||||
if not metadata or not metadata.userinfo_endpoint:
|
if not metadata or not metadata.userinfo_endpoint:
|
||||||
@ -334,7 +334,7 @@ class OIDCUtils:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def build_logout_url(cls, id_token: Optional[str] = None) -> Optional[str]:
|
async def build_logout_url(cls, id_token: str | None = None) -> str | None:
|
||||||
"""构建登出 URL"""
|
"""构建登出 URL"""
|
||||||
metadata = await cls.get_metadata()
|
metadata = await cls.get_metadata()
|
||||||
if not metadata or not metadata.end_session_endpoint:
|
if not metadata or not metadata.end_session_endpoint:
|
||||||
@ -418,9 +418,7 @@ async def find_user_by_oidc_sub(db, sub: str) -> User | None:
|
|||||||
return user
|
return user
|
||||||
|
|
||||||
legacy_result = await db.execute(
|
legacy_result = await db.execute(
|
||||||
select(User)
|
select(User).filter(User.user_id.like(f"{oidc_user_id}:%"), User.is_deleted == 0).order_by(User.id.asc())
|
||||||
.filter(User.user_id.like(f"{oidc_user_id}:%"), User.is_deleted == 0)
|
|
||||||
.order_by(User.id.asc())
|
|
||||||
)
|
)
|
||||||
legacy_users = list(legacy_result.scalars().all())
|
legacy_users = list(legacy_result.scalars().all())
|
||||||
if legacy_users:
|
if legacy_users:
|
||||||
@ -441,9 +439,7 @@ async def find_deleted_oidc_user_by_sub(db, sub: str) -> User | None:
|
|||||||
return deleted_user
|
return deleted_user
|
||||||
|
|
||||||
legacy_result = await db.execute(
|
legacy_result = await db.execute(
|
||||||
select(User)
|
select(User).filter(User.user_id.like(f"{oidc_user_id}:%"), User.is_deleted == 1).order_by(User.id.asc())
|
||||||
.filter(User.user_id.like(f"{oidc_user_id}:%"), User.is_deleted == 1)
|
|
||||||
.order_by(User.id.asc())
|
|
||||||
)
|
)
|
||||||
return legacy_result.scalar_one_or_none()
|
return legacy_result.scalar_one_or_none()
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user