|
1 | 1 | import asyncio |
| 2 | +import logging |
2 | 3 | import time |
3 | 4 | from collections.abc import Mapping, Sequence |
4 | 5 | from typing import Any, Optional, Union |
@@ -111,10 +112,23 @@ def __init__(self, options: ApiClientOptions): |
111 | 112 | raise ConfigurationError( |
112 | 113 | "organization_policy must be either 'required' or 'allow'" |
113 | 114 | ) |
114 | | - if options.organization_id is not None and options.organization_policy != "required": |
| 115 | + if options.organization_id is None: |
| 116 | + self._allowed_org_ids = None |
| 117 | + elif options.organization_policy != "required": |
115 | 118 | raise ConfigurationError( |
116 | 119 | "organization_id is only valid when organization_policy is 'required'" |
117 | 120 | ) |
| 121 | + else: |
| 122 | + org_ids = options.organization_id |
| 123 | + if isinstance(org_ids, str): |
| 124 | + org_ids = [org_ids] |
| 125 | + if not isinstance(org_ids, list) or not org_ids or not all( |
| 126 | + isinstance(o, str) and o.strip() for o in org_ids |
| 127 | + ): |
| 128 | + raise ConfigurationError( |
| 129 | + "organization_id must be a non-empty string or a non-empty list of non-empty strings" |
| 130 | + ) |
| 131 | + self._allowed_org_ids = frozenset(org_ids) |
118 | 132 |
|
119 | 133 | if options.cache_adapter: |
120 | 134 | self._discovery_cache = options.cache_adapter |
@@ -577,18 +591,13 @@ async def verify_access_token( |
577 | 591 | raise VerifyAccessTokenError(f"Missing required claim: {rc}") |
578 | 592 |
|
579 | 593 | # Organization policy enforcement |
580 | | - org_id = claims.get("org_id") |
581 | 594 | if self.options.organization_policy == "required": |
582 | | - if not org_id: |
| 595 | + org_id = claims.get("org_id") |
| 596 | + if not isinstance(org_id, str) or not org_id: |
583 | 597 | raise MissingOrganizationError("Token missing required 'org_id' claim") |
584 | | - allowed_orgs = self.options.organization_id |
585 | | - if allowed_orgs is not None: |
586 | | - if isinstance(allowed_orgs, str): |
587 | | - allowed_orgs = [allowed_orgs] |
588 | | - if org_id not in allowed_orgs: |
589 | | - raise OrganizationNotAllowedError( |
590 | | - f"Organization '{org_id}' is not in the allowed list" |
591 | | - ) |
| 598 | + if self._allowed_org_ids is not None and org_id not in self._allowed_org_ids: |
| 599 | + logging.warning("Rejected token with org_id %r not in the organization_id allowlist", org_id) |
| 600 | + raise OrganizationNotAllowedError("Token org_id is not in the allowed list") |
592 | 601 |
|
593 | 602 | return claims |
594 | 603 |
|
|
0 commit comments