329 lines
11 KiB
Python
329 lines
11 KiB
Python
"""Authentication, sessions, rate limiting and audit logging.
|
|
|
|
Passwords are hashed with PBKDF2-HMAC-SHA256 (stdlib ``hashlib.pbkdf2_hmac``)
|
|
and per-user random salts; plaintext passwords are never stored or logged.
|
|
Session tokens are random URL-safe strings; only their SHA-256 digest is
|
|
persisted, so a database leak does not expose usable tokens. Every login
|
|
attempt and every privileged action lands in ``audit_log``. The reasoning
|
|
behind these choices is recorded in ``docs/decisions/003-auth.md``.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import datetime, timedelta, timezone
|
|
import hashlib
|
|
import hmac
|
|
import os
|
|
import secrets
|
|
import sqlite3
|
|
import string
|
|
|
|
from .db import utc_now
|
|
|
|
|
|
MIN_PASSWORD_LENGTH = 8
|
|
PBKDF2_ITERATIONS = 260_000
|
|
SESSION_TTL_HOURS = 8
|
|
RATE_LIMIT_MAX_FAILURES = 5
|
|
RATE_LIMIT_WINDOW_MINUTES = 10
|
|
INITIAL_PASSWORD_LENGTH = 12
|
|
|
|
|
|
def _env_flag(value: str | None) -> bool:
|
|
"""Parse a yes/no style environment flag; blank/absent means False."""
|
|
return (value or "").strip().lower() in {"1", "true", "yes", "on"}
|
|
|
|
|
|
# Ops switch for internal test environments: APP_LOGIN_RATE_LIMIT_DISABLED=1
|
|
# turns off the login-failure lockout entirely. The default (unset) keeps the
|
|
# production policy — 5 failures within 10 minutes lock the (账号, IP) pair.
|
|
RATE_LIMIT_DISABLED = _env_flag(os.environ.get("APP_LOGIN_RATE_LIMIT_DISABLED"))
|
|
|
|
|
|
def hash_password(password: str) -> str:
|
|
"""Hash ``password`` as ``pbkdf2_sha256$<iterations>$<salt_hex>$<hash_hex>``."""
|
|
salt = secrets.token_bytes(16)
|
|
digest = hashlib.pbkdf2_hmac(
|
|
"sha256", password.encode("utf-8"), salt, PBKDF2_ITERATIONS
|
|
)
|
|
return f"pbkdf2_sha256${PBKDF2_ITERATIONS}${salt.hex()}${digest.hex()}"
|
|
|
|
|
|
def verify_password(password: str, stored: str) -> bool:
|
|
"""Constant-time check of ``password`` against a stored hash string."""
|
|
try:
|
|
scheme, iterations, salt_hex, hash_hex = stored.split("$")
|
|
if scheme != "pbkdf2_sha256":
|
|
return False
|
|
salt = bytes.fromhex(salt_hex)
|
|
expected = bytes.fromhex(hash_hex)
|
|
digest = hashlib.pbkdf2_hmac(
|
|
"sha256", password.encode("utf-8"), salt, int(iterations)
|
|
)
|
|
except (ValueError, TypeError):
|
|
return False
|
|
return hmac.compare_digest(digest, expected)
|
|
|
|
|
|
def generate_initial_password(exclude: str | None = None) -> str:
|
|
"""Generate a 12-char initial password with upper, lower and digit chars.
|
|
|
|
When ``exclude`` is given, the result is guaranteed to differ from it
|
|
(case-insensitive) so a fresh account never starts with a password equal
|
|
to its own username.
|
|
"""
|
|
alphabet = string.ascii_letters + string.digits
|
|
while True:
|
|
password = "".join(
|
|
secrets.choice(alphabet) for _ in range(INITIAL_PASSWORD_LENGTH)
|
|
)
|
|
if (
|
|
any(char.isupper() for char in password)
|
|
and any(char.islower() for char in password)
|
|
and any(char.isdigit() for char in password)
|
|
and (exclude is None or password.lower() != exclude.lower())
|
|
):
|
|
return password
|
|
|
|
|
|
def validate_password_policy(password: str, username: str) -> str | None:
|
|
"""Return an error message when ``password`` violates policy, else None."""
|
|
if len(password) < MIN_PASSWORD_LENGTH:
|
|
return f"密码长度至少为 {MIN_PASSWORD_LENGTH} 位。"
|
|
if password.lower() == username.lower():
|
|
return "密码不能与账号相同。"
|
|
if not any(char.isalpha() for char in password) or not any(
|
|
char.isdigit() for char in password
|
|
):
|
|
return "密码必须同时包含字母和数字。"
|
|
return None
|
|
|
|
|
|
def create_user(
|
|
connection: sqlite3.Connection,
|
|
username: str,
|
|
password: str,
|
|
role: str,
|
|
company_id: int | None = None,
|
|
must_change_password: bool = True,
|
|
) -> int:
|
|
"""Create a user, enforcing the role/company binding rules. Returns the id."""
|
|
username = username.strip()
|
|
if not username:
|
|
raise ValueError("用户名不能为空。")
|
|
if role not in ("admin", "company"):
|
|
raise ValueError("角色必须是 admin 或 company。")
|
|
if role == "company":
|
|
if company_id is None:
|
|
raise ValueError("公司账号必须绑定公司。")
|
|
company = connection.execute(
|
|
"SELECT id FROM companies WHERE id = ?", (company_id,)
|
|
).fetchone()
|
|
if company is None:
|
|
raise ValueError("绑定的公司不存在。")
|
|
elif company_id is not None:
|
|
raise ValueError("管理员账号不能绑定公司。")
|
|
|
|
now = utc_now()
|
|
try:
|
|
with connection:
|
|
cursor = connection.execute(
|
|
"""
|
|
INSERT INTO users (
|
|
username, password_hash, role, company_id,
|
|
must_change_password, created_at, updated_at
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
username,
|
|
hash_password(password),
|
|
role,
|
|
company_id,
|
|
1 if must_change_password else 0,
|
|
now,
|
|
now,
|
|
),
|
|
)
|
|
except sqlite3.IntegrityError as exc:
|
|
raise ValueError("用户名已存在。") from exc
|
|
return int(cursor.lastrowid)
|
|
|
|
|
|
def authenticate(
|
|
connection: sqlite3.Connection, username: str, password: str, ip: str
|
|
) -> tuple[sqlite3.Row | None, str | None]:
|
|
"""Verify credentials; returns ``(user_row, None)`` or ``(None, reason)``.
|
|
|
|
``reason`` is one of ``rate_limited``, ``disabled``, ``bad_credentials``.
|
|
Every non-rate-limited attempt is recorded in ``login_attempts`` and
|
|
``audit_log``; the password itself is never stored anywhere.
|
|
"""
|
|
if not RATE_LIMIT_DISABLED:
|
|
window_start = (
|
|
datetime.now(timezone.utc) - timedelta(minutes=RATE_LIMIT_WINDOW_MINUTES)
|
|
).isoformat()
|
|
failures = connection.execute(
|
|
"""
|
|
SELECT COUNT(*) AS n FROM login_attempts
|
|
WHERE username = ? AND ip = ? AND success = 0 AND created_at >= ?
|
|
""",
|
|
(username, ip, window_start),
|
|
).fetchone()
|
|
if failures["n"] >= RATE_LIMIT_MAX_FAILURES:
|
|
return None, "rate_limited"
|
|
|
|
user = connection.execute(
|
|
"SELECT * FROM users WHERE username = ?", (username,)
|
|
).fetchone()
|
|
|
|
if user is not None and user["status"] == "disabled":
|
|
_record_attempt(connection, username, ip, success=False)
|
|
audit(
|
|
connection,
|
|
"login_failed",
|
|
actor=user,
|
|
detail="账号已停用",
|
|
ip=ip,
|
|
)
|
|
return None, "disabled"
|
|
|
|
if user is None or not verify_password(password, user["password_hash"]):
|
|
_record_attempt(connection, username, ip, success=False)
|
|
audit(connection, "login_failed", actor=user, detail="账号或密码不正确", ip=ip)
|
|
return None, "bad_credentials"
|
|
|
|
_record_attempt(connection, username, ip, success=True)
|
|
audit(connection, "login_success", actor=user, ip=ip)
|
|
return user, None
|
|
|
|
|
|
def create_session(
|
|
connection: sqlite3.Connection, user_id: int, ttl_hours: int = SESSION_TTL_HOURS
|
|
) -> str:
|
|
"""Create a session with absolute expiry; returns the raw token."""
|
|
token = secrets.token_urlsafe(32)
|
|
token_hash = hashlib.sha256(token.encode("utf-8")).hexdigest()
|
|
now = datetime.now(timezone.utc)
|
|
with connection:
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO sessions (token_hash, user_id, created_at, expires_at)
|
|
VALUES (?, ?, ?, ?)
|
|
""",
|
|
(
|
|
token_hash,
|
|
user_id,
|
|
now.isoformat(),
|
|
(now + timedelta(hours=ttl_hours)).isoformat(),
|
|
),
|
|
)
|
|
return token
|
|
|
|
|
|
def resolve_session(connection: sqlite3.Connection, token: str) -> sqlite3.Row | None:
|
|
"""Return the user row for a live session token, else None.
|
|
|
|
Expired or revoked sessions and disabled users are all rejected.
|
|
"""
|
|
token_hash = hashlib.sha256(token.encode("utf-8")).hexdigest()
|
|
return connection.execute(
|
|
"""
|
|
SELECT u.*, s.id AS session_id
|
|
FROM sessions s
|
|
JOIN users u ON u.id = s.user_id
|
|
WHERE s.token_hash = ?
|
|
AND s.revoked_at IS NULL
|
|
AND s.expires_at > ?
|
|
AND u.status = 'active'
|
|
""",
|
|
(token_hash, utc_now()),
|
|
).fetchone()
|
|
|
|
|
|
def revoke_session(connection: sqlite3.Connection, token: str) -> None:
|
|
token_hash = hashlib.sha256(token.encode("utf-8")).hexdigest()
|
|
with connection:
|
|
connection.execute(
|
|
"UPDATE sessions SET revoked_at = ? WHERE token_hash = ? AND revoked_at IS NULL",
|
|
(utc_now(), token_hash),
|
|
)
|
|
|
|
|
|
def revoke_user_sessions(connection: sqlite3.Connection, user_id: int) -> None:
|
|
with connection:
|
|
connection.execute(
|
|
"UPDATE sessions SET revoked_at = ? WHERE user_id = ? AND revoked_at IS NULL",
|
|
(utc_now(), user_id),
|
|
)
|
|
|
|
|
|
def change_password(
|
|
connection: sqlite3.Connection,
|
|
user_id: int,
|
|
old_password: str,
|
|
new_password: str,
|
|
) -> str | None:
|
|
"""Change a user's password; returns an error message or None on success."""
|
|
user = connection.execute(
|
|
"SELECT * FROM users WHERE id = ?", (user_id,)
|
|
).fetchone()
|
|
if user is None:
|
|
return "用户不存在。"
|
|
if not verify_password(old_password, user["password_hash"]):
|
|
return "原密码不正确。"
|
|
error = validate_password_policy(new_password, user["username"])
|
|
if error is not None:
|
|
return error
|
|
with connection:
|
|
connection.execute(
|
|
"""
|
|
UPDATE users
|
|
SET password_hash = ?, must_change_password = 0, updated_at = ?
|
|
WHERE id = ?
|
|
""",
|
|
(hash_password(new_password), utc_now(), user_id),
|
|
)
|
|
audit(connection, "password_change", actor=user, target=f"user:{user_id}")
|
|
return None
|
|
|
|
|
|
def audit(
|
|
connection: sqlite3.Connection,
|
|
action: str,
|
|
actor: sqlite3.Row | None = None,
|
|
target: str | None = None,
|
|
detail: str | None = None,
|
|
ip: str | None = None,
|
|
) -> None:
|
|
"""Append an audit log entry. Never pass passwords in ``detail``."""
|
|
with connection:
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO audit_log (
|
|
actor_user_id, actor_username, action, target, detail, ip, created_at
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
actor["id"] if actor is not None else None,
|
|
actor["username"] if actor is not None else None,
|
|
action,
|
|
target,
|
|
detail,
|
|
ip,
|
|
utc_now(),
|
|
),
|
|
)
|
|
|
|
|
|
def _record_attempt(
|
|
connection: sqlite3.Connection, username: str, ip: str, success: bool
|
|
) -> None:
|
|
with connection:
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO login_attempts (username, ip, success, created_at)
|
|
VALUES (?, ?, ?, ?)
|
|
""",
|
|
(username, ip, 1 if success else 0, utc_now()),
|
|
)
|