Skip to content

Commit ded8cfe

Browse files
committed
refactor: validate bulk requests through scim2-models
1 parent 3cb43f7 commit ded8cfe

4 files changed

Lines changed: 174 additions & 113 deletions

File tree

‎scim2_server/bulk.py‎

Lines changed: 65 additions & 48 deletions
Original file line numberDiff line numberDiff line change
@@ -1,55 +1,71 @@
11
from collections.abc import Callable
22
from typing import Any
33

4+
from pydantic import BaseModel
5+
from scim2_models import BulkOperation
46
from scim2_models import InvalidValueException
57
from scim2_models import Resource
68
from werkzeug.exceptions import Conflict
79

810
BULK_ID_PREFIX = "bulkId:"
911

10-
Resolver = Callable[[Any], Any]
11-
"""Replaces the bulkId references of a raw operation."""
12+
Resolver = Callable[[BulkOperation], BulkOperation]
13+
"""Replaces the bulkId references of an operation."""
1214

13-
OperationRunner = Callable[[Any, Resolver], tuple[dict[str, Any], Resource | None]]
14-
"""Applies a raw operation once resolved, and returns its outcome and the resource it acted on."""
15-
16-
17-
def raw_attribute(payload: Any, name: str) -> Any:
18-
"""Return an attribute of a raw payload, whose names are case insensitive."""
19-
if not isinstance(payload, dict):
20-
return None
21-
return next(
22-
(value for key, value in payload.items() if key.casefold() == name.casefold()),
23-
None,
24-
)
15+
OperationRunner = Callable[
16+
[BulkOperation, Resolver], tuple[dict[str, Any], Resource | None]
17+
]
18+
"""Applies an operation once resolved, and returns its outcome and the resource it acted on."""
2519

2620

2721
def replace_bulk_ids(value: Any, replace: Callable[[str], str]) -> Any:
28-
"""Replace every "bulkId:" reference of a raw value."""
22+
"""Replace every "bulkId:" reference of a value.
23+
24+
A value without reference is returned as it is. Models are copied with
25+
only the changed fields, so the fields the client set stay the same.
26+
"""
2927
if isinstance(value, str) and value.startswith(BULK_ID_PREFIX):
3028
return replace(value.removeprefix(BULK_ID_PREFIX))
29+
3130
if isinstance(value, list):
32-
return [replace_bulk_ids(item, replace) for item in value]
31+
items = [replace_bulk_ids(item, replace) for item in value]
32+
changed = any(new is not old for new, old in zip(items, value, strict=True))
33+
return items if changed else value
34+
3335
if isinstance(value, dict):
34-
return {key: replace_bulk_ids(item, replace) for key, item in value.items()}
36+
entries = {key: replace_bulk_ids(item, replace) for key, item in value.items()}
37+
changed = any(entries[key] is not item for key, item in value.items())
38+
return entries if changed else value
39+
40+
if isinstance(value, BaseModel):
41+
updates = {}
42+
for name in type(value).model_fields:
43+
field = getattr(value, name)
44+
replaced = replace_bulk_ids(field, replace)
45+
if replaced is not field:
46+
updates[name] = replaced
47+
return value.model_copy(update=updates) if updates else value
48+
3549
return value
3650

3751

38-
def resolve_operation(payload: Any, replace: Callable[[str], str]) -> Any:
39-
"""Replace the "bulkId:" references of the path and the data of a raw bulk operation."""
40-
if not isinstance(payload, dict):
41-
return payload
52+
def resolve_operation(
53+
operation: BulkOperation, replace: Callable[[str], str]
54+
) -> BulkOperation:
55+
"""Replace the "bulkId:" references of the path and the data of a bulk operation."""
56+
updates: dict[str, Any] = {}
57+
if operation.path is not None:
58+
path = "/".join(
59+
replace_bulk_ids(segment, replace) for segment in operation.path.split("/")
60+
)
61+
if path != operation.path:
62+
updates["path"] = path
4263

43-
resolved = {}
44-
for key, value in payload.items():
45-
if key.casefold() == "path" and isinstance(value, str):
46-
value = "/".join(
47-
replace_bulk_ids(segment, replace) for segment in value.split("/")
48-
)
49-
elif key.casefold() == "data":
50-
value = replace_bulk_ids(value, replace)
51-
resolved[key] = value
52-
return resolved
64+
data = replace_bulk_ids(operation.data, replace)
65+
if data is not operation.data:
66+
updates["data"] = data
67+
68+
return operation.model_copy(update=updates) if updates else operation
5369

5470

5571
class BulkJob:
@@ -63,7 +79,7 @@ class BulkJob:
6379

6480
def __init__(
6581
self,
66-
operations: list[Any],
82+
operations: list[BulkOperation],
6783
fail_on_errors: int | None,
6884
run: OperationRunner,
6985
):
@@ -76,10 +92,12 @@ def __init__(
7692
self.errors = 0
7793

7894
self.creations: dict[str, int] = {}
79-
for index, payload in enumerate(operations):
80-
bulk_id = raw_attribute(payload, "bulkId")
81-
if raw_attribute(payload, "method") == "POST" and isinstance(bulk_id, str):
82-
self.creations.setdefault(bulk_id, index)
95+
for index, operation in enumerate(operations):
96+
if (
97+
operation.method == BulkOperation.Method.post
98+
and operation.bulk_id is not None
99+
):
100+
self.creations.setdefault(operation.bulk_id, index)
83101

84102
@property
85103
def stopped(self) -> bool:
@@ -105,15 +123,15 @@ def run_operation(self, index: int) -> None:
105123
if index in self.results or index in self.running or self.stopped:
106124
return
107125

108-
payload = self.operations[index]
126+
operation = self.operations[index]
109127
self.running.add(index)
110-
for bulk_id in self.references(payload):
128+
for bulk_id in self.references(operation):
111129
if bulk_id in self.creations:
112130
self.run_operation(self.creations[bulk_id])
113131

114132
if not self.stopped:
115133
result, resource = self.run_resolved(
116-
payload, lambda payload: self.resolve(index, payload)
134+
operation, lambda operation: self.resolve(index, operation)
117135
)
118136
self.results[index] = result
119137
if result["status"] >= 400:
@@ -127,31 +145,30 @@ def is_creation(self, index: int, bulk_id: str | None) -> bool:
127145
return bulk_id is not None and self.creations.get(bulk_id) == index
128146

129147
@staticmethod
130-
def references(payload: Any) -> list[str]:
148+
def references(operation: BulkOperation) -> list[str]:
131149
"""Return the bulkIds an operation references."""
132150
bulk_ids: list[str] = []
133151

134152
def collect(bulk_id: str) -> str:
135153
bulk_ids.append(bulk_id)
136154
return bulk_id
137155

138-
resolve_operation(payload, collect)
156+
resolve_operation(operation, collect)
139157
return bulk_ids
140158

141-
def resolve(self, index: int, payload: Any) -> Any:
159+
def resolve(self, index: int, operation: BulkOperation) -> BulkOperation:
142160
"""Replace the bulkId references of an operation with the identifiers of the created resources.
143161
144162
:raises Conflict: When a referenced resource was not created, as
145163
RFC 7644 §3.7.1 allows for circular references.
146164
"""
147-
bulk_id = raw_attribute(payload, "bulkId")
148165
if (
149-
raw_attribute(payload, "method") == "POST"
150-
and isinstance(bulk_id, str)
151-
and not self.is_creation(index, bulk_id)
166+
operation.method == BulkOperation.Method.post
167+
and operation.bulk_id is not None
168+
and not self.is_creation(index, operation.bulk_id)
152169
):
153170
raise InvalidValueException(
154-
detail=f"The bulkId {bulk_id} is not unique in the request"
171+
detail=f"The bulkId {operation.bulk_id} is not unique in the request"
155172
)
156173

157174
def replace(bulk_id: str) -> str:
@@ -161,4 +178,4 @@ def replace(bulk_id: str) -> str:
161178
raise Conflict(f"The bulkId {bulk_id} is part of a circular reference")
162179
raise Conflict(f"No resource was created with the bulkId {bulk_id}")
163180

164-
return resolve_operation(payload, replace)
181+
return resolve_operation(operation, replace)

‎scim2_server/provider.py‎

Lines changed: 28 additions & 56 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,6 @@
4747
from scim2_server.backend import Backend
4848
from scim2_server.bulk import BulkJob
4949
from scim2_server.bulk import Resolver
50-
from scim2_server.bulk import raw_attribute
5150
from scim2_server.utils import load_default_service_provider_config
5251

5352
SEARCH_REQUEST_PARAMETERS = (
@@ -516,7 +515,10 @@ def call_bulk(self, request: Request, **kwargs) -> Response:
516515
f"The payload exceeds the maxPayloadSize ({bulk.max_payload_size} bytes)"
517516
)
518517

519-
bulk_request, operations = self.read_bulk_request(request.json)
518+
bulk_request = BulkRequest[Union[tuple(self.get_models())]].model_validate( # noqa: UP007
519+
request.json, scim_ctx=Context.BULK_REQUEST
520+
)
521+
operations = cast(list[BulkOperation], bulk_request.operations)
520522
if bulk.max_operations is not None and len(operations) > bulk.max_operations:
521523
raise RequestEntityTooLarge(
522524
f"The number of operations exceeds the maxOperations ({bulk.max_operations})"
@@ -525,59 +527,48 @@ def call_bulk(self, request: Request, **kwargs) -> Response:
525527
results = BulkJob(
526528
operations,
527529
bulk_request.fail_on_errors,
528-
lambda payload, resolve: self.run_bulk_operation(request, payload, resolve),
530+
lambda operation, resolve: self.run_bulk_operation(
531+
request, operation, resolve
532+
),
529533
).run()
530534
return self.make_response(
531535
BulkResponse[Union[tuple(self.get_models())]]( # noqa: UP007
532536
operations=results
533537
).model_dump(scim_ctx=Context.BULK_RESPONSE)
534538
)
535539

536-
def read_bulk_request(self, payload: Any) -> tuple[BulkRequest, list[Any]]:
537-
"""Validate the envelope of a bulk request, and return it with its raw operations.
538-
539-
Each operation is validated on its own, so an invalid operation only
540-
fails itself (RFC 7644 §3.7.3).
541-
"""
542-
key = (
543-
next((key for key in payload if key.casefold() == "operations"), None)
544-
if isinstance(payload, dict)
545-
else None
546-
)
547-
operations = payload[key] if key else None
548-
envelope = {**payload, key: []} if isinstance(operations, list) else payload
549-
bulk_request = BulkRequest[Union[tuple(self.get_models())]].model_validate( # noqa: UP007
550-
envelope, scim_ctx=Context.BULK_REQUEST
551-
)
552-
return bulk_request, cast(list[Any], operations)
553-
554540
def run_bulk_operation(
555-
self, request: Request, payload: Any, resolve: Resolver
541+
self, request: Request, operation: BulkOperation, resolve: Resolver
556542
) -> tuple[dict[str, Any], Resource | None]:
557543
"""Apply one operation of a bulk job.
558544
545+
An operation that failed its validation keeps its error, once its
546+
references are resolved to locate it.
547+
559548
:return: The outcome of the operation, and the resource it created or updated.
560549
"""
561-
method = raw_attribute(payload, "method")
562-
bulk_id = raw_attribute(payload, "bulkId")
563550
result: dict[str, Any] = {
564-
"method": method
565-
if method in [member.value for member in BulkOperation.Method]
566-
else None,
567-
"bulk_id": bulk_id if isinstance(bulk_id, str) else None,
551+
"method": operation.method,
552+
"bulk_id": operation.bulk_id,
568553
}
569554

570555
try:
571-
payload = resolve(payload)
572-
resource_type, resource_id = self.get_bulk_target(payload)
573-
if resource_id:
556+
operation = resolve(operation)
557+
resource_type = self.get_resource_type_by_endpoint(operation.endpoint or "")
558+
if resource_type is not None and operation.resource_id:
574559
result["location"] = urljoin(
575-
request.url, f"{resource_type.endpoint.strip('/')}/{resource_id}"
560+
request.url,
561+
f"{resource_type.endpoint.strip('/')}/{operation.resource_id}",
576562
)
577-
operation = BulkOperation[self.get_model(resource_type)].model_validate(
578-
payload, scim_ctx=Context.BULK_REQUEST
563+
if isinstance(operation.response, Error):
564+
return {
565+
**result,
566+
"status": operation.status,
567+
"response": operation.response,
568+
}, None
569+
resource = self.apply_bulk_operation(
570+
cast(ResourceType, resource_type), operation
579571
)
580-
resource = self.apply_bulk_operation(resource_type, resource_id, operation)
581572
except Exception as exception:
582573
error = self.error_from(exception)
583574
return {**result, "status": error.status, "response": error}, None
@@ -591,34 +582,15 @@ def run_bulk_operation(
591582
result["version"] = resource.meta.version
592583
return result, resource
593584

594-
def get_bulk_target(self, payload: Any) -> tuple[ResourceType, str]:
595-
"""Return the resource type and the resource identifier of a bulk operation path.
596-
597-
:raises NotFound: When the path does not start with a resource type endpoint.
598-
"""
599-
path = raw_attribute(payload, "path")
600-
if not isinstance(path, str):
601-
raise InvalidValueException(
602-
detail="path is required for request operations"
603-
)
604-
605-
endpoint, _, resource_id = path.lstrip("/").partition("/")
606-
resource_type = self.get_resource_type_by_endpoint(endpoint)
607-
if resource_type is None:
608-
raise NotFound
609-
return resource_type, resource_id
610-
611585
def apply_bulk_operation(
612-
self,
613-
resource_type: ResourceType,
614-
resource_id: str,
615-
operation: BulkOperation,
586+
self, resource_type: ResourceType, operation: BulkOperation
616587
) -> Resource | None:
617588
"""Apply a validated bulk operation, and return the resource it acted on.
618589
619590
The data of the operation is already validated, and the resource
620591
operations take it as it is.
621592
"""
593+
resource_id = operation.resource_id
622594
if (operation.method == BulkOperation.Method.post) == bool(resource_id):
623595
raise InvalidValueException(
624596
detail="A POST path must target a resource type endpoint, other methods a resource"

0 commit comments

Comments
 (0)