ForcePilot/backend/server/utils/auth_utils.py

115 lines
3.8 KiB
Python
Raw Normal View History

2025-05-02 23:56:59 +08:00
import os
import secrets
from datetime import timedelta
from typing import Any
2025-05-02 23:56:59 +08:00
import jwt
from argon2 import PasswordHasher
from argon2.exceptions import InvalidHash, VerifyMismatchError, VerificationError
from yuxi.utils.datetime_utils import utc_now
2025-05-02 23:56:59 +08:00
JWT_ALGORITHM = "HS256"
JWT_EXPIRATION = 7 * 24 * 60 * 60 # 7天过期
JWT_AUDIENCE = "yuxi-know-api"
PUBLIC_DEFAULT_JWT_SECRET_KEY = "yuxi_know_secure_key"
PASSWORD_HASHER = PasswordHasher()
2025-05-02 23:56:59 +08:00
def _is_production_env() -> bool:
return os.environ.get("YUXI_ENV", "development").strip().lower() in {"prod", "production"}
def _get_or_create_dev_env(name: str, value_factory) -> str:
value = os.environ.get(name, "").strip()
if value:
return value
if _is_production_env():
raise ValueError(f"{name} 未配置,请在生产环境的 .env.prod 中设置持久化随机值")
value = value_factory()
os.environ[name] = value
print(f"{name} 未配置,开发环境已自动生成临时随机值,服务重启后会重新生成。")
return value
def _get_jwt_secret_key() -> str:
secret_key = _get_or_create_dev_env("JWT_SECRET_KEY", lambda: secrets.token_hex(32))
if _is_production_env() and secret_key == PUBLIC_DEFAULT_JWT_SECRET_KEY:
raise ValueError("JWT_SECRET_KEY 不能使用公开默认密钥,请重新生成随机强密钥")
return secret_key
def _get_jwt_issuer() -> str:
instance_id = _get_or_create_dev_env("YUXI_INSTANCE_ID", lambda: f"instance-{secrets.token_hex(8)}")
return f"yuxi-know:{instance_id}"
2025-05-02 23:56:59 +08:00
class AuthUtils:
"""认证工具类"""
@staticmethod
def hash_password(password: str) -> str:
"""使用 Argon2 哈希密码"""
return PASSWORD_HASHER.hash(password)
2025-05-02 23:56:59 +08:00
@staticmethod
def verify_password(stored_password: str, provided_password: str) -> bool:
"""验证密码"""
if not stored_password.startswith("$argon2"):
return False
try:
return PASSWORD_HASHER.verify(stored_password, provided_password)
except (InvalidHash, VerifyMismatchError, VerificationError):
2025-05-02 23:56:59 +08:00
return False
@staticmethod
def create_access_token(data: dict[str, Any], expires_delta: timedelta | None = None) -> str:
2025-05-02 23:56:59 +08:00
"""创建JWT访问令牌"""
to_encode = data.copy()
# 设置过期时间
if expires_delta:
expire = utc_now() + expires_delta
2025-05-02 23:56:59 +08:00
else:
expire = utc_now() + timedelta(seconds=JWT_EXPIRATION)
2025-05-02 23:56:59 +08:00
to_encode.update({"exp": expire, "iss": _get_jwt_issuer(), "aud": JWT_AUDIENCE})
2025-05-02 23:56:59 +08:00
# 编码JWT
encoded_jwt = jwt.encode(to_encode, _get_jwt_secret_key(), algorithm=JWT_ALGORITHM)
2025-05-02 23:56:59 +08:00
return encoded_jwt
@staticmethod
def decode_token(token: str) -> dict[str, Any] | None:
2025-05-02 23:56:59 +08:00
"""解码验证JWT令牌"""
try:
payload = jwt.decode(
token,
_get_jwt_secret_key(),
algorithms=[JWT_ALGORITHM],
issuer=_get_jwt_issuer(),
audience=JWT_AUDIENCE,
options={"require": ["exp", "sub", "iss", "aud"]},
)
2025-05-02 23:56:59 +08:00
return payload
except (jwt.PyJWTError, ValueError):
2025-05-02 23:56:59 +08:00
return None
@staticmethod
def verify_access_token(token: str) -> dict[str, Any]:
2025-05-02 23:56:59 +08:00
"""验证访问令牌,如果无效则抛出异常"""
try:
payload = jwt.decode(
token,
_get_jwt_secret_key(),
algorithms=[JWT_ALGORITHM],
issuer=_get_jwt_issuer(),
audience=JWT_AUDIENCE,
options={"require": ["exp", "sub", "iss", "aud"]},
)
2025-05-02 23:56:59 +08:00
return payload
except jwt.ExpiredSignatureError:
raise ValueError("令牌已过期")
except jwt.InvalidTokenError:
raise ValueError("无效的令牌")