| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357 |
- 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
|