test_phase3_wp02_enterprise_identity_postgres.py 21 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353
  1. from __future__ import annotations
  2. import os
  3. import threading
  4. import uuid
  5. from datetime import UTC, datetime, timedelta
  6. import pytest
  7. from sqlalchemy import create_engine, text
  8. from sqlalchemy.orm import Session
  9. from app.core.system.enterprise_identity import (
  10. DirectorySynchronizer,
  11. EmergencyAccess,
  12. IdentityPolicyError,
  13. sanitize_audit_detail,
  14. )
  15. from app.core.system.identity_repository import PostgresIdentityRepository
  16. from app.core.system.identity_sessions import SessionManager
  17. pytestmark = pytest.mark.integration
  18. def _uid() -> str:
  19. return str(uuid.uuid4())
  20. def _identity_attributes(username: str, roles: list[str] | None = None) -> dict:
  21. return {
  22. "provider_uid": "filled-by-synchronizer",
  23. "username": username,
  24. "display_name": username.title(),
  25. "department": "engineering",
  26. "groups": ["data-stewards"],
  27. "roles": roles or ["viewer"],
  28. "mapping_version": "integration-v1",
  29. "claims_digest": "a" * 64,
  30. "authorization_scope": {"environments": ["test"]},
  31. }
  32. @pytest.fixture
  33. def identity_database():
  34. database_url = os.environ.get("TEST_DATABASE_URL")
  35. if not database_url:
  36. pytest.skip("TEST_DATABASE_URL is required")
  37. engine = create_engine(database_url, pool_pre_ping=True)
  38. marker = uuid.uuid4().hex[:10]
  39. users = {name: _uid() for name in ("requester", "breakglass", "approver_a", "approver_b")}
  40. providers: list[str] = []
  41. with engine.begin() as connection:
  42. for name, user_uid in users.items():
  43. connection.execute(text("""
  44. INSERT INTO public.users(id,username,display_name,password_hash,status)
  45. VALUES(CAST(:uid AS uuid),:username,:display,'integration-only','active')
  46. """), {"uid": user_uid, "username": f"p3wp02-{name}-{marker}", "display": name})
  47. connection.execute(text("""
  48. INSERT INTO public.user_roles(user_id,role_id)
  49. SELECT CAST(:uid AS uuid),id FROM public.roles WHERE name='admin'
  50. """), {"uid": user_uid})
  51. state = {"engine": engine, "users": users, "providers": providers, "marker": marker}
  52. try:
  53. yield state
  54. finally:
  55. with engine.begin() as connection:
  56. linked_users: set[str] = set()
  57. for provider in providers:
  58. linked_users.update(str(value) for value in connection.execute(text(
  59. "SELECT user_uid FROM public.enterprise_identity_links WHERE provider_uid=CAST(:provider AS uuid)"
  60. ), {"provider": provider}).scalars())
  61. connection.execute(text("DELETE FROM public.identity_audit_events WHERE provider_uid=CAST(:provider AS uuid)"),
  62. {"provider": provider})
  63. connection.execute(text("DELETE FROM public.identity_directory_conflicts WHERE provider_uid=CAST(:provider AS uuid)"),
  64. {"provider": provider})
  65. connection.execute(text("DELETE FROM public.identity_directory_events WHERE provider_uid=CAST(:provider AS uuid)"),
  66. {"provider": provider})
  67. connection.execute(text("DELETE FROM public.identity_directory_checkpoints WHERE provider_uid=CAST(:provider AS uuid)"),
  68. {"provider": provider})
  69. connection.execute(text("DELETE FROM public.identity_organization_nodes WHERE provider_uid=CAST(:provider AS uuid)"),
  70. {"provider": provider})
  71. all_users = set(users.values()) | linked_users
  72. for user_uid in all_users:
  73. connection.execute(text("DELETE FROM public.identity_exchange_codes WHERE session_uid IN (SELECT uid FROM public.identity_sessions WHERE user_uid=CAST(:uid AS uuid))"),
  74. {"uid": user_uid})
  75. connection.execute(text("DELETE FROM public.identity_sessions WHERE user_uid=CAST(:uid AS uuid)"), {"uid": user_uid})
  76. for request_uid in connection.execute(text("SELECT uid::text FROM public.identity_emergency_requests WHERE requester_uid=CAST(:uid AS uuid)"),
  77. {"uid": users["requester"]}).scalars().all():
  78. connection.execute(text("DELETE FROM public.identity_emergency_approvals WHERE request_uid=CAST(:uid AS uuid)"), {"uid": request_uid})
  79. connection.execute(text("DELETE FROM public.identity_emergency_requests WHERE uid=CAST(:uid AS uuid)"), {"uid": request_uid})
  80. for provider in providers:
  81. connection.execute(text("DELETE FROM public.enterprise_identity_links WHERE provider_uid=CAST(:provider AS uuid)"),
  82. {"provider": provider})
  83. for user_uid in all_users:
  84. connection.execute(text("DELETE FROM public.identity_audit_events WHERE actor_uid=CAST(:uid AS uuid)"), {"uid": user_uid})
  85. connection.execute(text("DELETE FROM public.user_roles WHERE user_id=CAST(:uid AS uuid)"), {"uid": user_uid})
  86. connection.execute(text("DELETE FROM public.users WHERE id=CAST(:uid AS uuid)"), {"uid": user_uid})
  87. assert connection.execute(
  88. text("SELECT count(*) FROM public.users WHERE username LIKE :pattern"),
  89. {"pattern": f"%-{marker}"},
  90. ).scalar_one() == 0
  91. engine.dispose()
  92. def _provider(state) -> str:
  93. provider = _uid()
  94. state["providers"].append(provider)
  95. return provider
  96. def test_real_470_schema_constraints_and_provider_scoped_directory_unit_of_work(identity_database, monkeypatch):
  97. engine = identity_database["engine"]
  98. provider_a, provider_b = _provider(identity_database), _provider(identity_database)
  99. with Session(engine) as session:
  100. assert session.execute(text("SELECT version_num FROM public.alembic_version")).scalar_one() == "20260802_471"
  101. tables = set(session.execute(text("""
  102. SELECT table_name FROM information_schema.tables WHERE table_schema='public'
  103. AND table_name LIKE 'identity_%'
  104. """)).scalars())
  105. assert {"identity_directory_events", "identity_directory_checkpoints", "identity_directory_conflicts",
  106. "identity_organization_nodes", "identity_sessions", "identity_emergency_requests"} <= tables
  107. constraints = " ".join(session.execute(text("""
  108. SELECT pg_get_constraintdef(oid) FROM pg_constraint
  109. WHERE conrelid IN ('public.identity_sessions'::regclass,
  110. 'public.identity_directory_events'::regclass,
  111. 'public.identity_emergency_requests'::regclass,
  112. 'public.identity_idp_config_versions'::regclass)
  113. """)).scalars())
  114. assert "UNIQUE (refresh_hash)" in constraints
  115. assert "UNIQUE (provider_uid, source, source_event_id)" in constraints
  116. assert "expires_at > starts_at" in constraints
  117. assert "DATAOPS_OIDC_" in constraints
  118. repo = PostgresIdentityRepository(session)
  119. sync = DirectorySynchronizer(repo, SessionManager(repo, secret="integration-secret"))
  120. first = sync.apply(provider_uid=provider_a, source="scim", source_event_id="join-a", cursor="1",
  121. cursor_sequence=1, event_type="JOINER", subject="shared-subject",
  122. attributes=_identity_attributes(f"alice-a-{identity_database['marker']}"))
  123. second = sync.apply(provider_uid=provider_b, source="scim", source_event_id="join-b", cursor="1",
  124. cursor_sequence=1, event_type="JOINER", subject="shared-subject",
  125. attributes=_identity_attributes(f"alice-b-{identity_database['marker']}"))
  126. assert first["user_uid"] != second["user_uid"]
  127. assert repo.get_identity(provider_a, "shared-subject")["user_uid"] == first["user_uid"]
  128. assert repo.get_identity(provider_b, "shared-subject")["user_uid"] == second["user_uid"]
  129. session_a = SessionManager(repo, secret="integration-secret").create(
  130. provider_uid=provider_a, subject="shared-subject", user_uid=first["user_uid"],
  131. roles=["viewer"], identity_source="oidc")
  132. session_b = SessionManager(repo, secret="integration-secret").create(
  133. provider_uid=provider_b, subject="shared-subject", user_uid=second["user_uid"],
  134. roles=["viewer"], identity_source="oidc")
  135. assert {row["uid"] for row in repo.list_sessions(provider_a, "shared-subject")} == {session_a.session_uid}
  136. assert {row["uid"] for row in repo.list_sessions(provider_b, "shared-subject")} == {session_b.session_uid}
  137. sync.apply(provider_uid=provider_a, source="scim", source_event_id="move-a", cursor="2",
  138. cursor_sequence=2, event_type="MOVER", subject="shared-subject",
  139. attributes=_identity_attributes(f"alice-a-{identity_database['marker']}", ["editor"]))
  140. assert repo.get_session(session_a.session_uid)["status"] == "revoked"
  141. assert repo.get_session(session_b.session_uid)["status"] == "active"
  142. assert repo.get_directory_checkpoint(provider_a, "scim") == {"cursor": "2", "cursor_sequence": 2}
  143. with pytest.raises(IdentityPolicyError, match="idempotency conflict"):
  144. sync.apply(provider_uid=provider_a, source="scim", source_event_id="move-a", cursor="2",
  145. cursor_sequence=2, event_type="MOVER", subject="shared-subject",
  146. attributes=_identity_attributes(f"changed-{identity_database['marker']}", ["editor"]))
  147. assert session.execute(text("SELECT count(*) FROM public.identity_directory_conflicts WHERE provider_uid=CAST(:p AS uuid)"),
  148. {"p": provider_a}).scalar_one() == 1
  149. rollback_subject = "rollback-subject"
  150. original_insert = repo._insert_directory_event
  151. def fail_after_mutation(event, result):
  152. raise RuntimeError("injected event persistence failure")
  153. monkeypatch.setattr(repo, "_insert_directory_event", fail_after_mutation)
  154. with pytest.raises(RuntimeError, match="injected"):
  155. sync.apply(provider_uid=provider_a, source="rollback", source_event_id="rollback-1", cursor="1",
  156. cursor_sequence=1, event_type="JOINER", subject=rollback_subject,
  157. attributes=_identity_attributes(f"rollback-{identity_database['marker']}"))
  158. monkeypatch.setattr(repo, "_insert_directory_event", original_insert)
  159. assert repo.get_identity(provider_a, rollback_subject) is None
  160. assert repo.get_directory_checkpoint(provider_a, "rollback") is None
  161. def test_organization_node_incremental_lifecycle_is_provider_isolated_and_auditable(identity_database):
  162. engine = identity_database["engine"]
  163. provider_a, provider_b = _provider(identity_database), _provider(identity_database)
  164. with Session(engine) as session:
  165. repo = PostgresIdentityRepository(session)
  166. sync = DirectorySynchronizer(repo, SessionManager(repo, secret="integration-secret"))
  167. department = {"external_id": "dept-ops", "display_name": "Operations", "action": "UPSERT",
  168. "attributes": {"cost_center": "CC-1"}}
  169. created = sync.apply(provider_uid=provider_a, source="directory", source_event_id="dept-1", cursor="1",
  170. cursor_sequence=1, event_type="DEPARTMENT", attributes=department)
  171. assert created["status"] == "active"
  172. disabled = sync.apply(provider_uid=provider_a, source="directory", source_event_id="dept-2", cursor="2",
  173. cursor_sequence=2, event_type="DEPARTMENT",
  174. attributes={"external_id": "dept-ops", "action": "DISABLE"})
  175. assert disabled["status"] == "disabled"
  176. restored = sync.apply(provider_uid=provider_a, source="directory", source_event_id="dept-3", cursor="3",
  177. cursor_sequence=3, event_type="DEPARTMENT",
  178. attributes={**department, "action": "RESTORE", "display_name": "Operations Restored"})
  179. assert restored["status"] == "active"
  180. assert sync.apply(provider_uid=provider_a, source="directory", source_event_id="dept-3", cursor="3",
  181. cursor_sequence=3, event_type="DEPARTMENT",
  182. attributes={**department, "action": "RESTORE", "display_name": "Operations Restored"})["idempotent"] is True
  183. sync.apply(provider_uid=provider_a, source="directory", source_event_id="group-1", cursor="4",
  184. cursor_sequence=4, event_type="GROUP",
  185. attributes={"external_id": "group-stewards", "display_name": "Stewards", "action": "UPSERT"})
  186. sync.apply(provider_uid=provider_b, source="directory", source_event_id="dept-other", cursor="1",
  187. cursor_sequence=1, event_type="DEPARTMENT", attributes=department)
  188. node_a = repo.get_organization_node(provider_a, "department", "dept-ops")
  189. node_b = repo.get_organization_node(provider_b, "department", "dept-ops")
  190. assert node_a["display_name"] == "Operations Restored"
  191. assert node_b["display_name"] == "Operations"
  192. repo.audit({"event_type": "directory_delta", "outcome": "success", "provider_uid": provider_a,
  193. "resource_type": "organization_node", "resource_uid": created["node_uid"],
  194. "safe_detail": sanitize_audit_detail({"event_type": "DEPARTMENT", "token": "must-not-persist"})})
  195. audit = session.execute(text("SELECT safe_detail FROM public.identity_audit_events WHERE provider_uid=CAST(:p AS uuid)"),
  196. {"p": provider_a}).scalar_one()
  197. assert audit == {"event_type": "DEPARTMENT"}
  198. def test_refresh_rotation_and_session_limit_are_atomic_under_concurrency(identity_database):
  199. engine = identity_database["engine"]
  200. provider = _provider(identity_database)
  201. with Session(engine) as session:
  202. repo = PostgresIdentityRepository(session)
  203. sync = DirectorySynchronizer(repo, SessionManager(repo, secret="integration-secret"))
  204. identity = sync.apply(provider_uid=provider, source="scim", source_event_id="join", cursor="1",
  205. cursor_sequence=1, event_type="JOINER", subject="concurrent-subject",
  206. attributes=_identity_attributes(f"concurrent-{identity_database['marker']}"))
  207. original = SessionManager(repo, secret="integration-secret").create(
  208. provider_uid=provider, subject="concurrent-subject", user_uid=identity["user_uid"],
  209. roles=["viewer"], identity_source="oidc")
  210. family_uid = repo.get_session(original.session_uid)["family_uid"]
  211. barrier = threading.Barrier(2)
  212. outcomes: list[str] = []
  213. outcome_lock = threading.Lock()
  214. def rotate_once():
  215. with Session(engine) as thread_session:
  216. manager = SessionManager(PostgresIdentityRepository(thread_session), secret="integration-secret")
  217. barrier.wait()
  218. try:
  219. manager.refresh(original.refresh_token)
  220. outcome = "rotated"
  221. except IdentityPolicyError as exc:
  222. outcome = str(exc)
  223. with outcome_lock:
  224. outcomes.append(outcome)
  225. threads = [threading.Thread(target=rotate_once) for _ in range(2)]
  226. for thread in threads:
  227. thread.start()
  228. for thread in threads:
  229. thread.join(timeout=15)
  230. assert all(not thread.is_alive() for thread in threads)
  231. assert outcomes.count("rotated") <= 1
  232. assert "refresh reuse detected" in outcomes
  233. with Session(engine) as session:
  234. statuses = session.execute(text("SELECT status FROM public.identity_sessions WHERE family_uid=CAST(:family AS uuid)"),
  235. {"family": family_uid}).scalars().all()
  236. assert statuses and set(statuses) == {"revoked"}
  237. create_barrier = threading.Barrier(6)
  238. create_errors: list[Exception] = []
  239. def create_once(index: int):
  240. with Session(engine) as thread_session:
  241. manager = SessionManager(PostgresIdentityRepository(thread_session), secret="integration-secret", max_sessions=2)
  242. create_barrier.wait()
  243. try:
  244. manager.create(provider_uid=provider, subject="concurrent-subject", user_uid=identity["user_uid"],
  245. roles=["viewer"], identity_source="oidc")
  246. except Exception as exc: # captured and asserted below
  247. create_errors.append(exc)
  248. creators = [threading.Thread(target=create_once, args=(index,)) for index in range(6)]
  249. for thread in creators:
  250. thread.start()
  251. for thread in creators:
  252. thread.join(timeout=15)
  253. assert not create_errors
  254. with Session(engine) as session:
  255. active_count = session.execute(text("SELECT count(*) FROM public.identity_sessions WHERE user_uid=CAST(:u AS uuid) AND status='active'"),
  256. {"u": identity["user_uid"]}).scalar_one()
  257. assert active_count <= 2
  258. def test_emergency_session_rechecks_local_admin_and_close_revokes(identity_database):
  259. engine = identity_database["engine"]
  260. users = identity_database["users"]
  261. now = datetime.now(UTC)
  262. with Session(engine) as session:
  263. repo = PostgresIdentityRepository(session)
  264. service = EmergencyAccess(repo, clock=lambda: now)
  265. request = service.request(requester_uid=users["requester"], account_uid=users["breakglass"],
  266. reason="integration outage", expires_at=now + timedelta(minutes=30),
  267. account_is_local_active_admin=True)
  268. service.approve(request["uid"], approver_uid=users["approver_a"])
  269. service.approve(request["uid"], approver_uid=users["approver_b"])
  270. active = service.activate(request["uid"])
  271. manager = SessionManager(repo, secret="integration-secret", clock=lambda: now)
  272. credentials = manager.create(subject=users["breakglass"], user_uid=users["breakglass"], roles=["admin"],
  273. identity_source="emergency", emergency_request_uid=request["uid"],
  274. emergency_expires_at=active["expires_at"])
  275. session.execute(text("DELETE FROM public.user_roles WHERE user_id=CAST(:u AS uuid)"), {"u": users["breakglass"]})
  276. session.commit()
  277. with pytest.raises(IdentityPolicyError, match="active local administrator"):
  278. manager.create(subject=users["breakglass"], user_uid=users["breakglass"], roles=["admin"],
  279. identity_source="emergency", emergency_request_uid=request["uid"],
  280. emergency_expires_at=active["expires_at"])
  281. session.execute(text("INSERT INTO public.user_roles(user_id,role_id) SELECT CAST(:u AS uuid),id FROM public.roles WHERE name='admin'"),
  282. {"u": users["breakglass"]})
  283. session.commit()
  284. service.close(request["uid"], actor_uid=users["approver_a"])
  285. assert repo.get_session(credentials.session_uid)["status"] == "revoked"
  286. assert repo.get_session(credentials.session_uid)["emergency_request_uid"] == request["uid"]
  287. def test_oidc_login_persistence_rolls_back_identity_session_and_exchange_together(identity_database, monkeypatch):
  288. engine = identity_database["engine"]
  289. provider = _provider(identity_database)
  290. user_uid = _uid()
  291. subject = "atomic-callback-subject"
  292. with Session(engine) as session:
  293. repo = PostgresIdentityRepository(session)
  294. repo.put_identity(provider, subject, {
  295. "provider_uid": provider, "user_uid": user_uid,
  296. "username": f"atomic-{identity_database['marker']}", "display_name": "Atomic Callback",
  297. "department": "security", "groups": ["stewards"], "roles": ["viewer"],
  298. "authorization_scope": {"environments": ["test"]}, "mapping_version": "integration-v1",
  299. "claims_digest": "b" * 64, "status": "active", "token_version": 1,
  300. }, commit=False)
  301. credentials = SessionManager(repo, secret="integration-secret").create(
  302. provider_uid=provider, subject=subject, user_uid=user_uid, roles=["viewer"],
  303. identity_source="oidc", commit=False)
  304. def fail_exchange(record, *, commit=True):
  305. del record, commit
  306. raise RuntimeError("injected exchange persistence failure")
  307. monkeypatch.setattr(repo, "insert_exchange_code", fail_exchange)
  308. with pytest.raises(RuntimeError, match="injected"):
  309. repo.insert_exchange_code({"uid": _uid(), "code_hash": "c" * 64,
  310. "session_uid": credentials.session_uid,
  311. "redirect_uri": "https://dataops.example.com/login/callback",
  312. "expires_at": datetime.now(UTC) + timedelta(minutes=2)}, commit=False)
  313. session.rollback()
  314. assert repo.get_identity(provider, subject) is None
  315. assert repo.get_session(credentials.session_uid) is None
  316. assert session.execute(text("SELECT count(*) FROM public.users WHERE id=CAST(:u AS uuid)"),
  317. {"u": user_uid}).scalar_one() == 0