test_oidc_env_apply.py 5.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160
  1. """Upserting the env-managed OIDC provider (#2593).
  2. Startup applies BAMBUDDY_OIDC_* to the database. The row is updated in place,
  3. never delete-recreated: user_oidc_links.provider_id is FK ON DELETE CASCADE, so
  4. recreating the provider would silently unlink every account bound to it.
  5. """
  6. from __future__ import annotations
  7. import pytest
  8. from sqlalchemy import select
  9. from backend.app.core.oidc_env import apply_env_oidc_provider
  10. from backend.app.models.oidc_provider import OIDCProvider
  11. REQUIRED = {
  12. "BAMBUDDY_OIDC_NAME": "Keycloak",
  13. "BAMBUDDY_OIDC_ISSUER_URL": "https://sso.example.com/realms/main",
  14. "BAMBUDDY_OIDC_CLIENT_ID": "bambuddy",
  15. "BAMBUDDY_OIDC_CLIENT_SECRET": "s3cr3t",
  16. }
  17. ALL_VARS = (
  18. *REQUIRED,
  19. "BAMBUDDY_OIDC_SCOPES",
  20. "BAMBUDDY_OIDC_ENABLED",
  21. "BAMBUDDY_OIDC_AUTO_CREATE_USERS",
  22. "BAMBUDDY_OIDC_AUTO_LINK_EXISTING",
  23. "BAMBUDDY_OIDC_EMAIL_CLAIM",
  24. "BAMBUDDY_OIDC_REQUIRE_EMAIL_VERIFIED",
  25. "BAMBUDDY_OIDC_ICON_URL",
  26. "BAMBUDDY_OIDC_AUTOLOGIN",
  27. )
  28. @pytest.fixture(autouse=True)
  29. def clean_env(monkeypatch):
  30. for key in ALL_VARS:
  31. monkeypatch.delenv(key, raising=False)
  32. def _configure(monkeypatch, **overrides):
  33. for key, value in REQUIRED.items():
  34. monkeypatch.setenv(key, value)
  35. for key, value in overrides.items():
  36. monkeypatch.setenv(key, value)
  37. async def _env_provider(db_session) -> OIDCProvider | None:
  38. result = await db_session.execute(select(OIDCProvider).where(OIDCProvider.is_env_managed.is_(True)))
  39. return result.scalar_one_or_none()
  40. @pytest.mark.asyncio
  41. async def test_creates_the_provider_from_env(db_session, monkeypatch):
  42. _configure(monkeypatch)
  43. await apply_env_oidc_provider(db_session)
  44. provider = await _env_provider(db_session)
  45. assert provider is not None
  46. assert provider.name == "Keycloak"
  47. assert provider.client_id == "bambuddy"
  48. assert provider.is_env_managed is True
  49. assert provider.client_secret == "s3cr3t" # property decrypts
  50. @pytest.mark.asyncio
  51. async def test_a_changed_var_updates_the_same_row(db_session, monkeypatch):
  52. """The id must survive: user_oidc_links references it with ON DELETE
  53. CASCADE, so a delete-recreate would unlink every bound account."""
  54. _configure(monkeypatch)
  55. await apply_env_oidc_provider(db_session)
  56. original_id = (await _env_provider(db_session)).id
  57. monkeypatch.setenv("BAMBUDDY_OIDC_CLIENT_ID", "rotated")
  58. await apply_env_oidc_provider(db_session)
  59. provider = await _env_provider(db_session)
  60. assert provider.id == original_id
  61. assert provider.client_id == "rotated"
  62. @pytest.mark.asyncio
  63. async def test_removing_the_env_config_disables_but_keeps_the_row(db_session, monkeypatch):
  64. _configure(monkeypatch)
  65. await apply_env_oidc_provider(db_session)
  66. original_id = (await _env_provider(db_session)).id
  67. for key in ALL_VARS:
  68. monkeypatch.delenv(key, raising=False)
  69. await apply_env_oidc_provider(db_session)
  70. provider = await _env_provider(db_session)
  71. assert provider is not None, "deleting would cascade away every account link"
  72. assert provider.id == original_id
  73. assert provider.is_enabled is False
  74. @pytest.mark.asyncio
  75. async def test_env_autologin_clears_it_on_other_providers(db_session, monkeypatch):
  76. """Only one provider may be the autologin target; the env one wins."""
  77. ui_provider = OIDCProvider(
  78. name="UI provider",
  79. issuer_url="https://other.example.com",
  80. client_id="ui",
  81. is_autologin=True,
  82. )
  83. ui_provider.client_secret = "ui-secret"
  84. db_session.add(ui_provider)
  85. await db_session.commit()
  86. _configure(monkeypatch, BAMBUDDY_OIDC_AUTOLOGIN="true")
  87. await apply_env_oidc_provider(db_session)
  88. await db_session.refresh(ui_provider)
  89. assert (await _env_provider(db_session)).is_autologin is True
  90. assert ui_provider.is_autologin is False
  91. @pytest.mark.asyncio
  92. async def test_a_ui_provider_is_otherwise_left_alone(db_session, monkeypatch):
  93. ui_provider = OIDCProvider(name="UI provider", issuer_url="https://other.example.com", client_id="ui")
  94. ui_provider.client_secret = "ui-secret"
  95. db_session.add(ui_provider)
  96. await db_session.commit()
  97. _configure(monkeypatch)
  98. await apply_env_oidc_provider(db_session)
  99. await db_session.refresh(ui_provider)
  100. assert ui_provider.is_env_managed is False
  101. assert ui_provider.is_enabled is True
  102. assert ui_provider.client_id == "ui"
  103. @pytest.mark.asyncio
  104. async def test_an_unsafe_auto_link_config_is_skipped_not_raised(db_session, monkeypatch):
  105. """auto-link + unverified email is the SEC-1 account-takeover shape. The
  106. schema rejects it for the UI, and env config must not be a way around that
  107. -- but a bad variable must not stop the app from booting either."""
  108. _configure(
  109. monkeypatch,
  110. BAMBUDDY_OIDC_AUTO_LINK_EXISTING="true",
  111. BAMBUDDY_OIDC_REQUIRE_EMAIL_VERIFIED="false",
  112. )
  113. await apply_env_oidc_provider(db_session)
  114. assert await _env_provider(db_session) is None
  115. @pytest.mark.asyncio
  116. async def test_applying_twice_without_changes_is_a_no_op(db_session, monkeypatch):
  117. """Every boot re-applies; the second run must not create a second row."""
  118. _configure(monkeypatch)
  119. await apply_env_oidc_provider(db_session)
  120. await apply_env_oidc_provider(db_session)
  121. result = await db_session.execute(select(OIDCProvider).where(OIDCProvider.is_env_managed.is_(True)))
  122. assert len(result.scalars().all()) == 1