"""The WSGI application of the SCIM server, with no dependency."""
import json
from collections.abc import Callable
from collections.abc import Iterable
from http import HTTPStatus
from threading import Lock
from typing import TYPE_CHECKING
from urllib.parse import parse_qsl
from wsgiref.util import application_uri
from scim2_models import NotFoundException
from scim2_models import SCIMException
from scim2_models import ScimProvider
from scim2_server.applications.base import VERSION_PREFIX
from scim2_server.applications.base import BaseApplication
from scim2_server.handler import ScimHandler
from scim2_server.requests import ScimRequest
from scim2_server.responses import ScimResponse
from scim2_server.service import ScimService
from scim2_server.storage import ScimStorage
if TYPE_CHECKING:
from _typeshed.wsgi import StartResponse
from _typeshed.wsgi import WSGIApplication as WSGICallable
from _typeshed.wsgi import WSGIEnvironment
RESERVED_TENANTS = frozenset({"v2"})
def send_response(
response: ScimResponse, start_response: "StartResponse"
) -> Iterable[bytes]:
"""Send a SCIM response as a WSGI response."""
body = b"" if response.body is None else json.dumps(response.body).encode()
status = f"{response.status.value} {response.status.phrase}"
start_response(status, list(response.headers.items()))
return [body]
[docs]
class WSGIApplication(BaseApplication):
"""A WSGI application serving the SCIM protocol over a storage.
It reads the WSGI requests, serves them with a
:class:`~scim2_server.handler.ScimHandler`, and writes the responses.
:param storage: The storage of the resources.
:param provider: The description of the service.
:param service: The service serving the requests, built upon ``provider``.
"""
def __init__(
self,
storage: ScimStorage,
provider: ScimProvider,
service: ScimService | None = None,
):
super().__init__(provider, service)
self.storage = storage
self.handler = ScimHandler(self.service, storage)
[docs]
@staticmethod
def get_base_url(environ: "WSGIEnvironment") -> str:
"""Return the root URL of the SCIM endpoints, as the client sees it."""
return f"{application_uri(environ).rstrip('/')}{VERSION_PREFIX}"
[docs]
def read_request(self, environ: "WSGIEnvironment") -> ScimRequest:
"""Return the SCIM request of a WSGI request.
The body is read up to one byte more than the service accepts, so that
the service answers 413 to a larger body.
"""
headers = {
key.removeprefix("HTTP_").replace("_", "-"): value
for key, value in environ.items()
if key.startswith("HTTP_")
}
if environ.get("CONTENT_TYPE"):
headers["Content-Type"] = environ["CONTENT_TYPE"]
request = ScimRequest(
method=environ["REQUEST_METHOD"],
base_url=self.get_base_url(environ),
path=self.split_path(environ.get("PATH_INFO", "")),
query=dict(parse_qsl(environ.get("QUERY_STRING", ""))),
headers=headers,
)
request.body = self.read_body(environ, request)
return request
[docs]
def read_body(self, environ: "WSGIEnvironment", request: ScimRequest) -> bytes:
"""Read the body of a request, up to one byte more than the service accepts.
A body without :mdn:`Content-Length` is read up to its end, when the
server marks its input as terminated.
"""
try:
limit = self.service.max_body_size(request)
except SCIMException:
limit = None
if environ.get("CONTENT_LENGTH"):
size = int(environ["CONTENT_LENGTH"])
elif environ.get("wsgi.input_terminated"):
size = -1
else:
return b""
if limit is not None and (size < 0 or size > limit):
size = limit + 1
body: bytes = environ["wsgi.input"].read(size)
return body
[docs]
def dispatch_request(self, request: ScimRequest) -> ScimResponse:
"""Authenticate a request and serve it.
Override this method to act before or after a request is served.
"""
if self.needs_auth(request):
self.check_auth(request)
request.subject = self.get_subject(request)
return self.handler.handle(request)
[docs]
def serve(self, request: ScimRequest) -> ScimResponse:
"""Serve a request and return the response sent to the client."""
try:
response = self.dispatch_request(request)
except Exception as exception:
response = self.handle_exception(request, exception)
return self.finalize_response(request, response)
def __call__(
self, environ: "WSGIEnvironment", start_response: "StartResponse"
) -> Iterable[bytes]:
"""Serve a WSGI request."""
return send_response(self.serve(self.read_request(environ)), start_response)
[docs]
class TenantDispatcher:
"""A WSGI application serving each tenant with its own SCIM application.
The tenant is the first segment of the request path, as in the URL prefix
method of :rfc:`RFC 7644 §6.1 <7644#section-6.1>`: a request to ``/<tenant>/v2/Users`` is served by
the application of ``<tenant>``, mounted under ``/<tenant>``.
The application of a tenant is built on its first request, then kept for
the following ones.
:param factory: Build the application of a tenant from its name, or
return :data:`None` when the tenant does not exist.
"""
def __init__(self, factory: Callable[[str], "WSGICallable | None"]):
self.factory = factory
self.applications: dict[str, WSGICallable] = {}
self.lock = Lock()
[docs]
@staticmethod
def is_valid_tenant(tenant: str) -> bool:
"""Tell whether a name can identify a tenant.
The version segment is refused, so that a request without a tenant is
not served by a tenant named ``v2``.
"""
return bool(tenant) and "/" not in tenant and tenant not in RESERVED_TENANTS
[docs]
def select_tenant(self, environ: "WSGIEnvironment") -> str | None:
"""Return the tenant of a request, and move it from the path to the mount prefix.
Override this method to read the tenant from somewhere else, such as
a header or a sub-domain (:rfc:`RFC 7644 §6.1 <7644#section-6.1>`).
"""
path_info: str = environ.get("PATH_INFO", "")
_, _, path = path_info.partition("/")
tenant, separator, rest = path.partition("/")
if not self.is_valid_tenant(tenant):
return None
environ["SCRIPT_NAME"] = f"{environ.get('SCRIPT_NAME', '')}/{tenant}"
environ["PATH_INFO"] = separator + rest
return tenant
[docs]
def get_application(self, tenant: str) -> "WSGICallable | None":
"""Return the application of a tenant, building it on its first request."""
with self.lock:
if tenant not in self.applications:
application = self.factory(tenant)
if application is None:
return None
self.applications[tenant] = application
return self.applications[tenant]
def __call__(
self, environ: "WSGIEnvironment", start_response: "StartResponse"
) -> Iterable[bytes]:
"""Dispatch a request to the application of its tenant."""
tenant = self.select_tenant(environ)
application = self.get_application(tenant) if tenant is not None else None
if application is None:
error = NotFoundException(detail="Unknown tenant").to_error()
response = ScimResponse(HTTPStatus.NOT_FOUND, error.model_dump())
return send_response(response, start_response)
return application(environ, start_response)