| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280 |
- from __future__ import annotations
- import hashlib
- import hmac
- import logging
- import secrets
- from datetime import datetime, timedelta
- from typing import Any
- from sqlalchemy.exc import IntegrityError
- from supply_infra.config import get_infra_settings
- from supply_infra.db import get_session
- from supply_infra.db.models.auth_user import AuthUser
- from supply_infra.db.repositories.auth_repo import AuthRepository
- from supply_infra.pipeline.dates import CHINA_TIMEZONE
- logger = logging.getLogger(__name__)
- SESSION_COOKIE_NAME = "supply_session"
- ADMIN_ROLE = "admin"
- USER_ROLE = "user"
- ACTIVE_STATUS = "active"
- DISABLED_STATUS = "disabled"
- VALID_ROLES = {ADMIN_ROLE, USER_ROLE}
- VALID_STATUSES = {ACTIVE_STATUS, DISABLED_STATUS}
- _SCRYPT_N = 2**14
- _SCRYPT_R = 8
- _SCRYPT_P = 1
- _DUMMY_PASSWORD_HASH: str | None = None
- class InvalidCredentialsError(Exception):
- """Raised for every login rejection to avoid exposing account state."""
- class DuplicateUsernameError(Exception):
- """Raised when a local username already exists."""
- def _now() -> datetime:
- return datetime.now(CHINA_TIMEZONE).replace(tzinfo=None)
- def normalize_username(username: str) -> str:
- return username.strip().lower()
- def hash_password(password: str) -> str:
- salt = secrets.token_bytes(16)
- derived = hashlib.scrypt(
- password.encode("utf-8"),
- salt=salt,
- n=_SCRYPT_N,
- r=_SCRYPT_R,
- p=_SCRYPT_P,
- dklen=32,
- )
- return (
- f"scrypt${_SCRYPT_N}${_SCRYPT_R}${_SCRYPT_P}$"
- f"{salt.hex()}${derived.hex()}"
- )
- def verify_password(password: str, encoded: str) -> bool:
- try:
- algorithm, n, r, p, salt_hex, expected_hex = encoded.split("$", 5)
- if algorithm != "scrypt":
- return False
- derived = hashlib.scrypt(
- password.encode("utf-8"),
- salt=bytes.fromhex(salt_hex),
- n=int(n),
- r=int(r),
- p=int(p),
- dklen=len(bytes.fromhex(expected_hex)),
- )
- return hmac.compare_digest(derived.hex(), expected_hex)
- except (TypeError, ValueError):
- return False
- def _dummy_password_hash() -> str:
- global _DUMMY_PASSWORD_HASH
- if _DUMMY_PASSWORD_HASH is None:
- _DUMMY_PASSWORD_HASH = hash_password("supply-agent-invalid-password")
- return _DUMMY_PASSWORD_HASH
- def hash_session_token(token: str) -> str:
- return hashlib.sha256(token.encode("utf-8")).hexdigest()
- def serialize_user(user: AuthUser) -> dict[str, Any]:
- return {
- "id": user.id,
- "username": user.username,
- "display_name": user.display_name,
- "role": user.role,
- "status": user.status,
- "last_login_at": user.last_login_at.isoformat() if user.last_login_at else None,
- "created_at": user.created_at.isoformat() if user.created_at else None,
- "updated_at": user.updated_at.isoformat() if user.updated_at else None,
- }
- def authenticate(
- *,
- username: str,
- password: str,
- ip_address: str | None,
- user_agent: str | None,
- ) -> tuple[dict[str, Any], str]:
- now = _now()
- normalized = normalize_username(username)
- settings = get_infra_settings()
- with get_session() as db:
- repo = AuthRepository(db)
- repo.delete_expired_sessions(now)
- user = repo.get_user_by_username(normalized)
- if user is None:
- verify_password(password, _dummy_password_hash())
- raise InvalidCredentialsError
- if user.locked_until and user.locked_until > now:
- verify_password(password, user.password_hash)
- raise InvalidCredentialsError
- if user.locked_until and user.locked_until <= now:
- user.locked_until = None
- user.failed_login_count = 0
- password_valid = verify_password(password, user.password_hash)
- if not password_valid or user.status != ACTIVE_STATUS:
- if password_valid:
- raise InvalidCredentialsError
- user.failed_login_count += 1
- if user.failed_login_count >= 5:
- user.failed_login_count = 0
- user.locked_until = now + timedelta(minutes=15)
- raise InvalidCredentialsError
- user.failed_login_count = 0
- user.locked_until = None
- user.last_login_at = now
- token = secrets.token_urlsafe(32)
- repo.create_session(
- token_hash=hash_session_token(token),
- user_id=user.id,
- expires_at=now + timedelta(hours=settings.auth_session_hours),
- ip_address=(ip_address or "")[:64] or None,
- user_agent=(user_agent or "")[:512] or None,
- )
- db.flush()
- return serialize_user(user), token
- def resolve_session(token: str) -> dict[str, Any] | None:
- now = _now()
- with get_session() as db:
- repo = AuthRepository(db)
- user = repo.get_active_session(
- token_hash=hash_session_token(token),
- now=now,
- )
- if user is None:
- return None
- return serialize_user(user)
- def revoke_session(token: str) -> None:
- with get_session() as db:
- AuthRepository(db).delete_session_by_hash(hash_session_token(token))
- def list_users() -> list[dict[str, Any]]:
- with get_session() as db:
- return [serialize_user(user) for user in AuthRepository(db).list_users()]
- def create_user(
- *,
- username: str,
- password: str,
- display_name: str,
- role: str,
- ) -> dict[str, Any]:
- if role not in VALID_ROLES:
- raise ValueError("invalid role")
- try:
- with get_session() as db:
- user = AuthRepository(db).create_user(
- username=normalize_username(username),
- password_hash=hash_password(password),
- display_name=display_name.strip(),
- role=role,
- )
- db.flush()
- return serialize_user(user)
- except IntegrityError as exc:
- raise DuplicateUsernameError from exc
- def update_user(
- user_id: int,
- *,
- display_name: str | None = None,
- role: str | None = None,
- status: str | None = None,
- password: str | None = None,
- ) -> dict[str, Any] | None:
- if role is not None and role not in VALID_ROLES:
- raise ValueError("invalid role")
- if status is not None and status not in VALID_STATUSES:
- raise ValueError("invalid status")
- with get_session() as db:
- repo = AuthRepository(db)
- user = repo.get_user(user_id)
- if user is None:
- return None
- revoke_existing_sessions = False
- if display_name is not None:
- user.display_name = display_name.strip()
- if role is not None and role != user.role:
- user.role = role
- revoke_existing_sessions = True
- if status is not None and status != user.status:
- user.status = status
- revoke_existing_sessions = True
- if password is not None:
- user.password_hash = hash_password(password)
- user.failed_login_count = 0
- user.locked_until = None
- revoke_existing_sessions = True
- if revoke_existing_sessions:
- repo.delete_user_sessions(user.id)
- db.flush()
- return serialize_user(user)
- def delete_user(user_id: int) -> bool:
- with get_session() as db:
- repo = AuthRepository(db)
- user = repo.get_user(user_id)
- if user is None:
- return False
- repo.delete_user_sessions(user.id)
- repo.delete_user(user)
- return True
- def ensure_bootstrap_admin() -> None:
- settings = get_infra_settings()
- username = normalize_username(settings.auth_bootstrap_admin_username)
- password = settings.auth_bootstrap_admin_password.get_secret_value()
- if not username and not password:
- return
- if not username or not password:
- raise RuntimeError(
- "AUTH_BOOTSTRAP_ADMIN_USERNAME and AUTH_BOOTSTRAP_ADMIN_PASSWORD "
- "must be configured together"
- )
- if len(password) < 8:
- raise RuntimeError("AUTH_BOOTSTRAP_ADMIN_PASSWORD must contain at least 8 characters")
- with get_session() as db:
- repo = AuthRepository(db)
- if repo.get_user_by_username(username) is not None:
- return
- repo.create_user(
- username=username,
- password_hash=hash_password(password),
- display_name=settings.auth_bootstrap_admin_display_name.strip() or "系统管理员",
- role=ADMIN_ROLE,
- )
- logger.info("Created bootstrap administrator: username=%s", username)
|