74feca4467
- replace plaintext login and unlock requests with RSA-OAEP/AES-GCM encrypted payloads - add login challenge replay protection, production RSA key validation, and auth tests - wire compose to environment-driven dev/prod settings without committing local secrets - update setup-config smoke scripts and Postman docs for encrypted login - add visit schedule migrations/tests and update study/subject setup workflows
139 lines
4.9 KiB
Python
139 lines
4.9 KiB
Python
import base64
|
|
import json
|
|
import secrets
|
|
from dataclasses import dataclass
|
|
from datetime import datetime, timedelta, timezone
|
|
from typing import Any
|
|
|
|
from cryptography.hazmat.primitives import hashes, serialization
|
|
from cryptography.hazmat.primitives.asymmetric import padding, rsa
|
|
from cryptography.hazmat.primitives.asymmetric.rsa import RSAPrivateKey, RSAPublicKey
|
|
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
|
|
from pydantic import BaseModel, EmailStr, Field, ValidationError
|
|
|
|
from app.core.config import settings
|
|
|
|
|
|
class DecryptedLoginPayload(BaseModel):
|
|
email: EmailStr
|
|
password: str = Field(min_length=1, max_length=72)
|
|
challenge: str = Field(min_length=16)
|
|
|
|
|
|
class EncryptedLoginEnvelope(BaseModel):
|
|
encrypted_key: str = Field(min_length=1)
|
|
iv: str = Field(min_length=1)
|
|
data: str = Field(min_length=1)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class LoginChallenge:
|
|
value: str
|
|
expires_at: datetime
|
|
|
|
|
|
_private_key: RSAPrivateKey | None = None
|
|
_public_key_pem: str | None = None
|
|
_challenges: dict[str, LoginChallenge] = {}
|
|
|
|
|
|
def _normalize_pem(value: str) -> str:
|
|
return value.replace("\\n", "\n").strip()
|
|
|
|
|
|
def _load_or_create_private_key() -> RSAPrivateKey:
|
|
configured_key = settings.LOGIN_RSA_PRIVATE_KEY
|
|
if configured_key:
|
|
key = serialization.load_pem_private_key(_normalize_pem(configured_key).encode("utf-8"), password=None)
|
|
if not isinstance(key, RSAPrivateKey):
|
|
raise ValueError("LOGIN_RSA_PRIVATE_KEY 必须是 RSA 私钥")
|
|
return key
|
|
if settings.ENV == "production":
|
|
raise ValueError("生产环境必须配置 LOGIN_RSA_PRIVATE_KEY")
|
|
return rsa.generate_private_key(public_exponent=65537, key_size=2048)
|
|
|
|
|
|
def get_private_key() -> RSAPrivateKey:
|
|
global _private_key
|
|
if _private_key is None:
|
|
_private_key = _load_or_create_private_key()
|
|
return _private_key
|
|
|
|
|
|
def validate_login_crypto_configuration() -> None:
|
|
if settings.ENV == "production" and not settings.LOGIN_RSA_PRIVATE_KEY:
|
|
raise ValueError("生产环境必须配置 LOGIN_RSA_PRIVATE_KEY")
|
|
|
|
|
|
def get_public_key_pem() -> str:
|
|
global _public_key_pem
|
|
if _public_key_pem is not None:
|
|
return _public_key_pem
|
|
if settings.LOGIN_RSA_PUBLIC_KEY:
|
|
if not settings.LOGIN_RSA_PRIVATE_KEY:
|
|
raise ValueError("配置 LOGIN_RSA_PUBLIC_KEY 时必须同时配置 LOGIN_RSA_PRIVATE_KEY")
|
|
_public_key_pem = _normalize_pem(settings.LOGIN_RSA_PUBLIC_KEY)
|
|
return _public_key_pem
|
|
public_key: RSAPublicKey = get_private_key().public_key()
|
|
_public_key_pem = public_key.public_bytes(
|
|
encoding=serialization.Encoding.PEM,
|
|
format=serialization.PublicFormat.SubjectPublicKeyInfo,
|
|
).decode("utf-8")
|
|
return _public_key_pem
|
|
|
|
|
|
def create_login_challenge() -> LoginChallenge:
|
|
prune_expired_challenges()
|
|
while len(_challenges) >= settings.LOGIN_CHALLENGE_MAX_ACTIVE:
|
|
oldest = next(iter(_challenges))
|
|
_challenges.pop(oldest, None)
|
|
value = secrets.token_urlsafe(32)
|
|
expires_at = datetime.now(timezone.utc) + timedelta(seconds=settings.LOGIN_CHALLENGE_TTL_SECONDS)
|
|
challenge = LoginChallenge(value=value, expires_at=expires_at)
|
|
_challenges[value] = challenge
|
|
return challenge
|
|
|
|
|
|
def prune_expired_challenges() -> None:
|
|
now = datetime.now(timezone.utc)
|
|
expired = [value for value, challenge in _challenges.items() if challenge.expires_at <= now]
|
|
for value in expired:
|
|
_challenges.pop(value, None)
|
|
|
|
|
|
def consume_login_challenge(value: str) -> bool:
|
|
prune_expired_challenges()
|
|
challenge = _challenges.pop(value, None)
|
|
if not challenge:
|
|
return False
|
|
return challenge.expires_at > datetime.now(timezone.utc)
|
|
|
|
|
|
def decrypt_login_payload(*, key_id: str, challenge: str, ciphertext: str) -> DecryptedLoginPayload | None:
|
|
if key_id != settings.LOGIN_RSA_KEY_ID:
|
|
return None
|
|
if not consume_login_challenge(challenge):
|
|
return None
|
|
try:
|
|
envelope_data = json.loads(base64.b64decode(ciphertext, validate=True).decode("utf-8"))
|
|
envelope = EncryptedLoginEnvelope.model_validate(envelope_data)
|
|
encrypted_key = base64.b64decode(envelope.encrypted_key, validate=True)
|
|
iv = base64.b64decode(envelope.iv, validate=True)
|
|
encrypted_payload = base64.b64decode(envelope.data, validate=True)
|
|
aes_key = get_private_key().decrypt(
|
|
encrypted_key,
|
|
padding.OAEP(
|
|
mgf=padding.MGF1(algorithm=hashes.SHA256()),
|
|
algorithm=hashes.SHA256(),
|
|
label=None,
|
|
),
|
|
)
|
|
decrypted = AESGCM(aes_key).decrypt(iv, encrypted_payload, None)
|
|
payload: Any = json.loads(decrypted.decode("utf-8"))
|
|
parsed = DecryptedLoginPayload.model_validate(payload)
|
|
except (ValueError, TypeError, json.JSONDecodeError, ValidationError):
|
|
return None
|
|
if parsed.challenge != challenge:
|
|
return None
|
|
return parsed
|