test_b2_auth_and_directory.py 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351
  1. from __future__ import annotations
  2. from datetime import datetime, timedelta, timezone
  3. from uuid import uuid4
  4. import jwt
  5. import pytest
  6. from sqlalchemy import select
  7. from dms.extensions import db
  8. from dms.models import AuditLog, Organization, User
  9. from dms.security.passwords import hash_password, verify_password
  10. def _login(client, password, username="admin", keep_signed_in=False):
  11. return client.post(
  12. "/api/v1/auth/login",
  13. json={
  14. "username": username,
  15. "password": password,
  16. "keepSignedIn": keep_signed_in,
  17. },
  18. )
  19. def _auth(token):
  20. return {"Authorization": f"Bearer {token}"}
  21. def _custom_token(app, user, **overrides):
  22. now = datetime.now(timezone.utc)
  23. claims = {
  24. "sub": str(user.id),
  25. "username": user.username,
  26. "roleCode": user.role_code,
  27. "authVersion": user.auth_version,
  28. "iat": now,
  29. "exp": now + timedelta(hours=1),
  30. "jti": str(uuid4()),
  31. "iss": "dms",
  32. }
  33. claims.update(overrides)
  34. return jwt.encode(claims, app.config["DMS_JWT_SECRET"], algorithm="HS256")
  35. def test_argon2id_hash_and_verify(b2_password):
  36. password_hash = hash_password(b2_password)
  37. assert password_hash.startswith("$argon2id$")
  38. assert verify_password(password_hash, b2_password)
  39. assert not verify_password(password_hash, b2_password + "x")
  40. def test_login_success_and_safe_user_summary(b2_client, b2_password):
  41. response = _login(b2_client, b2_password)
  42. payload = response.get_json()
  43. assert response.status_code == 200
  44. assert payload["data"]["tokenType"] == "Bearer"
  45. assert payload["data"]["expiresIn"] == 7200
  46. user = payload["data"]["user"]
  47. assert isinstance(user["id"], str)
  48. assert user["allowedModules"] == ["DOCUMENT_BROWSER", "BACKEND_MANAGEMENT"]
  49. assert "passwordHash" not in user
  50. assert "authVersion" not in user
  51. @pytest.mark.parametrize("username,password_suffix", [("missing", ""), ("admin", "-bad")])
  52. def test_invalid_credentials_are_generic(
  53. b2_app, b2_client, b2_password, username, password_suffix
  54. ):
  55. response = _login(b2_client, b2_password + password_suffix, username=username)
  56. assert response.status_code == 401
  57. assert response.get_json()["code"] == "INVALID_CREDENTIALS"
  58. with b2_app.app_context():
  59. audit = db.session.scalars(
  60. select(AuditLog).where(AuditLog.action_type == "LOGIN")
  61. ).all()
  62. assert len(audit) == 1
  63. assert audit[0].operation_result == "FAILURE"
  64. def test_disabled_user_login(b2_client, b2_password):
  65. response = _login(b2_client, b2_password, username="disabled")
  66. assert response.status_code == 401
  67. assert response.get_json()["code"] == "USER_DISABLED"
  68. @pytest.mark.parametrize("keep_signed_in,expires_in", [(False, 7200), (True, 604800)])
  69. def test_keep_signed_in_expiration(b2_client, b2_password, keep_signed_in, expires_in):
  70. response = _login(b2_client, b2_password, keep_signed_in=keep_signed_in)
  71. assert response.get_json()["data"]["expiresIn"] == expires_in
  72. def test_login_rejects_non_exact_payload(b2_client, b2_password):
  73. response = b2_client.post(
  74. "/api/v1/auth/login",
  75. json={"username": "admin", "password": b2_password, "keepSignedIn": False, "x": 1},
  76. )
  77. assert response.status_code == 400
  78. assert response.get_json()["code"] == "INVALID_ARGUMENT"
  79. def test_empty_jwt_secret_fails_clearly(b2_app, b2_client, b2_password):
  80. b2_app.config["DMS_JWT_SECRET"] = ""
  81. response = _login(b2_client, b2_password)
  82. assert response.status_code == 500
  83. assert response.get_json()["code"] == "INTERNAL_ERROR"
  84. def test_jwt_has_required_claims(b2_app, b2_client, b2_password):
  85. token = _login(b2_client, b2_password).get_json()["data"]["accessToken"]
  86. claims = jwt.decode(
  87. token,
  88. b2_app.config["DMS_JWT_SECRET"],
  89. algorithms=["HS256"],
  90. issuer="dms",
  91. )
  92. assert {
  93. "sub", "username", "roleCode", "authVersion", "iat", "exp", "jti", "iss"
  94. } <= set(claims)
  95. assert isinstance(claims["sub"], str)
  96. def test_invalid_signature(b2_client):
  97. response = b2_client.get("/api/v1/users/me", headers=_auth("not-a-jwt"))
  98. assert response.status_code == 401
  99. assert response.get_json()["code"] == "TOKEN_INVALID"
  100. def test_expired_token(b2_app, b2_client):
  101. with b2_app.app_context():
  102. user = db.session.scalar(select(User).where(User.username == "admin"))
  103. token = _custom_token(
  104. b2_app,
  105. user,
  106. exp=datetime.now(timezone.utc) - timedelta(seconds=1),
  107. )
  108. response = b2_client.get("/api/v1/users/me", headers=_auth(token))
  109. assert response.status_code == 401
  110. assert response.get_json()["code"] == "TOKEN_EXPIRED"
  111. def test_auth_version_mismatch(b2_app, b2_client):
  112. with b2_app.app_context():
  113. user = db.session.scalar(select(User).where(User.username == "admin"))
  114. token = _custom_token(b2_app, user, authVersion=user.auth_version + 1)
  115. response = b2_client.get("/api/v1/users/me", headers=_auth(token))
  116. assert response.status_code == 401
  117. assert response.get_json()["code"] == "AUTH_VERSION_MISMATCH"
  118. def test_disabled_user_token(b2_app, b2_client):
  119. with b2_app.app_context():
  120. user = db.session.scalar(select(User).where(User.username == "disabled"))
  121. token = _custom_token(b2_app, user)
  122. response = b2_client.get("/api/v1/users/me", headers=_auth(token))
  123. assert response.status_code == 401
  124. assert response.get_json()["code"] == "USER_DISABLED"
  125. def test_logically_deleted_user_token(b2_app, b2_client, login_token):
  126. with b2_app.app_context():
  127. user = db.session.scalar(select(User).where(User.username == "admin"))
  128. user.is_deleted = True
  129. db.session.commit()
  130. response = b2_client.get("/api/v1/users/me", headers=_auth(login_token))
  131. assert response.status_code == 401
  132. assert response.get_json()["code"] == "TOKEN_INVALID"
  133. def test_missing_and_malformed_bearer(b2_client):
  134. assert b2_client.get("/api/v1/users/me").get_json()["code"] == "TOKEN_INVALID"
  135. response = b2_client.get(
  136. "/api/v1/users/me", headers={"Authorization": "Basic abc"}
  137. )
  138. assert response.status_code == 401
  139. assert response.get_json()["code"] == "TOKEN_INVALID"
  140. def test_current_user(b2_client, login_token):
  141. response = b2_client.get("/api/v1/users/me", headers=_auth(login_token))
  142. assert response.status_code == 200
  143. assert response.get_json()["data"]["username"] == "admin"
  144. def test_logout_invalidates_old_token_and_audits(b2_app, b2_client, login_token):
  145. with b2_app.app_context():
  146. before = db.session.scalar(
  147. select(User.auth_version).where(User.username == "admin")
  148. )
  149. response = b2_client.post("/api/v1/auth/logout", headers=_auth(login_token))
  150. assert response.status_code == 200
  151. second = b2_client.get("/api/v1/users/me", headers=_auth(login_token))
  152. assert second.status_code == 401
  153. assert second.get_json()["code"] == "AUTH_VERSION_MISMATCH"
  154. with b2_app.app_context():
  155. after = db.session.scalar(
  156. select(User.auth_version).where(User.username == "admin")
  157. )
  158. assert after == before + 1
  159. logout_audit = db.session.scalar(
  160. select(AuditLog).where(AuditLog.action_type == "LOGOUT")
  161. )
  162. assert logout_audit is not None
  163. assert logout_audit.operation_result == "SUCCESS"
  164. def test_audit_has_request_metadata_and_no_secrets(
  165. b2_app, b2_client, b2_password
  166. ):
  167. request_id = str(uuid4())
  168. response = b2_client.post(
  169. "/api/v1/auth/login",
  170. json={
  171. "username": "admin",
  172. "password": b2_password,
  173. "keepSignedIn": False,
  174. },
  175. headers={
  176. "X-Request-Id": request_id,
  177. "User-Agent": "b2-test-agent",
  178. "X-Forwarded-For": "192.0.2.10",
  179. },
  180. )
  181. token = response.get_json()["data"]["accessToken"]
  182. with b2_app.app_context():
  183. audit = db.session.scalar(
  184. select(AuditLog).where(AuditLog.action_type == "LOGIN")
  185. )
  186. rendered = " ".join(
  187. str(value)
  188. for value in (
  189. audit.failure_reason,
  190. audit.operation_detail,
  191. audit.target_name,
  192. )
  193. )
  194. assert audit.request_id == request_id
  195. assert audit.client_ip == "192.0.2.10"
  196. assert audit.user_agent == "b2-test-agent"
  197. assert b2_password not in rendered
  198. assert token not in rendered
  199. def test_organization_tree_and_stable_children(b2_client, login_token):
  200. response = b2_client.get(
  201. "/api/v1/organizations/tree", headers=_auth(login_token)
  202. )
  203. data = response.get_json()["data"]
  204. assert [item["orgCode"] for item in data] == ["ORG_ROOT"]
  205. assert [item["orgCode"] for item in data[0]["children"]] == [
  206. "ORG_OPS",
  207. "ORG_COMMS",
  208. ]
  209. assert all(isinstance(item["id"], str) for item in data[0]["children"])
  210. def test_organization_keyword_preserves_ancestor(b2_client, login_token):
  211. response = b2_client.get(
  212. "/api/v1/organizations/tree?keyword=通信",
  213. headers=_auth(login_token),
  214. )
  215. root = response.get_json()["data"][0]
  216. assert root["orgCode"] == "ORG_ROOT"
  217. assert [item["orgCode"] for item in root["children"]] == ["ORG_COMMS"]
  218. def test_organization_status_filter(b2_client, login_token):
  219. response = b2_client.get(
  220. "/api/v1/organizations/tree?status=DISABLED",
  221. headers=_auth(login_token),
  222. )
  223. root = response.get_json()["data"][0]
  224. assert [item["orgCode"] for item in root["children"]] == ["ORG_DISABLED"]
  225. def test_user_pagination_and_keyword(b2_client, login_token):
  226. response = b2_client.get(
  227. "/api/v1/users?page=1&pageSize=2&keyword=员",
  228. headers=_auth(login_token),
  229. )
  230. page = response.get_json()["data"]
  231. assert response.status_code == 200
  232. assert page["pageSize"] == 2
  233. assert page["total"] == 2
  234. assert len(page["items"]) == 2
  235. def test_include_descendants(b2_app, b2_client, login_token):
  236. with b2_app.app_context():
  237. root_id = db.session.scalar(
  238. select(Organization.id).where(Organization.org_code == "ORG_ROOT")
  239. )
  240. direct = b2_client.get(
  241. f"/api/v1/users?organizationId={root_id}",
  242. headers=_auth(login_token),
  243. ).get_json()["data"]
  244. recursive = b2_client.get(
  245. f"/api/v1/users?organizationId={root_id}&includeDescendants=true",
  246. headers=_auth(login_token),
  247. ).get_json()["data"]
  248. assert direct["total"] == 2
  249. assert recursive["total"] == 4
  250. def test_user_status_filter(b2_client, login_token):
  251. response = b2_client.get(
  252. "/api/v1/users?status=DISABLED", headers=_auth(login_token)
  253. )
  254. items = response.get_json()["data"]["items"]
  255. assert [item["username"] for item in items] == ["disabled"]
  256. def test_page_size_limit(b2_client, login_token):
  257. response = b2_client.get(
  258. "/api/v1/users?pageSize=101", headers=_auth(login_token)
  259. )
  260. assert response.status_code == 400
  261. assert response.get_json()["code"] == "INVALID_ARGUMENT"
  262. @pytest.mark.parametrize("username", ["user", "auditor"])
  263. def test_admin_role_required_for_directory(
  264. b2_client, b2_password, username
  265. ):
  266. token = _login(
  267. b2_client, b2_password, username=username
  268. ).get_json()["data"]["accessToken"]
  269. assert (
  270. b2_client.get("/api/v1/users", headers=_auth(token)).status_code == 403
  271. )
  272. assert (
  273. b2_client.get(
  274. "/api/v1/organizations/tree", headers=_auth(token)
  275. ).status_code
  276. == 403
  277. )
  278. @pytest.mark.parametrize(
  279. "username,modules",
  280. [
  281. ("admin", ["DOCUMENT_BROWSER", "BACKEND_MANAGEMENT"]),
  282. ("auditor", ["AUDIT_LOG"]),
  283. ("user", ["DOCUMENT_BROWSER"]),
  284. ],
  285. )
  286. def test_allowed_modules_mapping(
  287. b2_client, b2_password, username, modules
  288. ):
  289. response = _login(b2_client, b2_password, username=username)
  290. assert response.get_json()["data"]["user"]["allowedModules"] == modules