diff --git a/README.md b/README.md index c0f2e92..ebffaa2 100644 --- a/README.md +++ b/README.md @@ -19,14 +19,15 @@ The only optional feature currently missing is support for Bulk operations ([RFC ## Usage ```shell -$ scim2-server [-h] [--schema SCHEMA] [--resource-type RESOURCE_TYPE] [--bearer-token BEARER_TOKEN] [--hostname HOSTNAME] [--port PORT] [--reverse-proxy] [--dump-resources DUMP_RESOURCES] [--debug] +$ scim2-server [-h] [--schema SCHEMA] [--resource-type RESOURCE_TYPE] [--service-provider-config SERVICE_PROVIDER_CONFIG] [--bearer-token BEARER_TOKEN] [--hostname HOSTNAME] [--port PORT] [--reverse-proxy] [--dump-resources DUMP_RESOURCES] [--debug] ``` - `-h`/`--help`: Show help message - `--reverse-proxy`: Allow using the provider behind a Reverse Proxy (required for URL rewriting). - `--schema`: Register schemas from specified JSON file. If not provided, loads the default schemas from RFC 7643. - `--resource-type`: Register resource types from specified JSON file. If not provided, loads the default resource types from RFC 7643. -- `--bearer-token`: Registers a bearer token that can be used for accessing the service. If no tokens are provided, anonymous access without authentication is allowed. +- `--service-provider-config`: Load the service provider configuration from specified JSON file. If not provided, loads the default configuration. +- `--bearer-token`: Registers a bearer token that can be used for accessing the service, and announces the bearer token authentication scheme. If no tokens are provided, anonymous access without authentication is allowed. - `--hostname`: The hostname to listen on. Defaults to `127.0.0.1`. - `--port`: The port to listen on. Defaults to `8080`. - `--dump-resources`: Dump a JSON document containing all resources when the provider exits normally. diff --git a/pyproject.toml b/pyproject.toml index 545d239..a570df7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -31,7 +31,7 @@ classifiers = [ requires-python = ">= 3.10" dependencies = [ "scim2-filter-parser>=0.7.0", - "scim2-models>=0.8.1", + "scim2-models>=0.8.2", "werkzeug>=3.0.3", ] diff --git a/scim2_server/backend.py b/scim2_server/backend.py index 165d85b..2e1ea85 100644 --- a/scim2_server/backend.py +++ b/scim2_server/backend.py @@ -1,37 +1,31 @@ import dataclasses import datetime -import operator import pickle import uuid +from inspect import isclass from threading import Lock +from typing import Any from typing import Union -from scim2_models import Attribute from scim2_models import BaseModel from scim2_models import CaseExact from scim2_models import Extension from scim2_models import Meta from scim2_models import Resource from scim2_models import ResourceType -from scim2_models import Schema from scim2_models import ScimFilter from scim2_models import SearchRequest from scim2_models import Uniqueness from scim2_models import UniquenessException from werkzeug.http import generate_etag -from scim2_server.operators import ResolveSortOperator -from scim2_server.utils import get_by_alias - class Backend: - """The base class for a SCIM provider backend.""" + """The base class for a SCIM provider backend. - def __init__(self): - self.schemas: dict[str, Schema] = {} - self.resource_types: dict[str, ResourceType] = {} - self.resource_types_by_endpoint: dict[str, ResourceType] = {} - self.models_dict: dict[str, BaseModel] = {} + A backend only stores resources: what the service serves is described by the + :class:`~scim2_models.ScimProvider` of the application. + """ def __enter__(self): """Allow the backend to be used as a context manager. @@ -44,90 +38,33 @@ def __exit__(self, exc_type, exc_val, exc_tb): """Exit the transaction.""" pass - def register_schema(self, schema: Schema): - """Register a Schema for use with the backend.""" - self.schemas[schema.id] = schema - - def get_schemas(self): - """Return all schemas registered with the backend.""" - return self.schemas.values() - - def get_schema(self, schema_id: str) -> Schema | None: - """Get a schema by its id.""" - return self.schemas.get(schema_id) - - def register_resource_type(self, resource_type: ResourceType): - """Register a ResourceType for use with the backend. - - The schemas used for the resource and its extensions must have - been registered with the Backend beforehand. - """ - if resource_type.schema_ not in self.schemas: - raise RuntimeError(f"Unknown schema: {resource_type.schema_}") - for resource_extension in resource_type.schema_extensions or []: - if resource_extension.schema_ not in self.schemas: - raise RuntimeError(f"Unknown schema: {resource_extension.schema_}") - - self.resource_types[resource_type.id] = resource_type - self.resource_types_by_endpoint[resource_type.endpoint.lower()] = resource_type - - extensions = [ - Extension.from_schema(self.get_schema(se.schema_)) - for se in resource_type.schema_extensions or [] - ] - base_schema = self.get_schema(resource_type.schema_) - self.models_dict[resource_type.id] = Resource.from_schema(base_schema) - if extensions: - self.models_dict[resource_type.id] = self.models_dict[resource_type.id][ - Union[tuple(extensions)] # noqa: UP007 - ] - - def get_resource_types(self): - """Return all resource types registered with the backend.""" - return self.resource_types.values() - - def get_resource_type(self, resource_type_id: str) -> ResourceType | None: - """Return the resource type by its id.""" - return self.resource_types.get(resource_type_id) - - def get_resource_type_by_endpoint(self, endpoint: str) -> ResourceType | None: - """Return the resource type by its endpoint.""" - return self.resource_types_by_endpoint.get(endpoint.lower()) - - def get_model(self, resource_type_id: str) -> BaseModel | None: - """Return the Pydantic Python model for a given resource type.""" - return self.models_dict.get(resource_type_id) - - def get_models(self): - """Return all Pydantic Python models for all known resource types.""" - return self.models_dict.values() - def query_resources( self, search_request: SearchRequest, - resource_type_id: str | None = None, + resource_type: ResourceType | None = None, ) -> tuple[int, list[Resource]]: """Query the backend for a set of resources. :param search_request: SearchRequest instance describing the query. - :param resource_type_id: ID of the resource type to query. If - None, all resource types are queried. + :param resource_type: The resource type to query. If None, all + resource types are queried. :return: A tuple of "total results" and a List of found Resources. The List must contain a copy of resources. Mutating elements in the List must not modify the data stored in the backend. :raises TooManyException: If the backend only supports querying - for one resource type at a time, setting resource_type_id to + for one resource type at a time, setting resource_type to None the backend may raise TooManyException. """ raise NotImplementedError - def get_resource(self, resource_type_id: str, object_id: str) -> Resource | None: + def get_resource( + self, resource_type: ResourceType, object_id: str + ) -> Resource | None: """Query the backend for a resources by its ID. - :param resource_type_id: ID of the resource type to get the - object from. + :param resource_type: The resource type to get the object from. :param object_id: ID of the object to get. :return: The resource object if it exists, None otherwise. The resource must be a copy, modifying it must not change the @@ -135,22 +72,22 @@ def get_resource(self, resource_type_id: str, object_id: str) -> Resource | None """ raise NotImplementedError - def delete_resource(self, resource_type_id: str, object_id: str) -> bool: + def delete_resource(self, resource_type: ResourceType, object_id: str) -> bool: """Delete a resource. - :param resource_type_id: ID of the resource type to delete the - object from. + :param resource_type: The resource type to delete the object + from. :param object_id: ID of the object to delete. :return: True if the resource was deleted, False otherwise. """ raise NotImplementedError def create_resource( - self, resource_type_id: str, resource: Resource + self, resource_type: ResourceType, resource: Resource ) -> Resource | None: """Create a resource. - :param resource_type_id: ID of the resource type to create. + :param resource_type: The resource type to create. :param resource: Resource to create. :return: The created resource. Creation should set system- defined attributes (ID, Metadata). May be the same object @@ -159,11 +96,11 @@ def create_resource( raise NotImplementedError def update_resource( - self, resource_type_id: str, resource: Resource + self, resource_type: ResourceType, resource: Resource ) -> Resource | None: """Update a resource. The resource is identified by its ID. - :param resource_type_id: ID of the resource type to update. + :param resource_type: The resource type to update. :param resource: Resource to update. :return: The updated resource. Updating should update the "meta.lastModified" data. May be the same object that is @@ -181,56 +118,57 @@ class InMemoryBackend(Backend): implementation simple. """ - @dataclasses.dataclass + @dataclasses.dataclass(frozen=True) class UniquenessDescriptor: """Used to mimic uniqueness constraints e.g. from a SQL database.""" - schema: str | None - attribute_name: str + extension: str | None + field_name: str case_exact: bool + schema: str - def get_attribute(self, resource: Resource): - if self.schema is not None: - schema_field = get_by_alias(type(resource), self.schema) - resource = getattr(resource, schema_field) + def is_declared_by(self, resource: Resource) -> bool: + """Tell whether the model of a resource holds the schema of the attribute.""" + model = type(resource) + return self.schema == model.__schema__ or ( + self.schema in model.get_extension_models() + ) - attribute_field = get_by_alias(type(resource), self.attribute_name) - result = getattr(resource, attribute_field) - if not self.case_exact: - result = result.lower() - return result + def get_attribute(self, resource: Resource) -> Any: + holder = getattr(resource, self.extension) if self.extension else resource + value = getattr(holder, self.field_name, None) if holder else None + if isinstance(value, str) and not self.case_exact: + return value.casefold() + return value @classmethod def collect_unique_attrs( - cls, attributes: list[Attribute], schema: str | None + cls, model: type[BaseModel], extension: str | None = None ) -> list[UniquenessDescriptor]: - ret = [] - for attr in attributes: - if attr.uniqueness != Uniqueness.none: - ret.append( - cls.UniquenessDescriptor( - schema, attr.name, attr.case_exact == CaseExact.true - ) - ) - return ret + """Return the uniqueness constraints the annotations of a model declare. - @classmethod - def collect_resource_unique_attrs( - cls, resource_type: ResourceType, schemas: dict[str, Schema] - ) -> list[list[UniquenessDescriptor]]: - ret = cls.collect_unique_attrs(schemas[resource_type.schema_].attributes, None) - for extension in resource_type.schema_extensions or []: - ret.extend( - InMemoryBackend.collect_unique_attrs( - schemas[extension.schema_].attributes, extension.schema_ - ) + The ``id`` is left out: the backend issues it, so it cannot clash. + """ + descriptors = [ + cls.UniquenessDescriptor( + extension, + field_name, + model.get_field_annotation(field_name, CaseExact) == CaseExact.true, + str(model.__schema__), ) - return ret + for field_name in model.model_fields + if field_name != "id" + and model.get_field_annotation(field_name, Uniqueness) != Uniqueness.none + ] + for field_name in model.model_fields: + root_type = model.get_field_root_type(field_name) + if isclass(root_type) and issubclass(root_type, Extension): + descriptors.extend(cls.collect_unique_attrs(root_type, field_name)) + return descriptors def __init__(self): super().__init__() self.resources: list[Resource] = [] - self.unique_attributes: dict[str, list[list[str]]] = {} self.lock: Lock = Lock() def __enter__(self): @@ -247,56 +185,29 @@ def __exit__(self, exc_type, exc_val, exc_tb): super().__exit__(exc_type, exc_val, exc_tb) self.lock.release() - def register_resource_type(self, resource_type: ResourceType): - super().register_resource_type(resource_type) - self.unique_attributes[resource_type.id] = self.collect_resource_unique_attrs( - resource_type, self.schemas - ) - def query_resources( self, search_request: SearchRequest, - resource_type_id: str | None = None, + resource_type: ResourceType | None = None, ) -> tuple[int, list[Resource]]: start_index = (search_request.start_index or 1) - 1 + candidates = [ + r + for r in self.resources + if resource_type is None or self._is_of_type(r, resource_type) + ] + scim_filter = search_request.filter - if scim_filter is not None and not scim_filter.models: - models = ( - list(self.models_dict.values()) - if resource_type_id is None - else [self.models_dict[resource_type_id]] - ) - scim_filter = ScimFilter[Union[tuple(models)]](str(scim_filter)) # noqa: UP007 + if scim_filter is not None and not scim_filter.models and candidates: + models = tuple(dict.fromkeys(type(r) for r in candidates)) + scim_filter = ScimFilter[Union[models]](str(scim_filter)) # noqa: UP007 found_resources = [ - r - for r in self.resources - if (resource_type_id is None or r.meta.resource_type == resource_type_id) - and (scim_filter is None or scim_filter.match(r)) + r for r in candidates if scim_filter is None or scim_filter.match(r) ] - if search_request.sort_by is not None: - descending = search_request.sort_order == SearchRequest.SortOrder.descending - sort_operator = ResolveSortOperator(str(search_request.sort_by)) - - # To ensure that unset attributes are sorted last (when ascending, as defined in the RFC), - # we have to divide the result set into a set and unset subset. - unset_values = [] - set_values = [] - for resource in found_resources: - result = sort_operator(resource) - if result is None: - unset_values.append(resource) - else: - set_values.append((resource, result)) - - set_values.sort(key=operator.itemgetter(1), reverse=descending) - set_values = [value[0] for value in set_values] - if descending: - found_resources = unset_values + set_values - else: - found_resources = set_values + unset_values + found_resources = search_request.sort(found_resources) total_results = len(found_resources) found_resources = found_resources[start_index:] @@ -304,61 +215,81 @@ def query_resources( found_resources = found_resources[: search_request.count] return total_results, found_resources - def _get_resource_idx(self, resource_type_id: str, object_id: str) -> int | None: + def _is_of_type(self, resource: Resource, resource_type: ResourceType) -> bool: + """Tell whether a resource belongs to a resource type. + + RFC 7643 §3.1 has meta.resourceType carry the name of the resource type, + which may differ from its id. + """ + return resource.meta.resource_type == resource_type.name + + def _get_resource_idx( + self, resource_type: ResourceType, object_id: str + ) -> int | None: return next( ( idx for idx, r in enumerate(self.resources) - if r.meta.resource_type == resource_type_id and r.id == object_id + if self._is_of_type(r, resource_type) and r.id == object_id ), None, ) - def get_resource(self, resource_type_id: str, object_id: str) -> Resource | None: - resource_dict_idx = self._get_resource_idx(resource_type_id, object_id) + def get_resource( + self, resource_type: ResourceType, object_id: str + ) -> Resource | None: + resource_dict_idx = self._get_resource_idx(resource_type, object_id) if resource_dict_idx is not None: return self.resources[resource_dict_idx].model_copy(deep=True) return None - def delete_resource(self, resource_type_id: str, object_id: str) -> bool: - found = self.get_resource(resource_type_id, object_id) + def delete_resource(self, resource_type: ResourceType, object_id: str) -> bool: + found = self.get_resource(resource_type, object_id) if found: self.resources = [ r for r in self.resources - if not (r.meta.resource_type == resource_type_id and r.id == object_id) + if not (self._is_of_type(r, resource_type) and r.id == object_id) ] return True return False def create_resource( - self, resource_type_id: str, resource: Resource + self, resource_type: ResourceType, resource: Resource ) -> Resource | None: resource = resource.model_copy(deep=True) resource.id = uuid.uuid4().hex utcnow = datetime.datetime.now(datetime.timezone.utc) resource.meta = Meta( - resource_type=self.resource_types[resource_type_id].name, + resource_type=resource_type.name, created=utcnow, last_modified=utcnow, - location="/v2" - + self.resource_types[resource_type_id].endpoint - + "/" - + resource.id, + location="/v2" + resource_type.endpoint + "/" + resource.id, ) self._touch_resource(resource, utcnow) - - for unique_attribute in self.unique_attributes[resource_type_id]: - new_value = unique_attribute.get_attribute(resource) - for existing_resource in self.resources: - if existing_resource.meta.resource_type == resource_type_id: - existing_value = unique_attribute.get_attribute(existing_resource) - if existing_value == new_value: - raise UniquenessException() - + self._check_uniqueness(resource) self.resources.append(resource) return resource + def _check_uniqueness(self, resource: Resource): + """Refuse a resource sharing a unique value with another one of the same schema. + + RFC 7643 erratum 8279 scopes the uniqueness to the resources using the + schema that declares the attribute, whatever their resource type. A + missing value never clashes, as a SQL NULL does not. + """ + for unique_attribute in self.collect_unique_attrs(type(resource)): + value = unique_attribute.get_attribute(resource) + if value is None: + continue + for existing_resource in self.resources: + if ( + unique_attribute.is_declared_by(existing_resource) + and existing_resource.id != resource.id + and unique_attribute.get_attribute(existing_resource) == value + ): + raise UniquenessException() + @staticmethod def _touch_resource(resource: Resource, last_modified: datetime.datetime): """Touches a resource (updates last_modified and version). @@ -371,30 +302,16 @@ def _touch_resource(resource: Resource, last_modified: datetime.datetime): resource.meta.version = f'W/"{etag}"' def update_resource( - self, resource_type_id: str, resource: Resource + self, resource_type: ResourceType, resource: Resource ) -> Resource | None: - found_res_idx = self._get_resource_idx(resource_type_id, resource.id) + found_res_idx = self._get_resource_idx(resource_type, resource.id) if found_res_idx is not None: - updated_resource = self.models_dict[resource_type_id].model_validate( - resource.model_dump() - ) + updated_resource = type(resource).model_validate(resource.model_dump()) self._touch_resource( updated_resource, datetime.datetime.now(datetime.timezone.utc) ) - for unique_attribute in self.unique_attributes[resource_type_id]: - new_value = unique_attribute.get_attribute(updated_resource) - for existing_resource in self.resources: - if ( - existing_resource.meta.resource_type == resource_type_id - and existing_resource.id != updated_resource.id - ): - existing_value = unique_attribute.get_attribute( - existing_resource - ) - if existing_value == new_value: - raise UniquenessException() - + self._check_uniqueness(updated_resource) self.resources[found_res_idx] = updated_resource return updated_resource return None diff --git a/scim2_server/cli.py b/scim2_server/cli.py index 75cbd75..de59feb 100644 --- a/scim2_server/cli.py +++ b/scim2_server/cli.py @@ -3,14 +3,25 @@ import logging import pprint +from scim2_models import AuthenticationScheme from scim2_models import ResourceType from scim2_models import Schema +from scim2_models import ScimProvider +from scim2_models import ServiceProviderConfig from werkzeug.middleware.proxy_fix import ProxyFix from scim2_server.backend import InMemoryBackend -from scim2_server.provider import SCIMProvider +from scim2_server.provider import SCIMApplication from scim2_server.utils import load_default_resource_types from scim2_server.utils import load_default_schemas +from scim2_server.utils import load_default_service_provider_config + +BEARER_TOKEN_SCHEME = AuthenticationScheme( + type="oauthbearertoken", + name="bearer_token", + description="HTTP Bearer Token", + spec_uri="https://datatracker.ietf.org/doc/html/rfc6750", +) def log_environ(handler): @@ -31,6 +42,11 @@ def main(): parser.add_argument( "--resource-type", type=argparse.FileType("r"), help="Resource Type definitions" ) + parser.add_argument( + "--service-provider-config", + type=argparse.FileType("r"), + help="Service provider configuration", + ) parser.add_argument("--bearer-token", action="append", help="Add Bearer Token") parser.add_argument("--hostname", default="127.0.0.1", help="Hostname") parser.add_argument("--port", default=8080, type=int, help="Port number") @@ -55,28 +71,38 @@ def main(): from werkzeug.serving import run_simple - backend = InMemoryBackend() - app = SCIMProvider(backend) - if args.schema is None: - for schema in load_default_schemas().values(): - app.register_schema(schema) + schemas = load_default_schemas().values() else: - def_sch = json.load(args.schema) - for sc in def_sch: - schema = Schema.model_validate(sc) - app.register_schema(schema) - args.schema.close() + with args.schema: + schemas = [Schema.model_validate(sc) for sc in json.load(args.schema)] if args.resource_type is None: - for resource_type in load_default_resource_types().values(): - app.register_resource_type(resource_type) + resource_types = load_default_resource_types().values() + else: + with args.resource_type: + resource_types = [ + ResourceType.model_validate(rt) for rt in json.load(args.resource_type) + ] + + if args.service_provider_config is None: + config = load_default_service_provider_config() else: - def_rt = json.load(args.resource_type) - for rt in def_rt: - resource_type = ResourceType.model_validate(rt) - app.register_resource_type(resource_type) - args.resource_type.close() + with args.service_provider_config: + config = ServiceProviderConfig.model_validate( + json.load(args.service_provider_config) + ) + + if args.bearer_token is not None: + config.authentication_schemes = [ + *(config.authentication_schemes or []), + BEARER_TOKEN_SCHEME, + ] + + backend = InMemoryBackend() + app = SCIMApplication( + backend, ScimProvider.from_discovery(schemas, resource_types, config=config) + ) if args.bearer_token is not None: for bearer_token in args.bearer_token: diff --git a/scim2_server/operators.py b/scim2_server/operators.py index 3df8943..706d21d 100644 --- a/scim2_server/operators.py +++ b/scim2_server/operators.py @@ -4,7 +4,6 @@ from scim2_filter_parser.lexer import SCIMLexer from scim2_filter_parser.parser import SCIMParser from scim2_models import BaseModel -from scim2_models import CaseExact from scim2_models import InvalidPathException from scim2_models import InvalidValueException from scim2_models import Mutability @@ -386,91 +385,3 @@ def operation( value.add_result(model, alias) else: value.add_result_index(model, alias, index) - - -class ResolveSortOperator(ResolveOperator): - """Implement sorting in a helper Operator, according to RFC 7644, Section 3.4.2.3. - - The ResolveResult returned by this operator contains at most 1 value, according to - the specification: - "[...] if it's a multi-valued attribute, resources are sorted by the value of the - primary attribute (see Section 2.4 of [RFC7643]), if any, or else the first value - in the list, if any. [...]". - - Since a Query can result in resources of different types, sorting by an attribute - that is not defined for a certain resource type does not result in an error. No - value is returned and the resource is sorted as if the attribute on the resource - is not set. - """ - - def __init__(self, path: str | None): - super().__init__(path) - - def alias_forbidden(self, model: BaseModel, alias: str | None) -> bool: - return ( - not alias - or model.get_field_annotation(alias, Mutability) == Mutability.write_only - or model.get_field_annotation(alias, Returned) == Returned.never - ) - - def set_value_case_exact(self, value: Any, case_exact: CaseExact): - if isinstance(value, str) and case_exact == CaseExact.false: - value = value.lower() - self.value = value - - def evaluate_value_for_complex(self, model: BaseModel, alias: str): - sub_attribute_alias = get_by_alias(type(model), alias, True) - if self.alias_forbidden(model, sub_attribute_alias): - return - case_exact = model.get_field_annotation(sub_attribute_alias, CaseExact) - sub_attribute_value = getattr(model, sub_attribute_alias) - self.set_value_case_exact(sub_attribute_value, case_exact) - - def __call__(self, model: BaseModel): - self.value = None - if self.path: - model, path = self.parse_path(model) - if not path: - return - sub_attribute = path["sub_attribute"] or "value" - - attribute_alias = get_by_alias(type(model), path["attribute"], True) - if self.alias_forbidden(model, attribute_alias): - return - - case_exact = model.get_field_annotation(attribute_alias, CaseExact) - attribute_value = getattr(model, attribute_alias) - if not attribute_value: - return - - if isinstance(attribute_value, list): - if path["condition"]: - token_stream = SCIMLexer().tokenize(path["condition"]) - condition = SCIMParser().parse(token_stream) - attribute_value = [ - model - for model in attribute_value - if evaluate_filter(model, condition) - ] - candidate = self.select_candidate(attribute_value) - if isinstance(candidate, BaseModel): - self.evaluate_value_for_complex(candidate, sub_attribute) - else: - self.set_value_case_exact(candidate, case_exact) - elif isinstance(attribute_value, BaseModel): - if not path["condition"]: - self.evaluate_value_for_complex(attribute_value, sub_attribute) - else: - if not path["condition"] and not path["sub_attribute"]: - self.set_value_case_exact(attribute_value, case_exact) - return self.value - - def select_candidate(self, values: list[Any]) -> tuple[Any | None, int]: - """Select a viable candidate from a list of possible values.""" - for value in values: - primary = getattr(value, "primary", False) - if primary: - return value - if values: - return values[0] - return None diff --git a/scim2_server/provider.py b/scim2_server/provider.py index 7d14605..cc8e959 100644 --- a/scim2_server/provider.py +++ b/scim2_server/provider.py @@ -1,17 +1,14 @@ import itertools import json import logging -import traceback +from typing import Any from typing import Union +from typing import cast from urllib.parse import urljoin from pydantic import ValidationError -from scim2_models import AuthenticationScheme -from scim2_models import Bulk -from scim2_models import ChangePassword from scim2_models import Context from scim2_models import Error -from scim2_models import ETag from scim2_models import Filter from scim2_models import ListResponse from scim2_models import Meta @@ -22,6 +19,7 @@ from scim2_models import ResponseParameters from scim2_models import Schema from scim2_models import SCIMException +from scim2_models import ScimProvider from scim2_models import SearchRequest from scim2_models import ServiceProviderConfig from scim2_models import Sort @@ -40,7 +38,7 @@ from scim2_server.backend import Backend from scim2_server.operators import patch_resource -from scim2_server.utils import merge_resources +from scim2_server.utils import load_default_service_provider_config SEARCH_REQUEST_PARAMETERS = ( "attributes", @@ -53,16 +51,17 @@ ) -class SCIMProvider: +class SCIMApplication: """A WSGI application implementing a SCIM provider (server).""" - def __init__(self, backend: Backend): + def __init__(self, backend: Backend, provider: ScimProvider): self.bearer_tokens = set() self.backend = backend - self.page_size = 50 - self.log = logging.getLogger("SCIMProvider") + self.provider = provider + self.config = provider.config or load_default_service_provider_config() + self.log = logging.getLogger("SCIMApplication") - # Register the URL mapping. The endpoint refers to the name of the function to be called in this SCIMProvider ("call_" + endpoint). + # Register the URL mapping. The endpoint refers to the name of the function to be called in this SCIMApplication ("call_" + endpoint). rules = itertools.chain.from_iterable( [ Rule( @@ -123,93 +122,133 @@ def __init__(self, backend: Backend): self.url_map = Map(rules) - @staticmethod - def adjust_location( - request: Request, resource: Resource, cp=False - ) -> Resource | None: - """Adjust the "meta.location" attribute of a resource to match the hostname the client used to access this server. If a static URL is used,. - - :param request: The werkzeug request object - :param resource: The resource to modify - :param cp: Whether to return a modified copy of the resource or - to modify the resource in-place. + def get_model(self, resource_type: ResourceType) -> type[Resource]: + """Return the model of a resource type, its extensions included.""" + return cast(type[Resource], self.provider.model_for(resource_type)) + + def get_models(self) -> list[type[Resource]]: + """Return the models of every resource type.""" + return [self.get_model(rt) for rt in self.provider.resource_types] + + def get_resource_type_by_endpoint(self, endpoint: str) -> ResourceType | None: + """Return the resource type an endpoint serves.""" + return next( + ( + resource_type + for resource_type in self.provider.resource_types + if resource_type.endpoint.lstrip("/").casefold() + == endpoint.lstrip("/").casefold() + ), + None, + ) + + @property + def etag_supported(self) -> bool: + """Whether the configuration declares the resources versioned with ETags.""" + return bool(self.config.etag and self.config.etag.supported) + + def publish(self, request: Request, resource: Resource) -> Resource: + """Return a copy of a resource in the form sent to the client. + + Its location is made absolute from the URL the client requested, and + its version is left out when the service does not support ETags. """ - location = urljoin(request.url + "/", resource.meta.location) - if cp: - obj = resource.model_copy(deep=True) - obj.meta.location = location - return obj + update: dict[str, Any] = { + "location": urljoin(request.url + "/", resource.meta.location) + } + if not self.etag_supported: + update["version"] = None + return resource.model_copy( + update={"meta": resource.meta.model_copy(update=update)} + ) - resource.meta.location = location + @staticmethod + def etag_header(resource: Resource) -> dict[str, str]: + """Return the ETag header of a published resource, if it has a version.""" + return {"ETag": resource.meta.version} if resource.meta.version else {} def apply_patch_operation(self, resource: Resource, patch_operation): """Apply a PATCH operation to a resource.""" for op in patch_operation.operations: patch_resource(resource, op) - @staticmethod - def continue_etag(request: Request, resource: Resource) -> bool: - """Given a request and a resource, checks whether the ETag matches and allows continuing with the request. + def check_preconditions(self, request: Request, resource: Resource) -> bool: + """Evaluate the "If-Match" and "If-None-Match" headers against a resource. + + RFC 7232 §6 evaluates "If-Match" first: a failed "If-Match" answers + 412 whatever the method, a failed "If-None-Match" answers 304 to a GET + and 412 otherwise. - If the HTTP header "If-Match" is set, the request may only - continue if the ETag matches. If the HTTP header "If-None-Match" - is set, the request may only continue if the ETag does not - match. + :return: :data:`False` when a GET should answer 304 Not Modified. + :raises PreconditionFailed: When the method must not be performed. """ - cont = True - resource_version, _ = unquote_etag(resource.meta.version) - if request.if_none_match: - cont &= not request.if_none_match.contains_weak(resource_version) - if request.if_match: - cont &= request.if_match.contains_weak(resource_version) - return cont + # A service that does not support ETags has no tag to match: RFC 7232 + # §3.1 fails an If-Match listing tags, and lets "*" pass. + version, _ = ( + unquote_etag(resource.meta.version) if self.etag_supported else (None, None) + ) + # RFC 7232 §3.1 compares If-Match strongly, which would never match + # the weak ETags RFC 7644 §3.14 recommends and sends in its example. + if request.if_match and not request.if_match.contains_weak(version): + raise PreconditionFailed + + if request.if_none_match and request.if_none_match.contains_weak(version): + if request.method == "GET": + return False + raise PreconditionFailed + + return True def call_single_resource( self, request: Request, resource_endpoint: str, resource_id: str, **kwargs ) -> Response: - find_endpoint = "/" + resource_endpoint - resource_type = self.backend.get_resource_type_by_endpoint(find_endpoint) + resource_type = self.get_resource_type_by_endpoint(resource_endpoint) if not resource_type: raise NotFound match request.method: case "GET": - if resource := self.backend.get_resource(resource_type.id, resource_id): - if self.continue_etag(request, resource): - response_parameters = self.get_response_parameters( - request, self.backend.get_model(resource_type.id) - ) - self.adjust_location(request, resource) - return self.make_response( - resource.model_dump( - scim_ctx=Context.RESOURCE_QUERY_RESPONSE, - response_parameters=response_parameters, - ) - ) - else: - return self.make_response(None, status=304) - raise NotFound + resource = self.backend.get_resource(resource_type, resource_id) + if resource is None: + raise NotFound + resource = self.publish(request, resource) + if not self.check_preconditions(request, resource): + # RFC 7232 §4.1: a 304 carries the ETag a 200 would have + return self.make_response( + None, status=304, headers=self.etag_header(resource) + ) + + response_parameters = self.get_response_parameters( + request, self.get_model(resource_type) + ) + return self.make_response( + resource.model_dump( + scim_ctx=Context.RESOURCE_QUERY_RESPONSE, + response_parameters=response_parameters, + ) + ) case "DELETE": - if self.backend.delete_resource(resource_type.id, resource_id): - return self.make_response(None, 204) - else: + resource = self.backend.get_resource(resource_type, resource_id) + if resource is None: raise NotFound + self.check_preconditions(request, resource) + self.backend.delete_resource(resource_type, resource_id) + return self.make_response(None, 204) case "PUT": response_parameters = self.get_response_parameters( - request, self.backend.get_model(resource_type.id) + request, self.get_model(resource_type) ) - resource = self.backend.get_resource(resource_type.id, resource_id) + resource = self.backend.get_resource(resource_type, resource_id) if resource is None: raise NotFound - if not self.continue_etag(request, resource): - raise PreconditionFailed - - updated_attributes = self.backend.get_model( - resource_type.id - ).model_validate(request.json) - merge_resources(resource, updated_attributes) - updated = self.backend.update_resource(resource_type.id, resource) - self.adjust_location(request, updated) + self.check_preconditions(request, resource) + + replacement = self.get_model(resource_type).model_validate( + request.json, scim_ctx=Context.RESOURCE_REPLACEMENT_REQUEST + ) + replacement.replace(resource) + updated = self.backend.update_resource(resource_type, replacement) + updated = self.publish(request, updated) return self.make_response( updated.model_dump( scim_ctx=Context.RESOURCE_REPLACEMENT_RESPONSE, @@ -217,6 +256,7 @@ def call_single_resource( ) ) case _: # "PATCH" + self.ensure_supported(self.config.patch, "PATCH") payload = request.json # MS Entra sometimes passes a "id" attribute if "id" in payload: @@ -227,25 +267,24 @@ def call_single_resource( # MS Entra sometimes passes a "name" attribute del operation["name"] - ResourceModel = self.backend.get_model(resource_type.id) + ResourceModel = self.get_model(resource_type) patch_operation = PatchOp[ResourceModel].model_validate(payload) response_parameters = self.get_response_parameters( request, ResourceModel ) - resource = self.backend.get_resource(resource_type.id, resource_id) + resource = self.backend.get_resource(resource_type, resource_id) if resource is None: raise NotFound - if not self.continue_etag(request, resource): - raise PreconditionFailed + self.check_preconditions(request, resource) self.apply_patch_operation(resource, patch_operation) - updated = self.backend.update_resource(resource_type.id, resource) + updated = self.backend.update_resource(resource_type, resource) if ( response_parameters.attributes or response_parameters.excluded_attributes ): - self.adjust_location(request, updated) + updated = self.publish(request, updated) return self.make_response( updated.model_dump( scim_ctx=Context.RESOURCE_REPLACEMENT_RESPONSE, @@ -257,7 +296,9 @@ def call_single_resource( # A PATCH operation MAY return a 204 (no content) # if no attributes were requested return self.make_response( - None, 204, headers={"ETag": updated.meta.version} + None, + 204, + headers=self.etag_header(self.publish(request, updated)), ) @staticmethod @@ -291,33 +332,35 @@ def build_search_request( for key in SEARCH_REQUEST_PARAMETERS if key in request.args } + + # The filters of PATCH paths are part of the PATCH capability: the + # filter capability of RFC 7643 §5 refers to the search parameter of + # RFC 7644 §3.4.2.2 only. + parameters = {key.casefold() for key in payload} + if "filter" in parameters: + self.ensure_supported(self.config.filter, "Filtering") + if parameters & {"sortby", "sortorder"}: + self.ensure_supported(self.config.sort, "Sorting") + search_request = SearchRequest[Union[tuple(models)]].model_validate( # noqa: UP007 payload, scim_ctx=Context.SEARCH_REQUEST ) search_request.start_index = search_request.start_index or 1 - search_request.count = ( - self.page_size - if search_request.count is None - else min(search_request.count, self.page_size) - ) + max_results = self.config.filter.max_results if self.config.filter else None + if max_results is not None and ( + search_request.count is None or search_request.count > max_results + ): + search_request.count = max_results return search_request def query_resource(self, request: Request, resource: ResourceType | None): - models = ( - list(self.backend.get_models()) - if resource is None - else [self.backend.get_model(resource.id)] - ) + models = self.get_models() if resource is None else [self.get_model(resource)] search_request = self.build_search_request(request, models) - kwargs = {} - if resource is not None: - kwargs["resource_type_id"] = resource.id total_results, results = self.backend.query_resources( - search_request=search_request, **kwargs + search_request=search_request, resource_type=resource ) - for r in results: - self.adjust_location(request, r) + results = [self.publish(request, r) for r in results] resources = [ s.model_dump( @@ -327,7 +370,7 @@ def query_resource(self, request: Request, resource: ResourceType | None): for s in results ] - return ListResponse[Union[tuple(self.backend.get_models())]]( # noqa: UP007 + return ListResponse[Union[tuple(self.get_models())]]( # noqa: UP007 total_results=total_results, items_per_page=len(resources), start_index=search_request.start_index, @@ -337,9 +380,7 @@ def query_resource(self, request: Request, resource: ResourceType | None): def call_resource( self, request: Request, resource_endpoint: str, **kwargs ) -> Response: - resource_type = self.backend.get_resource_type_by_endpoint( - "/" + resource_endpoint - ) + resource_type = self.get_resource_type_by_endpoint(resource_endpoint) if not resource_type: raise NotFound @@ -352,14 +393,11 @@ def call_resource( ) case _: # "POST" payload = request.json - resource = self.backend.get_model(resource_type.id).model_validate( + resource = self.get_model(resource_type).model_validate( payload, scim_ctx=Context.RESOURCE_CREATION_REQUEST ) - created_resource = self.backend.create_resource( - resource_type.id, - resource, - ) - self.adjust_location(request, created_resource) + created_resource = self.backend.create_resource(resource_type, resource) + created_resource = self.publish(request, created_resource) return self.make_response( created_resource.model_dump( scim_ctx=Context.RESOURCE_CREATION_RESPONSE @@ -378,9 +416,7 @@ def call_query_all(self, request: Request, **kwargs) -> Response: def call_resource_search( self, request: Request, resource_endpoint: str, **kwargs ) -> Response: - resource_type = self.backend.get_resource_type_by_endpoint( - "/" + resource_endpoint - ) + resource_type = self.get_resource_type_by_endpoint(resource_endpoint) if not resource_type: raise NotFound return self.make_response( @@ -389,6 +425,20 @@ def call_resource_search( ) ) + @staticmethod + def ensure_supported(capability: Patch | Filter | Sort | None, operation: str): + """Refuse with a 501 an operation the configuration does not declare supported. + + RFC 7644 §3.12 answers 501 when the service provider does not support + the request operation. + """ + if capability is None or not capability.supported: + raise WerkzeugNotImplemented(f"{operation} is not supported") + + def call_bulk(self, request: Request, **kwargs): + """Implement the /Bulk endpoint, which this server does not support.""" + raise WerkzeugNotImplemented("Bulk operations are not supported") + def call_me(self, request: Request, **kwargs): """Implement the /Me endpoint. @@ -397,12 +447,6 @@ def call_me(self, request: Request, **kwargs): """ raise WerkzeugNotImplemented - def register_schema(self, schema: Schema): - self.backend.register_schema(schema) - - def register_resource_type(self, resource_type: ResourceType): - self.backend.register_resource_type(resource_type) - def register_bearer_token(self, token: str): """Register a static bearer token for authentication. @@ -449,89 +493,65 @@ def forbid_filter(request: Request): if "filter" in request.args: raise Forbidden - def get_service_provider_config(self): - """Build a ServiceProviderConfig object describing the server configuration.""" - auth_scheme = ( - [] - if not self.bearer_tokens - else [ - AuthenticationScheme( - type="oauthbearertoken", - name="bearer_token", - description="HTTP Bearer Token", - spec_uri="https://datatracker.ietf.org/doc/html/rfc6750", - ) - ] - ) - return ServiceProviderConfig( - documentation_uri="https://www.example.com/", - patch=Patch(supported=True), - bulk=Bulk(supported=False), - filter=Filter(supported=True, max_results=1000), - change_password=ChangePassword(supported=True), - sort=Sort(supported=True), - etag=ETag(supported=True), - authentication_schemes=auth_scheme, - meta=Meta( - resource_type="ServiceProviderConfig", - ), - ) - def call_service_provider_config(self, request: Request, **kwargs): """Return the ServiceProviderConfig.""" self.forbid_filter(request) - spc = self.get_service_provider_config() - spc.meta.location = request.url - return self.make_response(spc.model_dump()) + return self.make_response( + self.locate(self.config, request.base_url).model_dump() + ) + + @staticmethod + def locate(resource: ResourceType | Schema | ServiceProviderConfig, location: str): + """Return a copy of a discovery resource carrying its meta.""" + meta = Meta(resource_type=type(resource).__name__, location=location) + return resource.model_copy(update={"meta": meta}) def call_resource_type(self, request: Request, resource_type: str, **kwargs): """Return a single resource type.""" self.forbid_filter(request) - if res := self.backend.get_resource_type(resource_type): - cp = res.model_copy(deep=True) - cp.meta.location = request.url - return self.make_response(cp.model_dump()) + for res in self.provider.resource_types: + if res.id == resource_type: + return self.make_response( + self.locate(res, request.base_url).model_dump() + ) raise NotFound def call_schema(self, request: Request, schema_id: str): """Return a single schema.""" self.forbid_filter(request) - if res := self.backend.get_schema(schema_id): - cp = res.model_copy(deep=True) - cp.meta.location = request.url - return self.make_response(cp.model_dump()) + for res in self.provider.schemas: + if res.id == schema_id: + return self.make_response( + self.locate(res, request.base_url).model_dump() + ) raise NotFound def call_resource_types(self, request: Request, **kwargs): """Return a ListResponse of all known resource types.""" self.forbid_filter(request) - results = self.backend.get_resource_types() + results = self.provider.resource_types resp = ListResponse[ResourceType]( total_results=len(results), items_per_page=len(results), start_index=1, - resources=[self.adjust_location(request, s, True) for s in results], + resources=[self.locate(s, f"{request.base_url}/{s.id}") for s in results], ).model_dump() return self.make_response(resp) def call_schemas(self, request: Request, **kwargs): """Return a ListResponse of all known schemas.""" self.forbid_filter(request) - results = self.backend.get_schemas() + results = self.provider.schemas resp = ListResponse[Schema]( total_results=len(results), items_per_page=len(results), start_index=1, - resources=[self.adjust_location(request, s, True) for s in results], + resources=[self.locate(s, f"{request.base_url}/{s.id}") for s in results], ).model_dump() return self.make_response(resp) def wsgi_app(self, request: Request, environ): try: - if environ.get("PATH_INFO", "").endswith(".scim"): - # RFC 7644, Section 3.8 - # Just strip .scim suffix, the provider always returns application/scim+json - environ["PATH_INFO"], _, _ = environ["PATH_INFO"].rpartition(".scim") urls = self.url_map.bind_to_environ(environ) endpoint, args = urls.match() @@ -558,11 +578,14 @@ def wsgi_app(self, request: Request, environ): return self.make_error(Error.from_validation_errors(e)[0]) except Exception as e: self.log.exception(e) - tb = traceback.format_exc() - return self.make_error(Error(status=500, detail=str(e) + "\n" + tb)) + return self.make_error(Error(status=500, detail="Internal server error")) def __call__(self, environ, start_response): """Return the actual WSGI server implementation.""" + if environ.get("PATH_INFO", "").endswith(".scim"): + # RFC 7644, Section 3.8 + # Just strip .scim suffix, the provider always returns application/scim+json + environ["PATH_INFO"], _, _ = environ["PATH_INFO"].rpartition(".scim") request = Request(environ) response = self.wsgi_app(request, environ) if "Location" not in response.headers: diff --git a/scim2_server/resources/default-resource-types.json b/scim2_server/resources/default-resource-types.json index 8bf96f6..a567035 100644 --- a/scim2_server/resources/default-resource-types.json +++ b/scim2_server/resources/default-resource-types.json @@ -9,13 +9,9 @@ { "schema": "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User", - "required": true + "required": false } - ], - "meta": { - "location": "/v2/ResourceTypes/User", - "resourceType": "ResourceType" - } + ] }, { "schemas": ["urn:ietf:params:scim:schemas:core:2.0:ResourceType"], @@ -23,9 +19,5 @@ "name": "Group", "endpoint": "/Groups", "description": "Group", - "schema": "urn:ietf:params:scim:schemas:core:2.0:Group", - "meta": { - "location": "/v2/ResourceTypes/Group", - "resourceType": "ResourceType" - } + "schema": "urn:ietf:params:scim:schemas:core:2.0:Group" }] diff --git a/scim2_server/resources/default-schemas.json b/scim2_server/resources/default-schemas.json index 238a8b0..0cd0453 100644 --- a/scim2_server/resources/default-schemas.json +++ b/scim2_server/resources/default-schemas.json @@ -678,7 +678,6 @@ "description": "A label indicating the attribute's function.", "required": false, "caseExact": false, - "canonicalValues": [], "mutability": "readWrite", "returned": "default", "uniqueness": "none" @@ -733,7 +732,6 @@ "description": "A label indicating the attribute's function.", "required": false, "caseExact": false, - "canonicalValues": [], "mutability": "readWrite", "returned": "default", "uniqueness": "none" @@ -751,11 +749,7 @@ "mutability": "readWrite", "returned": "default" } - ], - "meta": { - "resourceType": "Schema", - "location": "/v2/Schemas/urn:ietf:params:scim:schemas:core:2.0:User" - } + ] }, { "id": "urn:ietf:params:scim:schemas:core:2.0:Group", @@ -830,11 +824,7 @@ "mutability": "readWrite", "returned": "default" } - ], - "meta": { - "resourceType": "Schema", - "location": "/v2/Schemas/urn:ietf:params:scim:schemas:core:2.0:Group" - } + ] }, { "id": "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User", @@ -941,10 +931,6 @@ "mutability": "readWrite", "returned": "default" } - ], - "meta": { - "resourceType": "Schema", - "location": "/v2/Schemas/urn:ietf:params:scim:schemas:extension:enterprise:2.0:User" - } + ] } ] diff --git a/scim2_server/resources/default-service-provider-config.json b/scim2_server/resources/default-service-provider-config.json new file mode 100644 index 0000000..5a7003e --- /dev/null +++ b/scim2_server/resources/default-service-provider-config.json @@ -0,0 +1,25 @@ +{ + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:ServiceProviderConfig" + ], + "patch": { + "supported": true + }, + "bulk": { + "supported": false + }, + "filter": { + "supported": true, + "maxResults": 1000 + }, + "changePassword": { + "supported": true + }, + "sort": { + "supported": true + }, + "etag": { + "supported": true + }, + "authenticationSchemes": [] +} diff --git a/scim2_server/utils.py b/scim2_server/utils.py index c4e2013..9be52b8 100644 --- a/scim2_server/utils.py +++ b/scim2_server/utils.py @@ -8,7 +8,6 @@ from pydantic import EmailStr from pydantic import ValidationError from scim2_models import BaseModel -from scim2_models import Extension from scim2_models import InvalidValueException from scim2_models import Mutability from scim2_models import MutabilityException @@ -16,9 +15,11 @@ from scim2_models import Resource from scim2_models import ResourceType from scim2_models import Schema +from scim2_models import ScimProvider +from scim2_models import ServiceProviderConfig -def load_json_resource(json_name: str) -> list: +def load_json_resource(json_name: str) -> Any: """Load a JSON document from the scim2_server package resources.""" fp = importlib.resources.files("scim2_server") / "resources" / json_name with open(fp) as f: @@ -45,27 +46,20 @@ def load_default_resource_types() -> dict[str, ResourceType]: return load_scim_resource("default-resource-types.json", ResourceType) -def merge_resources(target: Resource, updates: BaseModel): - """Merge a resource with another resource as specified for HTTP PUT (RFC 7644, section 3.5.1).""" - for set_attribute in updates.model_fields_set: - mutability = target.get_field_annotation(set_attribute, Mutability) - if mutability == Mutability.read_only: - continue - if isinstance(getattr(updates, set_attribute), Extension): - # This is a model extension, handle it as its own resource - # and don't simply overwrite it - target_extension = getattr(target, set_attribute) - if target_extension is None: - setattr(target, set_attribute, getattr(updates, set_attribute)) - else: - merge_resources(target_extension, getattr(updates, set_attribute)) - continue - new_value = getattr(updates, set_attribute) - if mutability == Mutability.immutable and getattr( - target, set_attribute - ) not in (None, new_value): - raise MutabilityException() - setattr(target, set_attribute, new_value) +def load_default_service_provider_config() -> ServiceProviderConfig: + """Load the default service provider configuration.""" + return ServiceProviderConfig.model_validate( + load_json_resource("default-service-provider-config.json") + ) + + +def load_default_provider() -> ScimProvider: + """Describe a service serving the default schemas, resource types and configuration.""" + return ScimProvider.from_discovery( + load_default_schemas().values(), + load_default_resource_types().values(), + config=load_default_service_provider_config(), + ) def get_by_alias( diff --git a/tests/conftest.py b/tests/conftest.py index bae4381..f3d9168 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,13 +3,25 @@ import httpx2 import pytest +from scim2_models import ScimProvider from scim2_server.backend import InMemoryBackend -from scim2_server.provider import SCIMProvider +from scim2_server.provider import SCIMApplication +from scim2_server.utils import load_default_provider from scim2_server.utils import load_default_resource_types from scim2_server.utils import load_default_schemas +@pytest.fixture(scope="session") +def scim_provider(): + return load_default_provider() + + +@pytest.fixture(scope="session") +def user_type(scim_provider): + return next(rt for rt in scim_provider.resource_types if rt.id == "User") + + @pytest.fixture def backend(): return InMemoryBackend() @@ -32,25 +44,41 @@ def fake_user_data(): @pytest.fixture -def provider(backend, static_data): - provider = SCIMProvider(backend) - for schema in static_data[0].values(): - provider.register_schema(schema) - for resource_type in static_data[1].values(): - provider.register_resource_type(resource_type) - return provider +def app(backend, scim_provider): + return SCIMApplication(backend, scim_provider) @pytest.fixture -def wsgi(provider): - transport = httpx2.WSGITransport(app=provider) +def wsgi(app): + transport = httpx2.WSGITransport(app=app) client = httpx2.Client(transport=transport, base_url="https://scim.example.com") client.__enter__() yield client - provider.backend.resources = [] + app.backend.resources = [] client.__exit__(None, None, None) +@pytest.fixture +def wsgi_with(backend, scim_provider): + """Build clients of applications serving the default resources under another configuration.""" + clients = [] + + def build(config): + provider = ScimProvider( + models=scim_provider.models, + resource_types=scim_provider.resource_types, + config=config, + ) + transport = httpx2.WSGITransport(app=SCIMApplication(backend, provider)) + client = httpx2.Client(transport=transport, base_url="https://scim.example.com") + clients.append(client) + return client + + yield build + for client in clients: + client.close() + + @pytest.fixture def first_fake_user(wsgi, fake_user_data): r = wsgi.post("/v2/Users", json=fake_user_data[0]) diff --git a/tests/integration/test_basic.py b/tests/integration/test_basic.py index eb97911..b21815a 100644 --- a/tests/integration/test_basic.py +++ b/tests/integration/test_basic.py @@ -7,7 +7,7 @@ from scim2_models import User -class TestSCIMProviderBasic: +class TestSCIMApplicationBasic: def test_user_creation(self, wsgi): payload = { "schemas": [ @@ -95,8 +95,8 @@ def test_unique_constraints(self, wsgi): assert r.status_code == 200 assert r.json()["userName"] == "bjensen2@example.com" - def test_sort(self, provider, wsgi): - TypedListResponse = ListResponse[Union[tuple(provider.backend.get_models())]] # noqa: UP007 + def test_sort(self, app, wsgi): + TypedListResponse = ListResponse[Union[tuple(app.get_models())]] # noqa: UP007 def assert_sorted(sort_by: str, sorted: list[str], endpoint: str = "/v2/Users"): for order_by, inverted in ( @@ -189,3 +189,100 @@ def test_sort_by_value_filter_is_refused(self, wsgi, sort_by): result = wsgi.get("/v2/Users", params={"sortBy": sort_by}) assert result.status_code == 400 assert result.json()["scimType"] == "invalidPath" + + def test_sort_on_the_root_puts_undeclared_attributes_last(self, wsgi): + """RFC 7644 §3.4.2.1 treats an attribute a resource type lacks as having no value.""" + user_id = wsgi.post( + "/v2/Users", + json={ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": "bjensen", + }, + ).json()["id"] + group_id = wsgi.post( + "/v2/Groups", + json={ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:Group"], + "displayName": "admins", + }, + ).json()["id"] + + r = wsgi.get("/v2/", params={"sortBy": "userName"}) + assert [resource["id"] for resource in r.json()["Resources"]] == [ + user_id, + group_id, + ] + r = wsgi.get("/v2/", params={"sortBy": "userName", "sortOrder": "descending"}) + assert [resource["id"] for resource in r.json()["Resources"]] == [ + group_id, + user_id, + ] + + @pytest.mark.parametrize( + ("sort_by", "first"), + [("externalId", "upper"), ("userName", "lower")], + ) + def test_sort_follows_the_case_exactness_of_the_attribute( + self, wsgi, sort_by, first + ): + """A case-exact attribute sorts upper case first, another one ignores the case.""" + ids = {} + for name, value in (("lower", "a"), ("upper", "B")): + ids[name] = wsgi.post( + "/v2/Users", + json={ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": value, + "externalId": value, + }, + ).json()["id"] + + r = wsgi.get("/v2/Users", params={"sortBy": sort_by}) + assert r.json()["Resources"][0]["id"] == ids[first] + + def test_sort_on_an_extension_attribute(self, wsgi): + """An attribute qualified by an extension URN is read from the extension.""" + enterprise = "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User" + ids = [ + wsgi.post( + "/v2/Users", + json={ + "schemas": [ + "urn:ietf:params:scim:schemas:core:2.0:User", + enterprise, + ], + "userName": user_name, + enterprise: {"employeeNumber": employee_number}, + }, + ).json()["id"] + for user_name, employee_number in (("a", "2"), ("b", "1")) + ] + + r = wsgi.get("/v2/Users", params={"sortBy": f"{enterprise}:employeeNumber"}) + assert [resource["id"] for resource in r.json()["Resources"]] == ids[::-1] + + def test_sort_reads_a_sub_attribute_from_the_primary_entry(self, wsgi): + """A sub-attribute of a multi-valued attribute is read from its primary entry.""" + ids = [ + wsgi.post( + "/v2/Users", + json={ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": user_name, + "emails": emails, + }, + ).json()["id"] + for user_name, emails in ( + ( + "a", + [ + {"value": "a@example.com", "type": "other"}, + {"value": "b@example.com", "type": "work", "primary": True}, + ], + ), + ("b", [{"value": "c@example.com", "type": "home"}]), + ) + ] + + r = wsgi.get("/v2/Users", params={"sortBy": "emails.type"}) + assert [resource["id"] for resource in r.json()["Resources"]] == ids[::-1] diff --git a/tests/integration/test_etags.py b/tests/integration/test_etags.py index dd2df90..b448752 100644 --- a/tests/integration/test_etags.py +++ b/tests/integration/test_etags.py @@ -1,4 +1,23 @@ -class TestSCIMProviderETags: +import pytest +from scim2_models import ETag + +from scim2_server.utils import load_default_service_provider_config + + +@pytest.fixture(params=[None, ETag(supported=False)], ids=["undeclared", "unsupported"]) +def unversioned(request, wsgi_with): + """Build a client of a service that does not support ETags.""" + config = load_default_service_provider_config() + config.etag = request.param + return wsgi_with(config) + + +@pytest.fixture +def unversioned_user(unversioned, fake_user_data): + return unversioned.post("/v2/Users", json=fake_user_data[0]).json()["id"] + + +class TestSCIMApplicationETags: def test_resource_get_etag_match(self, wsgi, first_fake_user): r = wsgi.get( f"/v2/Users/{first_fake_user}", @@ -24,7 +43,7 @@ def test_resource_get_etag_match(self, wsgi, first_fake_user): assert r.status_code == 200 r = wsgi.get(f"/v2/Users/{first_fake_user}", headers={"If-Match": 'W/"abc"'}) - assert r.status_code == 304 + assert r.status_code == 412 r = wsgi.get(f"/v2/Users/{first_fake_user}", headers={"If-Match": "*"}) assert r.status_code == 200 @@ -141,3 +160,107 @@ def test_resource_patch_etag_match(self, wsgi, first_fake_user): f"/v2/Users/{first_fake_user}", ) assert r.json()["userName"] == "Foo" + + def test_resource_get_not_modified_carries_the_etag(self, wsgi, first_fake_user): + """RFC 7232 §4.1: a 304 carries the ETag a 200 would have carried.""" + version = wsgi.get(f"/v2/Users/{first_fake_user}").headers["etag"] + r = wsgi.get(f"/v2/Users/{first_fake_user}", headers={"If-None-Match": version}) + assert r.status_code == 304 + assert r.headers["etag"] == version + + def test_resource_get_failed_if_match_wins_over_if_none_match( + self, wsgi, first_fake_user + ): + """RFC 7232 §6: If-Match is evaluated first, and its failure answers 412.""" + version = wsgi.get(f"/v2/Users/{first_fake_user}").headers["etag"] + r = wsgi.get( + f"/v2/Users/{first_fake_user}", + headers={"If-Match": 'W/"abc"', "If-None-Match": version}, + ) + assert r.status_code == 412 + + def test_resource_delete_etag_match(self, wsgi, first_fake_user): + """A DELETE only removes the version the client names.""" + version = wsgi.get(f"/v2/Users/{first_fake_user}").headers["etag"] + + r = wsgi.delete(f"/v2/Users/{first_fake_user}", headers={"If-Match": 'W/"abc"'}) + assert r.status_code == 412 + r = wsgi.delete(f"/v2/Users/{first_fake_user}", headers={"If-None-Match": "*"}) + assert r.status_code == 412 + assert wsgi.get(f"/v2/Users/{first_fake_user}").status_code == 200 + + r = wsgi.delete(f"/v2/Users/{first_fake_user}", headers={"If-Match": version}) + assert r.status_code == 204 + assert wsgi.get(f"/v2/Users/{first_fake_user}").status_code == 404 + + def test_resource_delete_unknown(self, wsgi): + """Deleting a resource that does not exist answers 404.""" + assert wsgi.delete("/v2/Users/unknown").status_code == 404 + + def test_resource_put_if_none_match(self, wsgi, first_fake_user): + """RFC 7232 §3.2: a failed If-None-Match on a PUT answers 412, not 304.""" + version = wsgi.get(f"/v2/Users/{first_fake_user}").headers["etag"] + r = wsgi.put( + f"/v2/Users/{first_fake_user}", + json={"userName": "Foo"}, + headers={"If-None-Match": version}, + ) + assert r.status_code == 412 + + +class TestSCIMApplicationWithoutETags: + def test_responses_carry_no_version(self, unversioned, fake_user_data): + """RFC 7644 §3.14: a service without ETags sends neither the header nor meta.version.""" + r = unversioned.post("/v2/Users", json=fake_user_data[0]) + assert "etag" not in r.headers + assert "version" not in r.json()["meta"] + user_id = r.json()["id"] + + r = unversioned.get(f"/v2/Users/{user_id}") + assert "etag" not in r.headers + assert "version" not in r.json()["meta"] + + r = unversioned.get("/v2/Users") + assert "version" not in r.json()["Resources"][0]["meta"] + + r = unversioned.put(f"/v2/Users/{user_id}", json={"userName": "Foo"}) + assert "etag" not in r.headers + assert "version" not in r.json()["meta"] + + r = unversioned.patch( + f"/v2/Users/{user_id}", + json={ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [{"op": "replace", "path": "userName", "value": "Bar"}], + }, + ) + assert r.status_code == 204 + assert "etag" not in r.headers + + @pytest.mark.parametrize("method", ["GET", "PUT", "DELETE"]) + def test_if_match_listing_tags_fails(self, unversioned, unversioned_user, method): + """RFC 7232 §3.1: no listed tag matches a resource that has none.""" + r = unversioned.request( + method, + f"/v2/Users/{unversioned_user}", + json={"userName": "Foo"} if method == "PUT" else None, + headers={"If-Match": 'W/"abc"'}, + ) + assert r.status_code == 412 + + def test_if_match_any_passes(self, unversioned, unversioned_user): + """RFC 7232 §3.1: "*" only requires the resource to exist.""" + r = unversioned.get(f"/v2/Users/{unversioned_user}", headers={"If-Match": "*"}) + assert r.status_code == 200 + + def test_if_none_match(self, unversioned, unversioned_user): + """A listed tag never matches, "*" does, and its 304 carries no ETag.""" + url = f"/v2/Users/{unversioned_user}" + assert ( + unversioned.get(url, headers={"If-None-Match": 'W/"abc"'}).status_code + == 200 + ) + + r = unversioned.get(url, headers={"If-None-Match": "*"}) + assert r.status_code == 304 + assert "etag" not in r.headers diff --git a/tests/integration/test_ms_entra.py b/tests/integration/test_ms_entra.py index 9e9d9a3..e2276ba 100644 --- a/tests/integration/test_ms_entra.py +++ b/tests/integration/test_ms_entra.py @@ -3,7 +3,7 @@ import pytest -class TestSCIMProviderMSEntraIntegration: +class TestSCIMApplicationMSEntraIntegration: """Tests based on Postman tests from Azure docs. https://github.com/AzureAD/SCIMReferenceCode/wiki/Test-Your-SCIM-Endpoint diff --git a/tests/integration/test_okta.py b/tests/integration/test_okta.py index cfade16..28c247e 100644 --- a/tests/integration/test_okta.py +++ b/tests/integration/test_okta.py @@ -1,7 +1,7 @@ import datetime -class TestSCIMProviderOktaIntegration: +class TestSCIMApplicationOktaIntegration: """Test based on Runscope Spec Test JSON from Okta docs. https://developer.okta.com/docs/guides/scim-provisioning-integration-prepare/main/#test-your-scim-api diff --git a/tests/integration/test_scim_provider.py b/tests/integration/test_scim_application.py similarity index 71% rename from tests/integration/test_scim_provider.py rename to tests/integration/test_scim_application.py index 84a23e4..0c08c26 100644 --- a/tests/integration/test_scim_provider.py +++ b/tests/integration/test_scim_application.py @@ -3,16 +3,27 @@ import httpx2 import pytest import time_machine +from scim2_models import Bulk +from scim2_models import ChangePassword +from scim2_models import ETag +from scim2_models import Filter +from scim2_models import Patch +from scim2_models import ScimProvider from scim2_models import SearchRequest +from scim2_models import ServiceProviderConfig +from scim2_models import Sort +from scim2_models import User +from scim2_server.provider import SCIMApplication +from scim2_server.utils import load_default_service_provider_config from tests.utils import compare_dicts -class TestSCIMProvider: - """End-to-end tests for the SCIMProvider.""" +class TestSCIMApplication: + """End-to-end tests for the SCIMApplication.""" - def test_location_mapping(self, provider): - transport = httpx2.WSGITransport(app=provider, script_name="/foo/bar") + def test_location_mapping(self, app): + transport = httpx2.WSGITransport(app=app, script_name="/foo/bar") with httpx2.Client( transport=transport, base_url="https://sub.testserver.company:1234" ) as client: @@ -42,7 +53,6 @@ def test_service_provider_configuration(self, wsgi): "supported": False, }, "changePassword": {"supported": True}, - "documentationUri": "https://www.example.com/", "etag": {"supported": True}, "filter": {"maxResults": 1000, "supported": True}, "meta": { @@ -65,7 +75,6 @@ def test_no_version_prefix(self, wsgi): "supported": False, }, "changePassword": {"supported": True}, - "documentationUri": "https://www.example.com/", "etag": {"supported": True}, "filter": {"maxResults": 1000, "supported": True}, "meta": { @@ -77,6 +86,36 @@ def test_no_version_prefix(self, wsgi): "sort": {"supported": True}, } + def test_service_provider_configuration_from_the_provider(self, backend): + """The configuration the provider carries is the one published.""" + config = ServiceProviderConfig( + patch=Patch(supported=False), + bulk=Bulk(supported=False), + filter=Filter(supported=False), + change_password=ChangePassword(supported=False), + sort=Sort(supported=False), + etag=ETag(supported=False), + ) + provider = ScimProvider(models=[User], config=config) + transport = httpx2.WSGITransport(app=SCIMApplication(backend, provider)) + with httpx2.Client( + transport=transport, base_url="https://scim.example.com" + ) as client: + published = client.get("/v2/ServiceProviderConfig").json() + assert published["patch"] == {"supported": False} + assert published["filter"] == {"supported": False} + + def test_service_provider_configuration_defaults(self, backend): + """A provider carrying no configuration is served the default one.""" + provider = ScimProvider(models=[User]) + transport = httpx2.WSGITransport(app=SCIMApplication(backend, provider)) + with httpx2.Client( + transport=transport, base_url="https://scim.example.com" + ) as client: + published = client.get("/v2/ServiceProviderConfig").json() + del published["meta"] + assert published == load_default_service_provider_config().model_dump() + def test_schemas(self, wsgi): r = wsgi.get("/v2/Schemas") assert r.status_code == 200 @@ -124,7 +163,12 @@ def test_resource_types(self, wsgi): j = r.json() assert j["schemas"] == ["urn:ietf:params:scim:schemas:core:2.0:ResourceType"] assert j["schema"] == "urn:ietf:params:scim:schemas:core:2.0:User" - assert len(j["schemaExtensions"]) == 1 + assert j["schemaExtensions"] == [ + { + "schema": "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User", + "required": False, + } + ] assert j["meta"]["location"] == "https://scim.example.com/v2/ResourceTypes/User" # RFC7644, Section 4 @@ -138,6 +182,29 @@ def test_resource_types(self, wsgi): assert j["status"] == "404" assert "not found" in j["detail"] + def test_discovery_resources_carry_their_meta(self, wsgi): + """Each schema and resource type is published with its type and its location.""" + base_url = "https://scim.example.com/v2" + for endpoint, resource_type in ( + ("Schemas", "Schema"), + ("ResourceTypes", "ResourceType"), + ): + for resource in wsgi.get(f"/v2/{endpoint}").json()["Resources"]: + location = f"{base_url}/{endpoint}/{resource['id']}" + assert resource["meta"] == { + "resourceType": resource_type, + "location": location, + } + single = wsgi.get(f"/v2/{endpoint}/{resource['id']}").json() + assert single["meta"] == resource["meta"] + + def test_discovery_location_leaves_out_the_scim_suffix(self, wsgi): + """RFC 7644 §3.8: the .scim suffix only selects the format, it is not part of the location.""" + r = wsgi.get("/v2/ResourceTypes/User.scim") + assert r.json()["meta"]["location"] == ( + "https://scim.example.com/v2/ResourceTypes/User" + ) + def test_me(self, wsgi): r = wsgi.get("/v2/Me") assert r.status_code == 501 @@ -319,10 +386,15 @@ def test_resource_put(self, wsgi, first_fake_user): ) assert r.status_code == 200 j = r.json() + # RFC 7644 §3.5.1 lets the omitted readWrite attributes be cleared. + assert "displayName" not in j + assert ( + "organization" + not in j["urn:ietf:params:scim:schemas:extension:enterprise:2.0:User"] + ) compare_dicts( { "id": first_fake_user, - "displayName": "Mx. Larry Hunt", "active": False, "userName": "foo@example.com", "name": { @@ -331,10 +403,8 @@ def test_resource_put(self, wsgi, first_fake_user): "phoneNumbers": [ {"value": "001-767-633-4744", "type": "home", "primary": True} ], - "addresses": [], "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User": { "employeeNumber": "512", - "organization": "Blake PLC", }, }, j, @@ -348,6 +418,7 @@ def test_resource_put_remove(self, wsgi, first_fake_user): r = wsgi.put( f"/v2/Users/{first_fake_user}", json={ + "userName": "joseph96@williams-brown.com", "name": None, "phoneNumbers": [], }, @@ -357,6 +428,80 @@ def test_resource_put_remove(self, wsgi, first_fake_user): assert not j.get("name") assert not j.get("phoneNumbers") + def test_resource_put_requires_the_required_attributes(self, wsgi, first_fake_user): + """RFC 7644 §3.5.1: a required attribute MUST be specified in a PUT.""" + r = wsgi.put(f"/v2/Users/{first_fake_user}", json={"displayName": "Foo"}) + assert r.status_code == 400 + assert r.json()["scimType"] == "invalidValue" + + def test_resource_put_clears_an_omitted_extension(self, wsgi, first_fake_user): + """An extension left out of the replacement is cleared like any readWrite attribute.""" + r = wsgi.put( + f"/v2/Users/{first_fake_user}", + json={ + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], + "userName": "joseph96@williams-brown.com", + }, + ) + assert r.status_code == 200 + assert ( + "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User" not in r.json() + ) + + def test_resource_put_keeps_an_omitted_password( + self, app, user_type, wsgi, first_fake_user + ): + """A client never gets the password back, so omitting it does not clear it.""" + stored = app.backend.get_resource(user_type, first_fake_user) + assert stored.password is not None + + r = wsgi.put( + f"/v2/Users/{first_fake_user}", + json={"userName": "joseph96@williams-brown.com"}, + ) + assert r.status_code == 200 + replaced = app.backend.get_resource(user_type, first_fake_user) + assert replaced.password == stored.password + + def test_resource_put_clears_a_password_set_to_null( + self, app, user_type, wsgi, first_fake_user + ): + """An explicit null is how RFC 7644 §3.5.1 lets a client clear a value.""" + r = wsgi.put( + f"/v2/Users/{first_fake_user}", + json={"userName": "joseph96@williams-brown.com", "password": None}, + ) + assert r.status_code == 200 + assert app.backend.get_resource(user_type, first_fake_user).password is None + + def test_resource_put_refuses_to_change_an_immutable_attribute(self, wsgi): + """RFC 7644 §3.5.1: an immutable value already set MUST match the input value.""" + group = { + "schemas": ["urn:ietf:params:scim:schemas:core:2.0:Group"], + "displayName": "admins", + "members": [{"value": "u1", "type": "User"}], + } + group_id = wsgi.post("/v2/Groups", json=group).json()["id"] + + group["members"][0]["type"] = "Group" + r = wsgi.put(f"/v2/Groups/{group_id}", json=group) + assert r.status_code == 400 + assert r.json()["scimType"] == "mutability" + + def test_resource_put_ignores_read_only_attributes(self, wsgi, first_fake_user): + """RFC 7644 §3.5.1: readOnly values provided SHALL be ignored.""" + r = wsgi.put( + f"/v2/Users/{first_fake_user}", + json={ + "userName": "joseph96@williams-brown.com", + "id": "another-id", + "meta": {"resourceType": "Group"}, + }, + ) + assert r.status_code == 200 + assert r.json()["id"] == first_fake_user + assert r.json()["meta"]["resourceType"] == "User" + def test_resource_delete(self, wsgi, first_fake_user): r = wsgi.delete(f"/v2/Users/{first_fake_user}") assert r.status_code == 204 @@ -603,11 +748,13 @@ def test_search_items_per_page_counts_the_returned_page(self, wsgi, fake_user_da assert len(r.json()["Resources"]) == 1 @pytest.mark.parametrize("payload", [{}, {"count": 10}]) - def test_search_post_is_capped_to_page_size( - self, provider, wsgi, fake_user_data, payload + def test_search_post_is_capped_to_max_results( + self, wsgi_with, fake_user_data, payload ): - """A POST search is paginated with the server page size as a GET is.""" - provider.page_size = 2 + """A page never holds more than the filter.maxResults the service declares.""" + config = load_default_service_provider_config() + config.filter.max_results = 2 + wsgi = wsgi_with(config) for user in fake_user_data[:3]: wsgi.post("/v2/Users", json=user) r = wsgi.post( @@ -623,6 +770,112 @@ def test_search_post_is_capped_to_page_size( assert r.json()["startIndex"] == 1 assert len(r.json()["Resources"]) == 2 + def test_search_without_max_results_is_not_paginated( + self, wsgi_with, fake_user_data + ): + """A service declaring no filter.maxResults returns every result when no count is given.""" + config = load_default_service_provider_config() + config.filter = Filter(supported=False) + wsgi = wsgi_with(config) + for user in fake_user_data[:3]: + wsgi.post("/v2/Users", json=user) + r = wsgi.get("/v2/Users") + assert r.json()["itemsPerPage"] == 3 + + def test_search_without_filter_capabilities_is_not_paginated( + self, wsgi_with, fake_user_data + ): + """A service declaring no filter capabilities returns every result when no count is given.""" + config = load_default_service_provider_config() + config.filter = None + wsgi = wsgi_with(config) + for user in fake_user_data[:3]: + wsgi.post("/v2/Users", json=user) + r = wsgi.get("/v2/Users") + assert r.json()["itemsPerPage"] == 3 + + def test_search_count_zero_returns_no_resource(self, wsgi, fake_user_data): + """RFC 7644 §3.4.2.4: a count of 0 only asks for the total results.""" + for user in fake_user_data[:3]: + wsgi.post("/v2/Users", json=user) + r = wsgi.get("/v2/Users", params={"count": 0}) + assert r.json()["totalResults"] == 3 + assert r.json()["itemsPerPage"] == 0 + + @pytest.mark.parametrize( + "capability,method,url,payload", + [ + ("filter", "GET", '/v2/Users?filter=userName eq "bjensen"', None), + ( + "filter", + "POST", + "/v2/Users/.search", + {"filter": 'userName eq "bjensen"'}, + ), + ("sort", "GET", "/v2/Users?sortBy=userName", None), + ("sort", "GET", "/v2/Users?sortOrder=descending", None), + ("sort", "POST", "/v2/.search", {"sortBy": "userName"}), + ], + ) + def test_search_refuses_an_unsupported_capability( + self, wsgi_with, capability, method, url, payload + ): + """RFC 7644 §3.12: a search using a capability the service does not declare answers 501.""" + config = load_default_service_provider_config() + setattr(config, capability, None) + wsgi = wsgi_with(config) + if payload is not None: + payload["schemas"] = ["urn:ietf:params:scim:api:messages:2.0:SearchRequest"] + r = wsgi.request(method, url, json=payload) + assert r.status_code == 501 + assert r.json()["schemas"] == ["urn:ietf:params:scim:api:messages:2.0:Error"] + assert r.json()["status"] == "501" + + @pytest.mark.parametrize("patch", [None, Patch(supported=False)]) + def test_patch_refuses_when_unsupported(self, wsgi_with, fake_user_data, patch): + """A service declaring no PATCH support answers 501 to a PATCH, before reading its body.""" + config = load_default_service_provider_config() + config.patch = patch + wsgi = wsgi_with(config) + user_id = wsgi.post("/v2/Users", json=fake_user_data[0]).json()["id"] + r = wsgi.patch(f"/v2/Users/{user_id}", json={}) + assert r.status_code == 501 + + def test_patch_path_filters_do_not_require_the_filter_capability( + self, wsgi_with, fake_user_data + ): + """The filter of a PATCH path belongs to PATCH, not to the search filter capability.""" + config = load_default_service_provider_config() + config.filter = Filter(supported=False) + wsgi = wsgi_with(config) + user_id = wsgi.post("/v2/Users", json=fake_user_data[0]).json()["id"] + r = wsgi.patch( + f"/v2/Users/{user_id}", + json={ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:PatchOp"], + "Operations": [ + { + "op": "replace", + "path": 'emails[type eq "work"].value', + "value": "new@example.com", + } + ], + }, + ) + assert r.status_code == 204 + + def test_bulk_is_not_implemented(self, wsgi): + """RFC 7644 §3.7: bulk is optional, and this server answers 501 to it.""" + r = wsgi.post( + "/v2/Bulk", + json={ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:BulkRequest"], + "Operations": [], + }, + ) + assert r.status_code == 501 + assert r.json()["detail"] == "Bulk operations are not supported" + def test_validation_error_carries_scim_type(self, wsgi): """A payload refused by validation answers with a SCIM error keyword.""" r = wsgi.post( @@ -770,10 +1023,10 @@ def update_user(user_id): assert j["meta"]["created"] == "2024-03-14T06:00:00Z" assert j["meta"]["lastModified"] == "2024-03-16T08:30:00Z" - def test_authentication(self, first_fake_user, provider, wsgi): + def test_authentication(self, first_fake_user, app, wsgi): r = wsgi.get("/v2/ServiceProviderConfig") assert "WWW-Authenticate" not in r.headers - provider.register_bearer_token("SuperSecretToken") + app.register_bearer_token("SuperSecretToken") r = wsgi.get("/v2/ServiceProviderConfig") assert "WWW-Authenticate" in r.headers diff --git a/tests/integration/test_shared_schema.py b/tests/integration/test_shared_schema.py new file mode 100644 index 0000000..b3b9a7e --- /dev/null +++ b/tests/integration/test_shared_schema.py @@ -0,0 +1,95 @@ +import httpx2 +import pytest +from scim2_models import ResourceType +from scim2_models import ScimProvider + +from scim2_server.backend import InMemoryBackend +from scim2_server.provider import SCIMApplication +from scim2_server.utils import load_default_resource_types +from scim2_server.utils import load_default_schemas +from scim2_server.utils import load_default_service_provider_config + +USER_SCHEMA = "urn:ietf:params:scim:schemas:core:2.0:User" + + +@pytest.fixture +def wsgi(): + """Build a client of a service serving users and admins with the same schema.""" + admin = ResourceType( + id="Admin", name="Admin", endpoint="/Admins", schema=USER_SCHEMA + ) + provider = ScimProvider.from_discovery( + load_default_schemas().values(), + [*load_default_resource_types().values(), admin], + config=load_default_service_provider_config(), + ) + transport = httpx2.WSGITransport(app=SCIMApplication(InMemoryBackend(), provider)) + with httpx2.Client( + transport=transport, base_url="https://scim.example.com" + ) as client: + yield client + + +def test_resource_types_sharing_a_schema_serve_disjoint_resources(wsgi): + """A resource belongs to the resource type it was created through only.""" + user_id = wsgi.post("/v2/Users", json={"userName": "bjensen"}).json()["id"] + admin = wsgi.post("/v2/Admins", json={"userName": "root"}).json() + assert admin["meta"]["resourceType"] == "Admin" + assert ( + admin["meta"]["location"] == f"https://scim.example.com/v2/Admins/{admin['id']}" + ) + + assert wsgi.get(f"/v2/Users/{admin['id']}").status_code == 404 + assert wsgi.get(f"/v2/Admins/{user_id}").status_code == 404 + assert wsgi.delete(f"/v2/Users/{admin['id']}").status_code == 404 + assert [r["id"] for r in wsgi.get("/v2/Users").json()["Resources"]] == [user_id] + assert [r["id"] for r in wsgi.get("/v2/Admins").json()["Resources"]] == [ + admin["id"] + ] + + +def test_a_root_search_spans_the_resource_types_sharing_a_schema(wsgi): + """RFC 7644 §3.4.2.1: a search on the root returns the resources of every type, each with its resourceType.""" + wsgi.post("/v2/Users", json={"userName": "bjensen"}) + wsgi.post("/v2/Admins", json={"userName": "root"}) + resources = wsgi.get("/v2/").json()["Resources"] + assert sorted(r["meta"]["resourceType"] for r in resources) == ["Admin", "User"] + + r = wsgi.post( + "/v2/.search", + json={ + "schemas": ["urn:ietf:params:scim:api:messages:2.0:SearchRequest"], + "filter": 'userName eq "root"', + }, + ) + assert [r["meta"]["resourceType"] for r in r.json()["Resources"]] == ["Admin"] + + +def test_a_shared_schema_is_published_once(wsgi): + """The schema two resource types share appears once on /Schemas.""" + schemas = wsgi.get("/v2/Schemas").json()["Resources"] + assert [s["id"] for s in schemas].count(USER_SCHEMA) == 1 + resource_types = wsgi.get("/v2/ResourceTypes").json()["Resources"] + assert {rt["id"] for rt in resource_types} == {"User", "Group", "Admin"} + + +@pytest.mark.parametrize("user_name", ["bjensen", "BJensen"]) +def test_uniqueness_spans_the_resource_types_sharing_a_schema(wsgi, user_name): + """RFC 7643 erratum 8279: uniqueness holds among the resources using the same schema.""" + wsgi.post("/v2/Users", json={"userName": "bjensen"}) + r = wsgi.post("/v2/Admins", json={"userName": user_name}) + assert r.status_code == 409 + assert r.json()["scimType"] == "uniqueness" + + +def test_a_group_holds_members_of_another_resource_type(wsgi): + """The type and $ref of a member may name any resource type.""" + admin_id = wsgi.post("/v2/Admins", json={"userName": "root"}).json()["id"] + member = { + "value": admin_id, + "$ref": f"https://scim.example.com/v2/Admins/{admin_id}", + "type": "Admin", + } + r = wsgi.post("/v2/Groups", json={"displayName": "Operators", "members": [member]}) + assert r.status_code == 201 + assert r.json()["members"] == [member] diff --git a/tests/test_backend.py b/tests/test_backend.py index e000329..6a48c2c 100644 --- a/tests/test_backend.py +++ b/tests/test_backend.py @@ -1,35 +1,25 @@ +from typing import Annotated + import pytest +from scim2_models import URN from scim2_models import Attribute from scim2_models import CaseExact from scim2_models import Extension +from scim2_models import Resource from scim2_models import ResourceType from scim2_models import Schema -from scim2_models import SchemaExtension +from scim2_models import ScimProvider from scim2_models import SearchRequest from scim2_models import Uniqueness +from scim2_models import UniquenessException from scim2_models import User from scim2_server.backend import InMemoryBackend class TestBackend: - def test_unique_attributes(self, provider): - backend = provider.backend - assert "Group" in backend.unique_attributes - assert "User" in backend.unique_attributes - assert len(backend.unique_attributes["User"]) == 1 - assert backend.unique_attributes["User"][ - 0 - ] == InMemoryBackend.UniquenessDescriptor( - schema=None, attribute_name="userName", case_exact=False - ) - - rt = ResourceType( - schema="urn:example:2.0:Foo", - schema_extensions=[ - SchemaExtension(schema="urn:example:2.0:Bar", required=True) - ], - ) + def test_unique_attributes(self, app): + """The uniqueness constraints are read from the annotations of the model, extensions included.""" foo_schema = Schema( id="urn:example:2.0:Foo", name="Foo", @@ -50,105 +40,176 @@ def test_unique_attributes(self, provider): name="a", type=Attribute.Type.string, uniqueness=Uniqueness.global_, - case_exact=CaseExact.true, ), ], ) + Bar = Extension.from_schema(bar_schema) + FooBar = Resource.from_schema(foo_schema)[Bar] - desc_1 = InMemoryBackend.UniquenessDescriptor( - schema=None, attribute_name="a", case_exact=True - ) - desc_2 = InMemoryBackend.UniquenessDescriptor( - schema="urn:example:2.0:Bar", attribute_name="a", case_exact=True - ) + assert InMemoryBackend.collect_unique_attrs(FooBar) == [ + InMemoryBackend.UniquenessDescriptor( + None, "a", True, "urn:example:2.0:Foo" + ), + InMemoryBackend.UniquenessDescriptor( + "Bar", "a", False, "urn:example:2.0:Bar" + ), + ] - assert InMemoryBackend.collect_resource_unique_attrs( - rt, - { - "urn:example:2.0:Foo": foo_schema, - "urn:example:2.0:Bar": bar_schema, - }, - ) == [desc_1, desc_2] + resource = FooBar.model_validate( + {"a": "ABC", "urn:example:2.0:Bar": {"a": "DEF"}} + ) + foo, bar = InMemoryBackend.collect_unique_attrs(FooBar) + assert foo.get_attribute(resource) == "ABC" + assert bar.get_attribute(resource) == "def" + assert bar.get_attribute(FooBar(a="ABC")) is None - ResType = ResourceType.from_schema(foo_schema)[ - Extension.from_schema(bar_schema) + def test_unique_attributes_of_the_default_user(self, app): + """The only uniqueness constraint checked on a User is userName, the id being assigned by the backend.""" + User = app.provider.model_for("User") + assert InMemoryBackend.collect_unique_attrs(User) == [ + InMemoryBackend.UniquenessDescriptor( + None, "user_name", False, "urn:ietf:params:scim:schemas:core:2.0:User" + ) ] - res = ResType.model_validate( - { - "a": "ABC", - "urn:example:2.0:Bar": { - "a": "DEF", - }, - } - ) - assert desc_1.get_attribute(res) == "ABC" - assert desc_2.get_attribute(res) == "DEF" - def test_query_resources_without_count_returns_every_resource(self, provider): + def test_a_missing_unique_value_does_not_clash(self): + """Two resources lacking a unique value do not conflict, as SQL NULLs do not.""" + + class Badge(Resource): + __schema__ = URN("urn:example:2.0:Badge") + code: Annotated[str | None, Uniqueness.server] = None + + backend = InMemoryBackend() + badge_type = ResourceType.from_resource(Badge) + backend.create_resource(badge_type, Badge()) + backend.create_resource(badge_type, Badge()) + backend.create_resource(badge_type, Badge(code="x")) + with pytest.raises(UniquenessException): + backend.create_resource(badge_type, Badge(code="x")) + + def test_unique_values_are_compared_with_unicode_case_folding(self, app, user_type): + """Unicode case folding makes "Straße" and "STRASSE" the same value.""" + backend = app.backend + User = app.provider.model_for("User") + backend.create_resource(user_type, User(user_name="Straße")) + with pytest.raises(UniquenessException): + backend.create_resource(user_type, User(user_name="STRASSE")) + + def test_query_resources_without_count_returns_every_resource(self, app, user_type): """A search request carrying no count is not paginated by the backend.""" - backend = provider.backend + backend = app.backend for user_name in ("a", "b", "c"): backend.create_resource( - "User", backend.get_model("User")(user_name=user_name) + user_type, app.provider.model_for("User")(user_name=user_name) ) - total_results, resources = backend.query_resources(SearchRequest(), "User") + total_results, resources = backend.query_resources(SearchRequest(), user_type) assert total_results == 3 assert len(resources) == 3 - def test_query_resources_total_results_counts_beyond_the_page(self, provider): + def test_query_resources_total_results_counts_beyond_the_page(self, app, user_type): """The total results count every matching resource, not only the returned page.""" - backend = provider.backend + backend = app.backend for user_name in ("a", "b", "c"): backend.create_resource( - "User", backend.get_model("User")(user_name=user_name) + user_type, app.provider.model_for("User")(user_name=user_name) ) total_results, resources = backend.query_resources( - SearchRequest(start_index=2, count=1), "User" + SearchRequest(start_index=2, count=1), user_type ) assert total_results == 3 assert len(resources) == 1 - def test_meta_resource_type_name(self, provider): - backend = provider.backend - backend.resource_types["User"] = backend.resource_types["User"].model_copy( - update={"name": "User RT Name"} - ) - resource = backend.get_model("User")(user_name="bjensen") - created = backend.create_resource("User", resource) - assert created.meta.resource_type == "User RT Name" - - def test_register_resource_type_unknown_schema(self): - backend = InMemoryBackend() - rt = ResourceType(schema="urn:unknown:Foo") - with pytest.raises(RuntimeError): - backend.register_resource_type(rt) - - schema = Schema(id="urn:unknown:Foo") - backend.register_schema(schema) - rt = ResourceType( - schema="urn:unknown:Foo", - schema_extensions=[ - SchemaExtension(schema="urn:unknown:Bar", required=True) - ], - ) - with pytest.raises(RuntimeError): - backend.register_resource_type(rt) + def test_delete_unknown_resource(self, backend, user_type): + """Deleting a resource the backend does not store reports that nothing was deleted.""" + assert backend.delete_resource(user_type, "unknown") is False - def test_update_unknown_resource(self): - backend = InMemoryBackend() + def test_update_unknown_resource(self, backend, user_type): resource = User(id="123") - assert backend.update_resource("User", resource) is None + assert backend.update_resource(user_type, resource) is None -@pytest.mark.parametrize("resource_type_id", ["User", None]) +@pytest.mark.parametrize("queried", ["User", None]) def test_query_resources_binds_a_filter_that_names_no_resource_type( - provider, resource_type_id + app, user_type, queried ): - """A filter left unbound is resolved against the resource types being queried.""" - backend = provider.backend + """A filter left unbound is resolved against the models of the stored resources.""" + backend = app.backend for user_name in ("alice", "bob"): - backend.create_resource("User", backend.get_model("User")(user_name=user_name)) + backend.create_resource( + user_type, app.provider.model_for("User")(user_name=user_name) + ) request = SearchRequest(filter='userName eq "bob"') - total_results, resources = backend.query_resources(request, resource_type_id) + resource_type = user_type if queried else None + total_results, resources = backend.query_resources(request, resource_type) assert total_results == 1 assert resources[0].user_name == "bob" + + +def test_query_resources_with_an_unbound_filter_and_no_resource(backend): + """With no stored resource, an unbound filter has nothing to be resolved against nor to match.""" + request = SearchRequest(filter='userName eq "bob"') + assert backend.query_resources(request) == (0, []) + + +def test_uniqueness_does_not_span_schemas(): + """Two schemas declaring a unique attribute of the same name do not constrain each other.""" + + class Badge(Resource): + __schema__ = URN("urn:example:2.0:Badge") + code: Annotated[str | None, Uniqueness.server] = None + + class Token(Resource): + __schema__ = URN("urn:example:2.0:Token") + code: Annotated[str | None, Uniqueness.server] = None + + backend = InMemoryBackend() + backend.create_resource(ResourceType.from_resource(Badge), Badge(code="x")) + backend.create_resource(ResourceType.from_resource(Token), Token(code="x")) + + +def test_uniqueness_of_an_extension_spans_the_resources_it_extends(): + """An extension attribute is unique among every resource carrying the extension.""" + + class Tag(Extension): + __schema__ = URN("urn:example:2.0:Tag") + code: Annotated[str | None, Uniqueness.server] = None + + class Badge(Resource): + __schema__ = URN("urn:example:2.0:Badge") + + class Token(Resource): + __schema__ = URN("urn:example:2.0:Token") + + backend = InMemoryBackend() + backend.create_resource( + ResourceType.from_resource(Badge[Tag]), + Badge[Tag].model_validate({"urn:example:2.0:Tag": {"code": "x"}}), + ) + backend.create_resource(ResourceType.from_resource(Token), Token()) + with pytest.raises(UniquenessException): + backend.create_resource( + ResourceType.from_resource(Token[Tag]), + Token[Tag].model_validate({"urn:example:2.0:Tag": {"code": "x"}}), + ) + + +def test_a_resource_type_named_apart_from_its_id(static_data): + """The resources of a resource type whose name differs from its id stay reachable.""" + resource_type = ResourceType( + id="Usr", + name="User", + endpoint="/Users", + schema="urn:ietf:params:scim:schemas:core:2.0:User", + ) + provider = ScimProvider.from_discovery(static_data[0].values(), [resource_type]) + backend = InMemoryBackend() + User = provider.model_for(resource_type) + created = backend.create_resource(resource_type, User(user_name="bjensen")) + assert created.meta.resource_type == "User" + + assert backend.get_resource(resource_type, created.id).user_name == "bjensen" + assert backend.query_resources(SearchRequest(), resource_type)[0] == 1 + with pytest.raises(UniquenessException): + backend.create_resource(resource_type, User(user_name="bjensen")) + assert backend.delete_resource(resource_type, created.id) + assert backend.get_resource(resource_type, created.id) is None diff --git a/tests/test_operators.py b/tests/test_operators.py index 6a9ae05..2db236c 100644 --- a/tests/test_operators.py +++ b/tests/test_operators.py @@ -22,7 +22,6 @@ from scim2_server.operators import ReplaceOperator from scim2_server.operators import ResolveOperator from scim2_server.operators import ResolveResult -from scim2_server.operators import ResolveSortOperator from scim2_server.operators import parse_attribute_path @@ -569,89 +568,3 @@ def test_resolve_operator_complex_multi_valued_attribute(self): assert self._resolve_value("emails[type pr]", u) == home_email assert self._resolve_value("emails[not (type pr)]", u) == work_email assert self._resolve_value("emails[primary eq true]", u) == work_email - - def _resolve_sort_value(self, path: str, model: BaseModel): - operator = ResolveSortOperator(path) - return operator(model) - - def test_resolve_sort_operator(self): - u = User[EnterpriseUser]( - id="123", - user_name="foo", - name=Name(formatted="Mr. Foo"), - emails=[ - Email(value="home@example.com", type="home", primary=False), - Email(value="work@example.com", type="work", primary=True), - ], - ) - u.EnterpriseUser = EnterpriseUser(employee_number="123") - assert self._resolve_sort_value("", u) is None - assert self._resolve_sort_value("id", u) == "123" - assert self._resolve_sort_value("externalId", u) is None - assert self._resolve_sort_value("userName", u) == "foo" - assert self._resolve_sort_value("USERNAME", u) == "foo" - assert self._resolve_sort_value("name.formatted", u) == "mr. foo" - assert self._resolve_sort_value("name.givenName", u) is None - assert self._resolve_sort_value("emails", u) == "work@example.com" - assert ( - self._resolve_sort_value( - "urn:ietf:params:scim:schemas:core:2.0:User:emails", u - ) - == "work@example.com" - ) - assert self._resolve_sort_value("emails.value", u) == "work@example.com" - assert ( - self._resolve_sort_value('emails[type eq "home"]', u) == "home@example.com" - ) - assert ( - self._resolve_sort_value('emails[type eq "home"].value', u) - == "home@example.com" - ) - assert self._resolve_sort_value('emails[type eq "home"].type', u) == "home" - assert self._resolve_sort_value('emails[type eq "other"]', u) is None - assert self._resolve_sort_value('emails[type eq "other"].value', u) is None - assert self._resolve_sort_value('emails[primary eq "True"].type', u) == "work" - assert self._resolve_sort_value('emails[primary eq "False"].type', u) == "home" - assert ( - self._resolve_sort_value( - "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User:employeeNumber", - u, - ) - == "123" - ) - assert ( - self._resolve_sort_value( - "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User:division", u - ) - is None - ) - - assert self._resolve_sort_value("invalidAttribute", u) is None - assert ( - self._resolve_sort_value("invalidAttribute.invalidSubAttribute", u) is None - ) - assert self._resolve_sort_value("name", u) is None - assert self._resolve_sort_value("name.invalidSubAttribute", u) is None - assert ( - self._resolve_sort_value("name[invalidSubAttribute pr].formatted", u) - is None - ) - assert self._resolve_sort_value("name[invalidSubAttribute pr]", u) is None - assert self._resolve_sort_value("id.formatted", u) is None - assert self._resolve_sort_value("id[invalidSubAttribute pr]", u) is None - assert ( - self._resolve_sort_value("id[invalidSubAttribute pr].formatted", u) is None - ) - assert ( - self._resolve_sort_value( - "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User", u - ) - is None - ) - assert ( - self._resolve_sort_value( - "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User:invalidAttribute", - u, - ) - is None - ) diff --git a/tests/test_patch.py b/tests/test_patch.py index 5e636f3..9f51428 100644 --- a/tests/test_patch.py +++ b/tests/test_patch.py @@ -8,8 +8,8 @@ class TestPatch: - def test_patch_operation_add_simple(self, provider): - user = provider.backend.get_model("User")(id="123") + def test_patch_operation_add_simple(self, app): + user = app.provider.model_for("User")(id="123") patch_resource( user, PatchOperation( @@ -63,8 +63,8 @@ def test_patch_operation_add_simple(self, provider): "userName": "Bar", } - def test_patch_operation_add_complex(self, provider): - user = provider.backend.get_model("User")(id="123") + def test_patch_operation_add_complex(self, app): + user = app.provider.model_for("User")(id="123") patch_resource( user, PatchOperation( @@ -143,8 +143,8 @@ def test_patch_operation_add_complex(self, provider): }, } - def test_patch_operation_add_multi_valued(self, provider): - user = provider.backend.get_model("User")(id="123") + def test_patch_operation_add_multi_valued(self, app): + user = app.provider.model_for("User")(id="123") patch_resource( user, PatchOperation( diff --git a/tests/test_provider.py b/tests/test_provider.py index c9d0bca..cb889db 100644 --- a/tests/test_provider.py +++ b/tests/test_provider.py @@ -4,8 +4,8 @@ class TestProvider: - def test_user_creation(self, provider): - user_model = provider.backend.get_model("User").model_validate( + def test_user_creation(self, app, user_type): + user_model = app.provider.model_for("User").model_validate( { "schemas": ["urn:ietf:params:scim:schemas:core:2.0:User"], "userName": "bjensen@example.com", @@ -21,11 +21,11 @@ def test_user_creation(self, provider): }, scim_ctx=Context.RESOURCE_CREATION_REQUEST, ) - ret = provider.backend.create_resource("User", user_model) + ret = app.backend.create_resource(user_type, user_model) assert ret.id is not None - def test_generic_exception_handling(self, provider): - """Test that generic exceptions are properly handled and return 500 status.""" + def test_generic_exception_handling(self, app): + """An unexpected error answers 500 without disclosing its message or traceback.""" from werkzeug import Request # Create a mock WSGI environ @@ -41,14 +41,13 @@ def test_generic_exception_handling(self, provider): # Mock to force a generic exception during request processing with patch.object( - provider, + app, "call_service_provider_config", side_effect=RuntimeError("Test error"), ): - response = provider.wsgi_app(request, environ) + response = app.wsgi_app(request, environ) # Should return a Response object with status 500 assert response.status_code == 500 - # The response should contain error details - response_data = response.get_data(as_text=True) - assert "Test error" in response_data + assert response.json["detail"] == "Internal server error" + assert "Traceback" not in response.get_data(as_text=True) diff --git a/tests/test_utils.py b/tests/test_utils.py index a66e430..86b572f 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -1,18 +1,13 @@ -from typing import Annotated - import pytest from scim2_filter_parser.lexer import SCIMLexer from scim2_filter_parser.parser import SCIMParser -from scim2_models import URN from scim2_models import Context from scim2_models import EnterpriseUser from scim2_models import InvalidFilterException from scim2_models import Meta -from scim2_models import Mutability from scim2_models import MutabilityException from scim2_models import Name from scim2_models import NoTargetException -from scim2_models import Resource from scim2_models import ResponseParameters from scim2_models import SensitiveException from scim2_models import User @@ -20,7 +15,8 @@ from scim2_server.filter import evaluate_filter from scim2_server.operators import ResolveOperator from scim2_server.utils import get_or_create -from scim2_server.utils import merge_resources +from scim2_server.utils import load_default_provider +from scim2_server.utils import load_default_schemas class TestUtils: @@ -53,8 +49,8 @@ def test_case_sensitivity(self): ), ) - def test_match_filter(self, provider): - user = provider.backend.get_model("User").model_validate( + def test_match_filter(self, app): + user = app.provider.model_for("User").model_validate( { "schemas": [ "urn:ietf:params:scim:schemas:core:2.0:User", @@ -179,8 +175,8 @@ def evaluate(filter_str: str) -> bool: 'emails[type eq "work" and value co "@example.com"] or ims[type eq "xmpp" and value co "@foo.com"]' ) - def test_attribute_resolving(self, provider): - user = provider.backend.get_model("User").model_validate( + def test_attribute_resolving(self, app): + user = app.provider.model_for("User").model_validate( { "schemas": [ "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User", @@ -281,7 +277,7 @@ def validate( "Emails", ) - def test_dump_creation(self, provider): + def test_dump_creation(self, app): user = User(id="1", user_name="ABC") user.name = Name(formatted="Barbara") user.meta = Meta( @@ -290,7 +286,7 @@ def test_dump_creation(self, provider): ) user.model_dump(scim_ctx=Context.RESOURCE_CREATION_RESPONSE) - def test_dump_extension(self, provider): + def test_dump_extension(self, app): user = User[EnterpriseUser].model_validate( { "userName": "thomas38@harding-herman.com", @@ -362,36 +358,14 @@ def test_dump_extension(self, provider): }, } - def test_merge_resources_immutable(self): - class Foo(Resource): - __schema__ = URN("urn:example:2.0:Foo") - immutable_string: Annotated[str | None, Mutability.immutable] = None - - stored = Foo() - merge_resources(stored, Foo(immutable_string="ABC")) - assert stored.immutable_string == "ABC" - merge_resources(stored, Foo(immutable_string="ABC")) - with pytest.raises(MutabilityException): - merge_resources(stored, Foo(immutable_string="D")) - def test_get_or_create_mutability(self): u = User() with pytest.raises(MutabilityException): get_or_create(u, "groups", True) - def test_merge_resources_none_extension(self): - """Test adding an extension parameter with merge_resources.""" - target = User[EnterpriseUser](user_name="test") - assert target[EnterpriseUser] is None - - payload = { - "userName": "test", - "urn:ietf:params:scim:schemas:extension:enterprise:2.0:User": { - "employeeNumber": "12345" - }, - } - update = User[EnterpriseUser].model_validate(payload) - - merge_resources(target, update) - assert target[EnterpriseUser].employee_number == "12345" +def test_the_default_provider_publishes_the_default_schemas(): + """The schemas rebuilt from the default models are the ones the package ships.""" + published = [schema.model_dump() for schema in load_default_provider().schemas] + shipped = [schema.model_dump() for schema in load_default_schemas().values()] + assert published == shipped diff --git a/uv.lock b/uv.lock index 8ef07f0..e36fe58 100644 --- a/uv.lock +++ b/uv.lock @@ -683,15 +683,15 @@ wheels = [ [[package]] name = "scim2-models" -version = "0.8.1" +version = "0.8.2" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "lark" }, { name = "pydantic", extra = ["email"] }, ] -sdist = { url = "https://files.pythonhosted.org/packages/df/1c/9010838ba60bac4c8ba6b76cbc8c6c74319f12f96e860cb720b0923f403b/scim2_models-0.8.1.tar.gz", hash = "sha256:8c269153ab252858c125a585565acde2d04be17d2554522315ec5d46d68780c7", size = 91966, upload-time = "2026-09-25T09:39:51.288Z" } +sdist = { url = "https://files.pythonhosted.org/packages/d5/61/bf0746a7ab542c937b3888375f4df8857824a37833df7a5c3eb5ea7d0bd0/scim2_models-0.8.2.tar.gz", hash = "sha256:8ee4b69687d2680c06827af24a50c93a1d33fbfe0be492e0b2db34f4c5233ecd", size = 94332, upload-time = "2026-09-25T19:07:50.984Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/72/59/4daf3488ec1aa235c38047316a4562f74e29083b73368f5b01833372be36/scim2_models-0.8.1-py3-none-any.whl", hash = "sha256:4fc5ed3e6ea250a507d7d371c47bf4ade4141788af8fd5af6ae7282e08783b7f", size = 114332, upload-time = "2026-09-25T09:39:50.024Z" }, + { url = "https://files.pythonhosted.org/packages/e0/8e/6a793d8f2e2cb7f1f180f277b47ca104dbf87a14b8d53abd4e7f3ae37db9/scim2_models-0.8.2-py3-none-any.whl", hash = "sha256:b27eec4847cfdc57e173f78e81004932e7d7298e3192d05f5fe1ec579f4735e9", size = 116865, upload-time = "2026-09-25T19:07:49.035Z" }, ] [[package]] @@ -717,7 +717,7 @@ dev = [ [package.metadata] requires-dist = [ { name = "scim2-filter-parser", specifier = ">=0.7.0" }, - { name = "scim2-models", specifier = ">=0.8.1" }, + { name = "scim2-models", specifier = ">=0.8.2" }, { name = "werkzeug", specifier = ">=3.0.3" }, ]