oidc_group_sync.py 6.3 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153
  1. """OIDC group sync (#3107).
  2. Mirrors the LDAP group sync semantics from api/routes/auth.py
  3. (`_sync_ldap_user`) for OIDC logins: the provider's ``group_mapping``
  4. configures which Bambuddy groups the IdP is allowed to manage, and every
  5. login replaces only that managed slice — manual assignments to any other
  6. group survive (#1292, same fix the LDAP path needed).
  7. Differences from LDAP worth stating:
  8. - LDAP reads group DNs from the directory entry. OIDC reads group values
  9. from a JWT claim, and providers disagree on the shape: Keycloak ships a
  10. JSON array, Authentik ships an array, Logto and some legacy setups ship a
  11. space- or comma-separated string. ``extract_idp_groups`` accepts both.
  12. - LDAP has a ``default_group`` fallback when no mapped group matches. The
  13. OIDC path already has ``default_group_id`` applied at account creation,
  14. and re-asserting it on every login would fight manual upgrades: an admin
  15. who promotes an auto-created user out of Viewers would see the promotion
  16. reverted at the next SSO login. So the OIDC sync has no fallback — an
  17. empty resolved set simply means "the IdP grants none of the mapped
  18. groups", which removes exactly the mapped groups and nothing else.
  19. """
  20. from __future__ import annotations
  21. import contextlib
  22. import logging
  23. from sqlalchemy import select
  24. from sqlalchemy.ext.asyncio import AsyncSession
  25. from backend.app.models.group import Group
  26. from backend.app.models.user import User
  27. logger = logging.getLogger(__name__)
  28. # Bound on the claim parsing below. A legitimate groups claim holds tens of
  29. # entries; anything past this is a malformed or hostile token, and iterating
  30. # it would just burn cycles before the mapping lookup ignores the extras.
  31. _MAX_CLAIM_ITEMS = 500
  32. def extract_idp_groups(claim_value: object) -> list[str]:
  33. """Normalise a raw JWT claim value into a list of IdP group strings.
  34. Accepts the shapes seen in the wild:
  35. - list of strings (Keycloak, Authentik, most modern providers)
  36. - single string, space- or comma-separated (Logto, some legacy setups)
  37. - a single group name as a bare string
  38. Non-string entries, empty fragments, and obvious non-group payloads
  39. (dicts, numbers) are dropped rather than rejected: a provider adding an
  40. unexpected claim shape must not lock users out of their mapped groups.
  41. Duplicates are removed while preserving order (first occurrence wins).
  42. The result is bounded by _MAX_CLAIM_ITEMS for both shapes — a string
  43. claim splits into arbitrarily many fragments, so the slice applies after
  44. splitting, not only on the list path.
  45. """
  46. if claim_value is None:
  47. return []
  48. if isinstance(claim_value, list):
  49. raw_items = [item for item in claim_value[:_MAX_CLAIM_ITEMS] if isinstance(item, str)]
  50. elif isinstance(claim_value, str):
  51. # Space-separated is the OIDC convention (scope-style); commas are a
  52. # pragmatic extra since some IdPs stringify arrays that way.
  53. raw_items = claim_value.replace(",", " ").split(" ")[:_MAX_CLAIM_ITEMS]
  54. else:
  55. return []
  56. seen: set[str] = set()
  57. result: list[str] = []
  58. for item in raw_items:
  59. cleaned = item.strip()
  60. if cleaned and cleaned not in seen:
  61. seen.add(cleaned)
  62. result.append(cleaned)
  63. return result
  64. def resolve_oidc_group_mapping(idp_groups: list[str], group_mapping: dict[str, str]) -> list[str]:
  65. """Map IdP group values to Bambuddy group names (case-insensitive on the key).
  66. Same contract as ldap_service.resolve_group_mapping: returns the Bambuddy
  67. group names the user should hold among the mapped set. Values are compared
  68. case-insensitively because IdP group casing is not stable across providers
  69. (Keycloak preserves case; some LDAP-backed OIDC deployments downcase),
  70. and a case mismatch silently dropping a group is the failure mode an
  71. admin can least diagnose from the UI.
  72. """
  73. if not group_mapping:
  74. return []
  75. mapping_lower = {k.lower(): v for k, v in group_mapping.items()}
  76. result: list[str] = []
  77. for idp_group in idp_groups:
  78. mapped = mapping_lower.get(idp_group.lower())
  79. if mapped and mapped not in result:
  80. result.append(mapped)
  81. return result
  82. async def sync_oidc_user_groups(
  83. db: AsyncSession,
  84. user: User,
  85. *,
  86. group_claim: str,
  87. group_mapping: dict[str, str],
  88. claims: dict,
  89. ) -> None:
  90. """Apply the provider's group mapping to ``user`` after a successful login.
  91. Only Bambuddy groups named in ``group_mapping`` values are managed; every
  92. other group on the user is a manual assignment and is preserved. Commits
  93. only when something actually changed (the LDAP sync logs on change; same
  94. here). Never raises: a group-sync failure must not abort the login the
  95. token exchange already authenticated — the exception is logged and the
  96. user keeps the groups they had.
  97. """
  98. if not group_mapping:
  99. # No mapping configured: nothing is managed, so nothing may change.
  100. # This is the default state and must remain a no-op for upgrades.
  101. return
  102. try:
  103. mapped_names = resolve_oidc_group_mapping(extract_idp_groups(claims.get(group_claim)), group_mapping)
  104. # Only groups that exist locally can be granted; a mapping entry
  105. # pointing at a deleted group is skipped (the same dangling-FK
  106. # tolerance default_group_id documents for SQLite).
  107. if mapped_names:
  108. groups_result = await db.execute(select(Group).where(Group.name.in_(mapped_names)))
  109. target_groups = list(groups_result.scalars().all())
  110. else:
  111. target_groups = []
  112. managed_names = set(group_mapping.values())
  113. preserved = [g for g in user.groups if g.name not in managed_names]
  114. new_groups = preserved + target_groups
  115. current_ids = {g.id for g in user.groups}
  116. new_ids = {g.id for g in new_groups}
  117. if current_ids == new_ids:
  118. return
  119. user.groups = new_groups
  120. await db.commit()
  121. logger.info(
  122. "OIDC group sync: user %s groups -> %s",
  123. user.username,
  124. sorted(g.name for g in new_groups),
  125. )
  126. except Exception: # noqa: BLE001 -- login must survive a sync failure
  127. logger.exception("OIDC group sync failed for user %s; groups left unchanged", user.username)
  128. with contextlib.suppress(Exception):
  129. await db.rollback()