ForcePilot/backend/package/yuxi/utils/auth_utils.py

97 lines
3.4 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
2026-05-29 22:19:58 +08:00
from argon2.exceptions import InvalidHash, VerificationError, VerifyMismatchError
from yuxi.utils.datetime_utils import utc_now
2025-05-02 23:56:59 +08:00
JWT_ALGORITHM = "HS256"
2026-05-29 22:19:58 +08:00
JWT_EXPIRATION = 7 * 24 * 60 * 60
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:
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
to_encode = data.copy()
2026-05-29 22:19:58 +08:00
expire = utc_now() + (expires_delta or timedelta(seconds=JWT_EXPIRATION))
to_encode.update({"exp": expire, "iss": _get_jwt_issuer(), "aud": JWT_AUDIENCE})
2026-05-29 22:19:58 +08:00
return jwt.encode(to_encode, _get_jwt_secret_key(), algorithm=JWT_ALGORITHM)
2025-05-02 23:56:59 +08:00
@staticmethod
def decode_token(token: str) -> dict[str, Any] | None:
2025-05-02 23:56:59 +08:00
try:
2026-05-29 22:19:58 +08:00
return jwt.decode(
token,
_get_jwt_secret_key(),
algorithms=[JWT_ALGORITHM],
issuer=_get_jwt_issuer(),
audience=JWT_AUDIENCE,
options={"require": ["exp", "sub", "iss", "aud"]},
)
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:
2026-05-29 22:19:58 +08:00
return 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
except jwt.ExpiredSignatureError:
raise ValueError("令牌已过期")
except jwt.InvalidTokenError:
raise ValueError("无效的令牌")