11from collections .abc import Callable
22from typing import Any
33
4+ from pydantic import BaseModel
5+ from scim2_models import BulkOperation
46from scim2_models import InvalidValueException
57from scim2_models import Resource
68from werkzeug .exceptions import Conflict
79
810BULK_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
2721def 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
5571class 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 )
0 commit comments