test_b2_auth_and_directory.py 12 KB

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