@@ -51,8 +51,13 @@ def _metadata_lance_schema(task: DocumentBatch) -> pa.Schema | None:
5151 return json_to_schema (schema )
5252
5353
54- def _schema_for_table (schema : pa .Schema , table : pa .Table ) -> pa .Schema :
55- fields = [schema .field (name ) if name in schema .names else table .schema .field (name ) for name in table .column_names ]
54+ def _schema_for_table (lance_schema : pa .Schema , table : pa .Table ) -> pa .Schema :
55+ fields = []
56+ for table_field in table .schema :
57+ if table_field .name in lance_schema .names :
58+ fields .append (lance_schema .field (table_field .name ))
59+ else :
60+ fields .append (table_field )
5661 return pa .schema (fields )
5762
5863
@@ -220,8 +225,6 @@ def process(self, task: DocumentBatch) -> FileGroupTask:
220225 fragment_ids = sorted (int (value ) for value in pc .unique (table [LANCE_FRAGID_COLUMN ].combine_chunks ()).to_pylist ())
221226 for fragment_id in fragment_ids :
222227 update_table = self ._update_table_for_fragment (table , fragment_id )
223- if update_table .num_rows == 0 :
224- continue
225228 fragment = dataset .get_fragment (fragment_id )
226229 updated_fragment , fields_modified = fragment .update_columns (
227230 update_table ,
@@ -256,6 +259,49 @@ def process(self, task: DocumentBatch) -> FileGroupTask:
256259 )
257260
258261
262+ def _validate_checkpoint_path (records : list [dict [str , Any ]], path : str ) -> None :
263+ dataset_paths = {record ["dataset_path" ] for record in records }
264+ if dataset_paths != {path }:
265+ msg = f"Checkpoint records are for { sorted (dataset_paths )} , not { path } "
266+ raise ValueError (msg )
267+
268+
269+ def _single_checkpoint_value (records : list [dict [str , Any ]], key : str , label : str ) -> object :
270+ values = {record [key ] for record in records }
271+ if len (values ) != 1 :
272+ msg = f"Expected one { label } ; got { sorted (values )} "
273+ raise ValueError (msg )
274+ return next (iter (values ))
275+
276+
277+ def _decode_write_fragments (records : list [dict [str , Any ]]) -> list [tuple [object , pa .Schema ]]:
278+ from lance .schema import json_to_schema
279+
280+ return [
281+ (pickle .loads (base64 .b64decode (record ["fragment" ])), json_to_schema (record ["schema" ])) # noqa: S301
282+ for record in records
283+ ]
284+
285+
286+ def _annotation_records_by_fragment (records : list [dict [str , Any ]]) -> dict [int , dict [str , Any ]]:
287+ records_by_fragment = {int (record ["fragment_id" ]): record for record in records }
288+ if len (records_by_fragment ) != len (records ):
289+ msg = "Ensure each Lance fragment is updated by at most one writer task."
290+ raise ValueError (msg )
291+ return records_by_fragment
292+
293+
294+ def _decode_updated_fragments (records : list [dict [str , Any ]]) -> list [object ]:
295+ return [
296+ pickle .loads (base64 .b64decode (record ["updated_fragment" ])) # noqa: S301
297+ for record in records
298+ ]
299+
300+
301+ def _fields_modified (records : list [dict [str , Any ]]) -> list [str ]:
302+ return sorted ({field for record in records for field in record ["fields_modified" ]})
303+
304+
259305def commit_lance_checkpoint (
260306 path : str ,
261307 commit_path : str ,
@@ -265,32 +311,22 @@ def commit_lance_checkpoint(
265311) -> int :
266312 """Commit records written by ``LanceWriter`` and return the Lance version."""
267313 import lance
268- from lance .schema import json_to_schema
269314 from lance_ray import LanceFragmentCommitter
270315
271316 records , committed_version = read_lance_checkpoint (commit_path , "lance_write" , checkpoint_storage_options )
272317 if committed_version is not None :
273318 return committed_version
274319
275- dataset_paths = {record ["dataset_path" ] for record in records }
276- if dataset_paths != {path }:
277- msg = f"Checkpoint records are for { sorted (dataset_paths )} , not { path } "
278- raise ValueError (msg )
279- modes = {record ["mode" ] for record in records }
280- if len (modes ) != 1 :
281- msg = f"Expected one write mode; got { sorted (modes )} "
282- raise ValueError (msg )
283- mode = str (next (iter (modes )))
284- fragments = [
285- (pickle .loads (base64 .b64decode (record ["fragment" ])), json_to_schema (record ["schema" ])) # noqa: S301
286- for record in records
287- ]
320+ _validate_checkpoint_path (records , path )
321+ mode = str (_single_checkpoint_value (records , "mode" , "write mode" ))
322+ fragments = _decode_write_fragments (records )
288323 schema = fragments [0 ][1 ]
289324
290325 committer = LanceFragmentCommitter (path , schema = schema , mode = mode , storage_options = storage_options )
291326 if mode == "append" :
292327 committer .on_write_start (schema )
293- committer .on_write_complete ([[(pickle .dumps (fragment ), pickle .dumps (schema )) for fragment , schema in fragments ]])
328+ fragment_payloads = [(pickle .dumps (fragment ), pickle .dumps (schema )) for fragment , schema in fragments ]
329+ committer .on_write_complete ([fragment_payloads ])
294330 version = lance .dataset (path , storage_options = storage_options ).version
295331 write_lance_checkpoint_marker (commit_path , version , checkpoint_storage_options )
296332 return version
@@ -312,24 +348,12 @@ def commit_lance_annotation_checkpoint(
312348 if committed_version is not None :
313349 return committed_version
314350
315- dataset_paths = {record ["dataset_path" ] for record in records }
316- if dataset_paths != {path }:
317- msg = f"Checkpoint records are for { sorted (dataset_paths )} , not { path } "
318- raise ValueError (msg )
319- read_versions = {int (record ["dataset_version" ]) for record in records }
320- if len (read_versions ) != 1 :
321- msg = f"Expected one dataset version; got { sorted (read_versions )} "
322- raise ValueError (msg )
323- read_version = next (iter (read_versions ))
324- records_by_fragment = {int (record ["fragment_id" ]): record for record in records }
325- if len (records_by_fragment ) != len (records ):
326- msg = "Ensure each Lance fragment is updated by at most one writer task."
327- raise ValueError (msg )
328- updated_fragments = [
329- pickle .loads (base64 .b64decode (record ["updated_fragment" ])) # noqa: S301
330- for record in records_by_fragment .values ()
331- ]
332- fields_modified = sorted ({field for record in records_by_fragment .values () for field in record ["fields_modified" ]})
351+ _validate_checkpoint_path (records , path )
352+ read_version = int (_single_checkpoint_value (records , "dataset_version" , "dataset version" ))
353+ records_by_fragment = _annotation_records_by_fragment (records )
354+ fragment_records = list (records_by_fragment .values ())
355+ updated_fragments = _decode_updated_fragments (fragment_records )
356+ fields_modified = _fields_modified (fragment_records )
333357 operation = lance .LanceOperation .Update (updated_fragments = updated_fragments , fields_modified = fields_modified )
334358 version = lance .LanceDataset .commit (
335359 path ,
0 commit comments