from __future__ import annotations from datetime import datetime, timedelta, timezone from uuid import uuid4 import jwt import pytest from sqlalchemy import select from dms.extensions import db from dms.models import AuditLog, Organization, User from dms.security.passwords import hash_password, verify_password def _login(client, password, username="admin", keep_signed_in=False): return client.post( "/api/v1/auth/login", json={ "username": username, "password": password, "keepSignedIn": keep_signed_in, }, ) def _auth(token): return {"Authorization": f"Bearer {token}"} def _custom_token(app, user, **overrides): now = datetime.now(timezone.utc) claims = { "sub": str(user.id), "username": user.username, "roleCode": user.role_code, "authVersion": user.auth_version, "iat": now, "exp": now + timedelta(hours=1), "jti": str(uuid4()), "iss": "dms", } claims.update(overrides) return jwt.encode(claims, app.config["DMS_JWT_SECRET"], algorithm="HS256") def test_argon2id_hash_and_verify(b2_password): password_hash = hash_password(b2_password) assert password_hash.startswith("$argon2id$") assert verify_password(password_hash, b2_password) assert not verify_password(password_hash, b2_password + "x") def test_login_success_and_safe_user_summary(b2_client, b2_password): response = _login(b2_client, b2_password) payload = response.get_json() assert response.status_code == 200 assert payload["data"]["tokenType"] == "Bearer" assert payload["data"]["expiresIn"] == 7200 user = payload["data"]["user"] assert isinstance(user["id"], str) assert user["allowedModules"] == [ "DOCUMENT_BROWSER", "BACKEND_MANAGEMENT", "AUDIT_LOG", ] assert "passwordHash" not in user assert "authVersion" not in user @pytest.mark.parametrize("username,password_suffix", [("missing", ""), ("admin", "-bad")]) def test_invalid_credentials_are_generic( b2_app, b2_client, b2_password, username, password_suffix ): response = _login(b2_client, b2_password + password_suffix, username=username) assert response.status_code == 401 assert response.get_json()["code"] == "INVALID_CREDENTIALS" with b2_app.app_context(): audit = db.session.scalars( select(AuditLog).where(AuditLog.action_type == "LOGIN") ).all() assert len(audit) == 1 assert audit[0].operation_result == "FAILURE" def test_disabled_user_login(b2_client, b2_password): response = _login(b2_client, b2_password, username="disabled") assert response.status_code == 401 assert response.get_json()["code"] == "USER_DISABLED" @pytest.mark.parametrize("keep_signed_in,expires_in", [(False, 7200), (True, 604800)]) def test_keep_signed_in_expiration(b2_client, b2_password, keep_signed_in, expires_in): response = _login(b2_client, b2_password, keep_signed_in=keep_signed_in) assert response.get_json()["data"]["expiresIn"] == expires_in def test_login_rejects_non_exact_payload(b2_client, b2_password): response = b2_client.post( "/api/v1/auth/login", json={"username": "admin", "password": b2_password, "keepSignedIn": False, "x": 1}, ) assert response.status_code == 400 assert response.get_json()["code"] == "INVALID_ARGUMENT" def test_empty_jwt_secret_fails_clearly(b2_app, b2_client, b2_password): b2_app.config["DMS_JWT_SECRET"] = "" response = _login(b2_client, b2_password) assert response.status_code == 500 assert response.get_json()["code"] == "INTERNAL_ERROR" def test_jwt_has_required_claims(b2_app, b2_client, b2_password): token = _login(b2_client, b2_password).get_json()["data"]["accessToken"] claims = jwt.decode( token, b2_app.config["DMS_JWT_SECRET"], algorithms=["HS256"], issuer="dms", ) assert { "sub", "username", "roleCode", "authVersion", "iat", "exp", "jti", "iss" } <= set(claims) assert isinstance(claims["sub"], str) def test_invalid_signature(b2_client): response = b2_client.get("/api/v1/users/me", headers=_auth("not-a-jwt")) assert response.status_code == 401 assert response.get_json()["code"] == "TOKEN_INVALID" def test_expired_token(b2_app, b2_client): with b2_app.app_context(): user = db.session.scalar(select(User).where(User.username == "admin")) token = _custom_token( b2_app, user, exp=datetime.now(timezone.utc) - timedelta(seconds=1), ) response = b2_client.get("/api/v1/users/me", headers=_auth(token)) assert response.status_code == 401 assert response.get_json()["code"] == "TOKEN_EXPIRED" def test_auth_version_mismatch(b2_app, b2_client): with b2_app.app_context(): user = db.session.scalar(select(User).where(User.username == "admin")) token = _custom_token(b2_app, user, authVersion=user.auth_version + 1) response = b2_client.get("/api/v1/users/me", headers=_auth(token)) assert response.status_code == 401 assert response.get_json()["code"] == "AUTH_VERSION_MISMATCH" def test_disabled_user_token(b2_app, b2_client): with b2_app.app_context(): user = db.session.scalar(select(User).where(User.username == "disabled")) token = _custom_token(b2_app, user) response = b2_client.get("/api/v1/users/me", headers=_auth(token)) assert response.status_code == 401 assert response.get_json()["code"] == "USER_DISABLED" def test_logically_deleted_user_token(b2_app, b2_client, login_token): with b2_app.app_context(): user = db.session.scalar(select(User).where(User.username == "admin")) user.is_deleted = True db.session.commit() response = b2_client.get("/api/v1/users/me", headers=_auth(login_token)) assert response.status_code == 401 assert response.get_json()["code"] == "TOKEN_INVALID" def test_missing_and_malformed_bearer(b2_client): assert b2_client.get("/api/v1/users/me").get_json()["code"] == "TOKEN_INVALID" response = b2_client.get( "/api/v1/users/me", headers={"Authorization": "Basic abc"} ) assert response.status_code == 401 assert response.get_json()["code"] == "TOKEN_INVALID" def test_current_user(b2_client, login_token): response = b2_client.get("/api/v1/users/me", headers=_auth(login_token)) assert response.status_code == 200 assert response.get_json()["data"]["username"] == "admin" def test_logout_invalidates_old_token_and_audits(b2_app, b2_client, login_token): with b2_app.app_context(): before = db.session.scalar( select(User.auth_version).where(User.username == "admin") ) response = b2_client.post("/api/v1/auth/logout", headers=_auth(login_token)) assert response.status_code == 200 second = b2_client.get("/api/v1/users/me", headers=_auth(login_token)) assert second.status_code == 401 assert second.get_json()["code"] == "AUTH_VERSION_MISMATCH" with b2_app.app_context(): after = db.session.scalar( select(User.auth_version).where(User.username == "admin") ) assert after == before + 1 logout_audit = db.session.scalar( select(AuditLog).where(AuditLog.action_type == "LOGOUT") ) assert logout_audit is not None assert logout_audit.operation_result == "SUCCESS" def test_audit_has_request_metadata_and_no_secrets( b2_app, b2_client, b2_password ): request_id = str(uuid4()) response = b2_client.post( "/api/v1/auth/login", json={ "username": "admin", "password": b2_password, "keepSignedIn": False, }, headers={ "X-Request-Id": request_id, "User-Agent": "b2-test-agent", "X-Forwarded-For": "192.0.2.10", }, ) token = response.get_json()["data"]["accessToken"] with b2_app.app_context(): audit = db.session.scalar( select(AuditLog).where(AuditLog.action_type == "LOGIN") ) rendered = " ".join( str(value) for value in ( audit.failure_reason, audit.operation_detail, audit.target_name, ) ) assert audit.request_id == request_id assert audit.client_ip == "192.0.2.10" assert audit.user_agent == "b2-test-agent" assert b2_password not in rendered assert token not in rendered def test_organization_tree_and_stable_children(b2_client, login_token): response = b2_client.get( "/api/v1/organizations/tree", headers=_auth(login_token) ) data = response.get_json()["data"] assert [item["orgCode"] for item in data] == ["ORG_ROOT"] assert [item["orgCode"] for item in data[0]["children"]] == [ "ORG_OPS", "ORG_COMMS", ] assert all(isinstance(item["id"], str) for item in data[0]["children"]) def test_organization_keyword_preserves_ancestor(b2_client, login_token): response = b2_client.get( "/api/v1/organizations/tree?keyword=通信", headers=_auth(login_token), ) root = response.get_json()["data"][0] assert root["orgCode"] == "ORG_ROOT" assert [item["orgCode"] for item in root["children"]] == ["ORG_COMMS"] def test_organization_status_filter(b2_client, login_token): response = b2_client.get( "/api/v1/organizations/tree?status=DISABLED", headers=_auth(login_token), ) root = response.get_json()["data"][0] assert [item["orgCode"] for item in root["children"]] == ["ORG_DISABLED"] def test_user_pagination_and_keyword(b2_client, login_token): response = b2_client.get( "/api/v1/users?page=1&pageSize=2&keyword=员", headers=_auth(login_token), ) page = response.get_json()["data"] assert response.status_code == 200 assert page["pageSize"] == 2 assert page["total"] == 1 assert len(page["items"]) == 1 def test_include_descendants(b2_app, b2_client, login_token): with b2_app.app_context(): root_id = db.session.scalar( select(Organization.id).where(Organization.org_code == "ORG_ROOT") ) direct = b2_client.get( f"/api/v1/users?organizationId={root_id}", headers=_auth(login_token), ).get_json()["data"] recursive = b2_client.get( f"/api/v1/users?organizationId={root_id}&includeDescendants=true", headers=_auth(login_token), ).get_json()["data"] assert direct["total"] == 1 assert recursive["total"] == 3 def test_user_status_filter(b2_client, login_token): response = b2_client.get( "/api/v1/users?status=DISABLED", headers=_auth(login_token) ) items = response.get_json()["data"]["items"] assert [item["username"] for item in items] == ["disabled"] def test_page_size_limit(b2_client, login_token): response = b2_client.get( "/api/v1/users?pageSize=101", headers=_auth(login_token) ) assert response.status_code == 400 assert response.get_json()["code"] == "INVALID_ARGUMENT" @pytest.mark.parametrize("username", ["user"]) def test_admin_role_required_for_directory( b2_client, b2_password, username ): token = _login( b2_client, b2_password, username=username ).get_json()["data"]["accessToken"] assert ( b2_client.get("/api/v1/users", headers=_auth(token)).status_code == 403 ) assert ( b2_client.get( "/api/v1/organizations/tree", headers=_auth(token) ).status_code == 403 ) @pytest.mark.parametrize( "username,modules", [ ( "admin", ["DOCUMENT_BROWSER", "BACKEND_MANAGEMENT", "AUDIT_LOG"], ), ("user", ["DOCUMENT_BROWSER"]), ], ) def test_allowed_modules_mapping( b2_client, b2_password, username, modules ): response = _login(b2_client, b2_password, username=username) assert response.get_json()["data"]["user"]["allowedModules"] == modules