scim2-server 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
File without changes
@@ -0,0 +1,398 @@
1
+ import dataclasses
2
+ import datetime
3
+ import operator
4
+ import pickle
5
+ import uuid
6
+ from threading import Lock
7
+ from typing import Dict
8
+ from typing import List
9
+ from typing import Optional
10
+ from typing import Tuple
11
+ from typing import Union
12
+
13
+ from scim2_filter_parser import lexer
14
+ from scim2_filter_parser.parser import SCIMParser
15
+ from scim2_models import Attribute
16
+ from scim2_models import BaseModel
17
+ from scim2_models import CaseExact
18
+ from scim2_models import Error
19
+ from scim2_models import Extension
20
+ from scim2_models import Meta
21
+ from scim2_models import Resource
22
+ from scim2_models import ResourceType
23
+ from scim2_models import Schema
24
+ from scim2_models import SearchRequest
25
+ from scim2_models import Uniqueness
26
+ from werkzeug.http import generate_etag
27
+
28
+ from scim2_server.filter import evaluate_filter
29
+ from scim2_server.operators import ResolveSortOperator
30
+ from scim2_server.utils import SCIMException
31
+ from scim2_server.utils import get_by_alias
32
+
33
+
34
+ class Backend:
35
+ """The base class for a SCIM provider backend."""
36
+
37
+ def __init__(self):
38
+ self.schemas: Dict[str, Schema] = {}
39
+ self.resource_types: Dict[str, ResourceType] = {}
40
+ self.resource_types_by_endpoint: Dict[str, ResourceType] = {}
41
+ self.models_dict: Dict[str, BaseModel] = {}
42
+
43
+ def __enter__(self):
44
+ """Allows the backend to be used as a context manager.
45
+
46
+ This enables support for transactions.
47
+ """
48
+ return self
49
+
50
+ def __exit__(self, exc_type, exc_val, exc_tb):
51
+ """Exits the transaction."""
52
+ pass
53
+
54
+ def register_schema(self, schema: Schema):
55
+ """Registers a Schema for use with the backend."""
56
+ self.schemas[schema.id] = schema
57
+
58
+ def get_schemas(self):
59
+ """Returns all schemas registered with the backend."""
60
+ return self.schemas.values()
61
+
62
+ def get_schema(self, schema_id: str) -> Optional[Schema]:
63
+ """Gets a schema by its id."""
64
+ return self.schemas.get(schema_id)
65
+
66
+ def register_resource_type(self, resource_type: ResourceType):
67
+ """Registers a ResourceType for use with the backend.
68
+
69
+ The schemas used for the resource and its extensions must have
70
+ been registered with the Backend beforehand.
71
+ """
72
+ if resource_type.schema_ not in self.schemas:
73
+ raise RuntimeError(f"Unknown schema: {resource_type.schema_}")
74
+ for resource_extension in resource_type.schema_extensions or []:
75
+ if resource_extension.schema_ not in self.schemas:
76
+ raise RuntimeError(f"Unknown schema: {resource_extension.schema_}")
77
+
78
+ self.resource_types[resource_type.id] = resource_type
79
+ self.resource_types_by_endpoint[resource_type.endpoint.lower()] = resource_type
80
+
81
+ extensions = [
82
+ Extension.from_schema(self.get_schema(se.schema_))
83
+ for se in resource_type.schema_extensions or []
84
+ ]
85
+ base_schema = self.get_schema(resource_type.schema_)
86
+ self.models_dict[resource_type.id] = Resource.from_schema(base_schema)
87
+ if extensions:
88
+ self.models_dict[resource_type.id] = self.models_dict[resource_type.id][
89
+ Union[*extensions]
90
+ ]
91
+
92
+ def get_resource_types(self):
93
+ """Returns all resource types registered with the backend."""
94
+ return self.resource_types.values()
95
+
96
+ def get_resource_type(self, resource_type_id: str) -> Optional[ResourceType]:
97
+ """Returns the resource type by its id."""
98
+ return self.resource_types.get(resource_type_id)
99
+
100
+ def get_resource_type_by_endpoint(self, endpoint: str) -> Optional[ResourceType]:
101
+ """Returns the resource type by its endpoint."""
102
+ return self.resource_types_by_endpoint.get(endpoint.lower())
103
+
104
+ def get_model(self, resource_type_id: str) -> Optional[BaseModel]:
105
+ """Returns the Pydantic Python model for a given resource type."""
106
+ return self.models_dict.get(resource_type_id)
107
+
108
+ def get_models(self):
109
+ """Returns all Pydantic Python models for all known resource types."""
110
+ return self.models_dict.values()
111
+
112
+ def query_resources(
113
+ self,
114
+ search_request: SearchRequest,
115
+ resource_type_id: Optional[str] = None,
116
+ ) -> Tuple[int, List[Resource]]:
117
+ """Queries the backend for a set of resources.
118
+
119
+ :param search_request: SearchRequest instance describing the
120
+ query.
121
+ :param resource_type_id: ID of the resource type to query. If
122
+ None, all resource types are queried.
123
+ :return: A tuple of "total results" and a List of found
124
+ Resources. The List must contain a copy of resources.
125
+ Mutating elements in the List must not modify the data
126
+ stored in the backend.
127
+ :raises SCIMException: If the backend only supports querying for
128
+ one resource type at a time, setting resource_type_id to
129
+ None the backend may raise a
130
+ SCIMException(Error.make_too_many_error()).
131
+ """
132
+ raise NotImplementedError
133
+
134
+ def get_resource(self, resource_type_id: str, object_id: str) -> Optional[Resource]:
135
+ """Queries the backend for a resources by its ID.
136
+
137
+ :param resource_type_id: ID of the resource type to get the
138
+ object from.
139
+ :param object_id: ID of the object to get.
140
+ :return: The resource object if it exists, None otherwise. The
141
+ resource must be a copy, modifying it must not change the
142
+ data stored in the backend.
143
+ """
144
+ raise NotImplementedError
145
+
146
+ def delete_resource(self, resource_type_id: str, object_id: str) -> bool:
147
+ """Deletes a resource.
148
+
149
+ :param resource_type_id: ID of the resource type to delete the
150
+ object from.
151
+ :param object_id: ID of the object to delete.
152
+ :return: True if the resource was deleted, False otherwise.
153
+ """
154
+ raise NotImplementedError
155
+
156
+ def create_resource(
157
+ self, resource_type_id: str, resource: Resource
158
+ ) -> Optional[Resource]:
159
+ """Creates a resource.
160
+
161
+ :param resource_type_id: ID of the resource type to create.
162
+ :param resource: Resource to create.
163
+ :return: The created resource. Creation should set system-
164
+ defined attributes (ID, Metadata). May be the same object
165
+ that is passed in.
166
+ """
167
+ raise NotImplementedError
168
+
169
+ def update_resource(
170
+ self, resource_type_id: str, resource: Resource
171
+ ) -> Optional[Resource]:
172
+ """Updates a resource. The resource is identified by its ID.
173
+
174
+ :param resource_type_id: ID of the resource type to update.
175
+ :param resource: Resource to update.
176
+ :return: The updated resource. Updating should update the
177
+ "meta.lastModified" data. May be the same object that is
178
+ passed in.
179
+ """
180
+ raise NotImplementedError
181
+
182
+
183
+ class InMemoryBackend(Backend):
184
+ """This is an example in-memory backend for the SCIM provider.
185
+
186
+ It is not optimized for performance. Many operations are O(n) or
187
+ worse, whereas they would perform better with an actual production
188
+ database in the backend. This is intentional to keep the
189
+ implementation simple.
190
+ """
191
+
192
+ @dataclasses.dataclass
193
+ class UniquenessDescriptor:
194
+ """Used to mimic uniqueness constraints e.g. from a SQL database."""
195
+
196
+ schema: Optional[str]
197
+ attribute_name: str
198
+ case_exact: bool
199
+
200
+ def get_attribute(self, resource: Resource):
201
+ if self.schema is not None:
202
+ resource = getattr(resource, get_by_alias(resource, self.schema))
203
+ result = getattr(resource, get_by_alias(resource, self.attribute_name))
204
+ if not self.case_exact:
205
+ result = result.lower()
206
+ return result
207
+
208
+ @classmethod
209
+ def collect_unique_attrs(
210
+ cls, attributes: List[Attribute], schema: Optional[str]
211
+ ) -> List[UniquenessDescriptor]:
212
+ ret = []
213
+ for attr in attributes:
214
+ if attr.uniqueness != Uniqueness.none:
215
+ ret.append(
216
+ cls.UniquenessDescriptor(
217
+ schema, attr.name, attr.case_exact == CaseExact.true
218
+ )
219
+ )
220
+ return ret
221
+
222
+ @classmethod
223
+ def collect_resource_unique_attrs(
224
+ cls, resource_type: ResourceType, schemas: Dict[str, Schema]
225
+ ) -> List[List[UniquenessDescriptor]]:
226
+ ret = cls.collect_unique_attrs(schemas[resource_type.schema_].attributes, None)
227
+ for extension in resource_type.schema_extensions or []:
228
+ ret.extend(
229
+ InMemoryBackend.collect_unique_attrs(
230
+ schemas[extension.schema_].attributes, extension.schema_
231
+ )
232
+ )
233
+ return ret
234
+
235
+ def __init__(self):
236
+ super().__init__()
237
+ self.resources: List[Resource] = []
238
+ self.unique_attributes: Dict[str, List[List[str]]] = {}
239
+ self.lock: Lock = Lock()
240
+
241
+ def __enter__(self):
242
+ """See super docs.
243
+
244
+ The InMemoryBackend uses a simple Lock to synchronize all
245
+ access.
246
+ """
247
+ super().__enter__()
248
+ self.lock.acquire()
249
+ return self
250
+
251
+ def __exit__(self, exc_type, exc_val, exc_tb):
252
+ super().__exit__(exc_type, exc_val, exc_tb)
253
+ self.lock.release()
254
+
255
+ def register_resource_type(self, resource_type: ResourceType):
256
+ super().register_resource_type(resource_type)
257
+ self.unique_attributes[resource_type.id] = self.collect_resource_unique_attrs(
258
+ resource_type, self.schemas
259
+ )
260
+
261
+ def query_resources(
262
+ self,
263
+ search_request: SearchRequest,
264
+ resource_type_id: Optional[str] = None,
265
+ ) -> Tuple[int, List[Resource]]:
266
+ start_index = (search_request.start_index or 1) - 1
267
+
268
+ tree = None
269
+ if search_request.filter is not None:
270
+ token_stream = lexer.SCIMLexer().tokenize(search_request.filter)
271
+ tree = SCIMParser().parse(token_stream)
272
+
273
+ found_resources = [
274
+ r
275
+ for r in self.resources
276
+ if (resource_type_id is None or r.meta.resource_type == resource_type_id)
277
+ and (tree is None or evaluate_filter(r, tree))
278
+ ]
279
+
280
+ if search_request.sort_by is not None:
281
+ descending = search_request.sort_order == SearchRequest.SortOrder.descending
282
+ sort_operator = ResolveSortOperator(search_request.sort_by)
283
+
284
+ # To ensure that unset attributes are sorted last (when ascending, as defined in the RFC),
285
+ # we have to divide the result set into a set and unset subset.
286
+ unset_values = []
287
+ set_values = []
288
+ for resource in found_resources:
289
+ result = sort_operator(resource)
290
+ if result is None:
291
+ unset_values.append(resource)
292
+ else:
293
+ set_values.append((resource, result))
294
+
295
+ set_values.sort(key=operator.itemgetter(1), reverse=descending)
296
+ set_values = [value[0] for value in set_values]
297
+ if descending:
298
+ found_resources = unset_values + set_values
299
+ else:
300
+ found_resources = set_values + unset_values
301
+
302
+ found_resources = found_resources[start_index:]
303
+ if search_request.count is not None:
304
+ found_resources = found_resources[: search_request.count]
305
+ return len(found_resources), found_resources
306
+
307
+ def _get_resource_idx(self, resource_type_id: str, object_id: str) -> Optional[int]:
308
+ return next(
309
+ (
310
+ idx
311
+ for idx, r in enumerate(self.resources)
312
+ if r.meta.resource_type == resource_type_id and r.id == object_id
313
+ ),
314
+ None,
315
+ )
316
+
317
+ def get_resource(self, resource_type_id: str, object_id: str) -> Optional[Resource]:
318
+ resource_dict_idx = self._get_resource_idx(resource_type_id, object_id)
319
+ if resource_dict_idx is not None:
320
+ return self.resources[resource_dict_idx].model_copy(deep=True)
321
+ return None
322
+
323
+ def delete_resource(self, resource_type_id: str, object_id: str) -> bool:
324
+ found = self.get_resource(resource_type_id, object_id)
325
+ if found:
326
+ self.resources = [
327
+ r
328
+ for r in self.resources
329
+ if not (r.meta.resource_type == resource_type_id and r.id == object_id)
330
+ ]
331
+ return True
332
+ return False
333
+
334
+ def create_resource(
335
+ self, resource_type_id: str, resource: Resource
336
+ ) -> Optional[Resource]:
337
+ resource = resource.model_copy(deep=True)
338
+ resource.id = uuid.uuid4().hex
339
+ utcnow = datetime.datetime.now(datetime.UTC)
340
+ resource.meta = Meta(
341
+ resource_type=resource_type_id,
342
+ created=utcnow,
343
+ last_modified=utcnow,
344
+ location="/v2"
345
+ + self.resource_types[resource_type_id].endpoint
346
+ + "/"
347
+ + resource.id,
348
+ )
349
+ self._touch_resource(resource, utcnow)
350
+
351
+ for unique_attribute in self.unique_attributes[resource_type_id]:
352
+ new_value = unique_attribute.get_attribute(resource)
353
+ for existing_resource in self.resources:
354
+ if existing_resource.meta.resource_type == resource_type_id:
355
+ existing_value = unique_attribute.get_attribute(existing_resource)
356
+ if existing_value == new_value:
357
+ raise SCIMException(Error.make_uniqueness_error())
358
+
359
+ self.resources.append(resource)
360
+ return resource
361
+
362
+ @staticmethod
363
+ def _touch_resource(resource: Resource, last_modified: datetime.datetime):
364
+ """Touches a resource (updates last_modified and version).
365
+
366
+ Version is generated by hashing last_modified. Another option
367
+ would be to hash the entire resource instead.
368
+ """
369
+ resource.meta.last_modified = last_modified
370
+ etag = generate_etag(pickle.dumps(resource.meta.last_modified))
371
+ resource.meta.version = f'W/"{etag}"'
372
+
373
+ def update_resource(
374
+ self, resource_type_id: str, resource: Resource
375
+ ) -> Optional[Resource]:
376
+ found_res_idx = self._get_resource_idx(resource_type_id, resource.id)
377
+ if found_res_idx is not None:
378
+ updated_resource = self.models_dict[resource_type_id].model_validate(
379
+ resource.model_dump()
380
+ )
381
+ self._touch_resource(updated_resource, datetime.datetime.now(datetime.UTC))
382
+
383
+ for unique_attribute in self.unique_attributes[resource_type_id]:
384
+ new_value = unique_attribute.get_attribute(updated_resource)
385
+ for existing_resource in self.resources:
386
+ if (
387
+ existing_resource.meta.resource_type == resource_type_id
388
+ and existing_resource.id != updated_resource.id
389
+ ):
390
+ existing_value = unique_attribute.get_attribute(
391
+ existing_resource
392
+ )
393
+ if existing_value == new_value:
394
+ raise SCIMException(Error.make_uniqueness_error())
395
+
396
+ self.resources[found_res_idx] = updated_resource
397
+ return updated_resource
398
+ return None
scim2_server/cli.py ADDED
@@ -0,0 +1,98 @@
1
+ import argparse
2
+ import json
3
+ import logging
4
+ import pprint
5
+
6
+ from scim2_models import ResourceType
7
+ from scim2_models import Schema
8
+ from werkzeug.middleware.proxy_fix import ProxyFix
9
+
10
+ from scim2_server.backend import InMemoryBackend
11
+ from scim2_server.provider import SCIMProvider
12
+ from scim2_server.utils import load_default_resource_types
13
+ from scim2_server.utils import load_default_schemas
14
+
15
+
16
+ def log_environ(handler):
17
+ """A simple decorator to log all WSGI environment variables."""
18
+
19
+ def _inner(environ, start_fn):
20
+ logging.getLogger("log_environ").debug(pprint.pformat(environ))
21
+ return handler(environ, start_fn)
22
+
23
+ return _inner
24
+
25
+
26
+ def main():
27
+ parser = argparse.ArgumentParser()
28
+ parser.add_argument(
29
+ "--schema", type=argparse.FileType("r"), help="Schema definitions"
30
+ )
31
+ parser.add_argument(
32
+ "--resource-type", type=argparse.FileType("r"), help="Resource Type definitions"
33
+ )
34
+ parser.add_argument("--bearer-token", action="append", help="Add Bearer Token")
35
+ parser.add_argument("--hostname", default="127.0.0.1", help="Hostname")
36
+ parser.add_argument("--port", default=8080, type=int, help="Port number")
37
+ parser.add_argument(
38
+ "--reverse-proxy",
39
+ action="store_true",
40
+ help='Allow running behind a reverse proxy (respect "X-Forwarded-*" HTTP headers)',
41
+ )
42
+ parser.add_argument(
43
+ "--dump-resources",
44
+ type=argparse.FileType("w"),
45
+ help="Dump resources to a JSON file on exit",
46
+ )
47
+ args = parser.parse_args()
48
+
49
+ logging.basicConfig(level=logging.DEBUG)
50
+
51
+ from werkzeug.serving import run_simple
52
+
53
+ backend = InMemoryBackend()
54
+ app = SCIMProvider(backend)
55
+
56
+ if args.schema is None:
57
+ for schema in load_default_schemas().values():
58
+ app.register_schema(schema)
59
+ else:
60
+ def_sch = json.load(args.schema)
61
+ for sc in def_sch:
62
+ schema = Schema.model_validate(sc)
63
+ app.register_schema(schema)
64
+ args.schema.close()
65
+
66
+ if args.resource_type is None:
67
+ for resource_type in load_default_resource_types().values():
68
+ app.register_resource_type(resource_type)
69
+ else:
70
+ def_rt = json.load(args.resource_type)
71
+ for rt in def_rt:
72
+ resource_type = ResourceType.model_validate(rt)
73
+ app.register_resource_type(resource_type)
74
+ args.resource_type.close()
75
+
76
+ if args.bearer_token is not None:
77
+ for bearer_token in args.bearer_token:
78
+ app.register_bearer_token(bearer_token)
79
+
80
+ app = log_environ(app)
81
+ if args.reverse_proxy:
82
+ app = ProxyFix(app, x_for=1, x_proto=1, x_host=1, x_port=1, x_prefix=1)
83
+
84
+ run_simple(
85
+ args.hostname,
86
+ args.port,
87
+ app,
88
+ use_debugger=True,
89
+ use_reloader=True,
90
+ )
91
+
92
+ if args.dump_resources:
93
+ with args.dump_resources as f:
94
+ f.write(json.dumps([r.model_dump() for r in backend.resources], indent=2))
95
+
96
+
97
+ if __name__ == "__main__":
98
+ main()
scim2_server/filter.py ADDED
@@ -0,0 +1,128 @@
1
+ from types import NoneType
2
+ from typing import List
3
+
4
+ from scim2_filter_parser import ast as scim2ast
5
+ from scim2_models import BaseModel
6
+ from scim2_models import CaseExact
7
+ from scim2_models import Error
8
+
9
+ from scim2_server.utils import SCIMException
10
+ from scim2_server.utils import get_by_alias
11
+ from scim2_server.utils import parse_new_value
12
+
13
+
14
+ def evaluate_filter(
15
+ obj: BaseModel | List[BaseModel], tree: scim2ast.AST
16
+ ) -> bool | List[bool]:
17
+ """This implementation is limited by the specifics of the
18
+ scim2_filter_parser module.
19
+
20
+ It works well enough for simple cases, though. It should be re-
21
+ implemented in the future. Probably once
22
+ https://github.com/yaal-coop/scim2-models/issues/17
23
+ is implemented.
24
+ """
25
+ from scim2_server.operators import ResolveOperator
26
+ from scim2_server.operators import ResolveResult
27
+
28
+ match type(tree):
29
+ case scim2ast.Filter:
30
+ if tree.namespace is not None:
31
+ obj = ResolveOperator(tree.namespace.attr_name)(obj).get_values()
32
+ if isinstance(obj, List):
33
+ return [
34
+ o
35
+ for o in obj
36
+ if bool(evaluate_filter(o, tree.expr)) != tree.negated
37
+ ]
38
+ return bool(evaluate_filter(obj, tree.expr)) != tree.negated
39
+ case scim2ast.LogExpr:
40
+ match tree.op:
41
+ case "and":
42
+ return evaluate_filter(obj, tree.expr1) and evaluate_filter(
43
+ obj, tree.expr2
44
+ )
45
+ case _: # "or"
46
+ return evaluate_filter(obj, tree.expr1) or evaluate_filter(
47
+ obj, tree.expr2
48
+ )
49
+ case _: # scim2ast.AttrExpr
50
+ path = tree.attr_path.attr_name
51
+ sub_attribute_name = None
52
+ if isinstance(path, scim2ast.Filter):
53
+ resolved = evaluate_filter(obj, path)
54
+ model = resolved[0]
55
+ attribute_name = ""
56
+
57
+ # FIXME: Best guesses since there is no way to know for sure by this point
58
+ case_sensitivity = CaseExact.false
59
+ sub_attribute_name = path.namespace.sub_attr.value
60
+ else:
61
+ if tree.attr_path.uri:
62
+ path = tree.attr_path.uri + ":" + path
63
+ if tree.attr_path.sub_attr:
64
+ path += "." + tree.attr_path.sub_attr.value
65
+ resolved = ResolveOperator(path)(obj)
66
+ case_sensitivity = resolved.get_field_annotation(CaseExact)
67
+ model = resolved.model
68
+ attribute_name = resolved.attribute
69
+
70
+ if isinstance(resolved, ResolveResult):
71
+ value = resolved.get_values()
72
+ else:
73
+ value = [
74
+ getattr(v, get_by_alias(v, sub_attribute_name)) for v in resolved
75
+ ]
76
+
77
+ compare_value = None
78
+ if tree.comp_value:
79
+ if attribute_name:
80
+ compare_value = parse_new_value(
81
+ model, attribute_name, tree.comp_value.value
82
+ )
83
+ else:
84
+ compare_value = tree.comp_value.value
85
+ if not case_sensitivity and isinstance(value, str):
86
+ value = value.lower()
87
+ if compare_value:
88
+ compare_value = compare_value.lower()
89
+
90
+ match tree.value:
91
+ case "eq":
92
+ return value == compare_value
93
+ case "ne":
94
+ return value != compare_value
95
+ case "sw":
96
+ return value.startswith(compare_value)
97
+ case "ew":
98
+ return value.endswith(compare_value)
99
+ case "pr":
100
+ return bool(value)
101
+ case "co":
102
+ if value is None:
103
+ return False
104
+ return compare_value in value
105
+ case "gt":
106
+ check_comparable_value(value)
107
+ return value > compare_value
108
+ case "lt":
109
+ check_comparable_value(value)
110
+ return value < compare_value
111
+ case "ge":
112
+ check_comparable_value(value)
113
+ return value >= compare_value
114
+ case _: # "le"
115
+ check_comparable_value(value)
116
+ return value <= compare_value
117
+ return False
118
+
119
+
120
+ def check_comparable_value(value):
121
+ """Certain values may not be compared in a filter, see RFC 7644, section
122
+ 3.4.2.2:
123
+
124
+ "Boolean and Binary attributes SHALL cause a failed response (HTTP
125
+ status code 400) with "scimType" of "invalidFilter"."
126
+ """
127
+ if isinstance(value, (bytes, bool, NoneType)):
128
+ raise SCIMException(Error.make_invalid_filter_error())