diff --git a/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_specs_e2e_test.py b/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_specs_e2e_test.py index cf39bb386..e8c7c28b0 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_specs_e2e_test.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_specs_e2e_test.py @@ -15,7 +15,9 @@ """End-to-end tests for distributed arrays serialization with OCDBT.""" import asyncio +import contextlib import dataclasses +import tempfile from typing import TypeAlias import unittest @@ -158,6 +160,9 @@ async def _write_array( ts_context: ts.Context, *, store_ocdbt_metadata_and_values_separately: bool = False, + temporary_metadata_context: ( + tensorstore_utils.OcdbtTemporaryMetadataContext | None + ) = None, ) -> None: """Writes array fragments to the given path with the given process id.""" array_write_tspec = tensorstore_utils.ArrayWriteSpec( @@ -171,6 +176,7 @@ async def _write_array( store_ocdbt_metadata_and_values_separately=( store_ocdbt_metadata_and_values_separately ), + ocdbt_temporary_metadata_context=temporary_metadata_context, ).json if _should_create_ts(array_fragments): @@ -300,6 +306,126 @@ async def _write(): ) await _verify_array_data(test_data, test_dir) + @parameterized.product( + store_ocdbt_metadata_and_values_separately=(False, True), + ) + async def test_write_with_temporary_metadata_context( + self, + store_ocdbt_metadata_and_values_separately: bool, + ): + test_dir = epath.Path(self.create_tempdir()) / "test_data" + test_dir.mkdir(parents=True, exist_ok=True) + + test_data = build_test_data() + + # Create process-specific persistent subdirectories. + for process_id in all_process_ids(test_data): + spec = ocdbt_process_spec.OcdbtProcessSpec(process_id=process_id) + (test_dir / str(spec)).mkdir(parents=False, exist_ok=False) + + ts_context = tensorstore_utils.get_ts_context(use_ocdbt=True) + + async def _write(): + exit_stack = contextlib.ExitStack() + with exit_stack: + # Allocate temporary directories for each process's temporary metadata. + tmp_metadata_context_by_process_id: dict[ + str, tensorstore_utils.OcdbtTemporaryMetadataContext + ] = {} + + def _get_tmp_metadata_context_by_process_id( + process_id: str, + ) -> tensorstore_utils.OcdbtTemporaryMetadataContext: + if process_id not in tmp_metadata_context_by_process_id: + tmp_context = exit_stack.enter_context( + tempfile.TemporaryDirectory() + ) + tmp_metadata_context_by_process_id[process_id] = ( + tensorstore_utils.OcdbtTemporaryMetadataContext( + path=epath.Path(tmp_context) + ) + ) + return tmp_metadata_context_by_process_id[process_id] + + write_futures = [] + for array in test_data: + for process_id, fragments in array.fragments_by_process_id.items(): + write_futures.append( + _write_array( + array.name, + test_dir, + fragments, + process_id, + ts_context, + store_ocdbt_metadata_and_values_separately=( + store_ocdbt_metadata_and_values_separately + ), + temporary_metadata_context=( + _get_tmp_metadata_context_by_process_id(process_id) + ), + ) + ) + await asyncio.gather(*write_futures) + + # Verify that temporary metadata has been written as expected. + for process_id in all_process_ids(test_data): + process_spec = ocdbt_process_spec.OcdbtProcessSpec( + process_id=process_id + ) + process_dir = test_dir / str(process_spec) + tmp_metadata_dir = _get_tmp_metadata_context_by_process_id( + process_id + ).path + # Manifest should not have been written to the final destination, but + # should exist in the temporary metadata context. + self.assertFalse((process_dir / "manifest.ocdbt").exists()) + tmp_manifest_path = ( + tmp_metadata_dir / "ocdbt_tmp_meta/manifest.ocdbt" + ) + self.assertTrue(tmp_manifest_path.exists()) + # Check that metadata is not written to the final destination. We can + # only check this easily if we're writing metadata and values + # separately. + if store_ocdbt_metadata_and_values_separately: + self.assertFalse((process_dir / "ocdbt_meta").is_dir()) + self.assertTrue( + (tmp_metadata_dir / "ocdbt_tmp_meta").is_dir() + ) + # b_large should have generated files written to ocdbt_data/ subdir + # directly, bypassing the temporary metadata context. + if process_id in ("h0", "h1"): + self.assertTrue((process_dir / "ocdbt_data").is_dir()) + + # Commit temporary metadata to persistent storage. + commit_metadata_futures = [] + for process_id in all_process_ids(test_data): + process_spec = ocdbt_process_spec.OcdbtProcessSpec( + process_id=process_id + ) + commit_metadata_futures.append( + ocdbt_utils.commit_temporary_ocdbt_metadata( + test_dir / str(process_spec), + _get_tmp_metadata_context_by_process_id(process_id), + ts_context, + store_ocdbt_metadata_and_values_separately=( + store_ocdbt_metadata_and_values_separately + ), + ) + ) + await asyncio.gather(*commit_metadata_futures) + + await _write() + self._verify_per_process_ocdbt_files( + test_data, + test_dir, + store_ocdbt_metadata_and_values_separately, + ) + + await ocdbt_utils.merge_ocdbt_per_process_files( + test_dir, ts_context, use_zarr3=False, enable_validation=False + ) + await _verify_array_data(test_data, test_dir) + if __name__ == "__main__": absltest.main() diff --git a/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_utils.py b/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_utils.py index 8d74ec420..28c179234 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_utils.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/ocdbt_utils.py @@ -220,6 +220,37 @@ async def merge_ocdbt_per_process_files( ) +async def commit_temporary_ocdbt_metadata( + persistent_path: epath.Path, + temporary_metadata_context: ts_utils.OcdbtTemporaryMetadataContext, + ts_context: ts.Context, + *, + store_ocdbt_metadata_and_values_separately: bool = False, +) -> None: + """Commits temporary OCDBT metadata to the persistent metadata directory.""" + target_kvstore_tspec = ts_utils.build_kvstore_tspec( + persistent_path.as_posix(), + use_ocdbt=True, + ocdbt_write_options=ts_utils.OcdbtKvStoreWriteOptions( + mode=ts_utils.OcdbtWriteMode.COMMIT_TEMPORARY, + store_ocdbt_metadata_and_values_separately=( + store_ocdbt_metadata_and_values_separately + ), + ), + ocdbt_temporary_metadata_context=temporary_metadata_context, + ) + source_kvstore_tspec = ts_utils.build_kvstore_tspec( + persistent_path.as_posix(), + use_ocdbt=True, + ocdbt_temporary_metadata_context=temporary_metadata_context, + ) + target_kvstore, source_kvstore = await asyncio.gather( + ts_utils.open_kv_store(target_kvstore_tspec, ts_context), + ts_utils.open_kv_store(source_kvstore_tspec, ts_context), + ) + await source_kvstore.experimental_copy_range_to(target_kvstore) + + def get_process_index_for_subdir( use_ocdbt: bool, override_ocdbt_process_id: Optional[str] = None, diff --git a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py index 9fe37b4c3..aa0d15366 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils.py @@ -69,6 +69,7 @@ # 'ocdbt_data/' subdirectory. _OCDBT_SPLIT_VALUE_DATA_PREFIX = 'ocdbt_data/' _OCDBT_SPLIT_META_DATA_PREFIX = 'ocdbt_meta/' +_OCDBT_TMP_METADATA_PREFIX = 'ocdbt_tmp_meta/' ZARR_VER2 = 'zarr' ZARR_VER3 = 'zarr3' @@ -155,10 +156,13 @@ class OcdbtWriteMode(enum.Enum): WRITE: Used when writing checkpoint data. MERGE: Used for target (parent) KvStore when merging OCDBT metadata from all per-process subdirectories. + COMMIT_TEMPORARY: Used when committing metadata accumulated in a temporary + metadata directory to its target persistent location. """ WRITE = 'write' MERGE = 'merge' + COMMIT_TEMPORARY = 'commit_temporary' @dataclasses.dataclass(frozen=True) @@ -178,6 +182,34 @@ class OcdbtKvStoreWriteOptions: store_ocdbt_metadata_and_values_separately: bool = False +@dataclasses.dataclass(frozen=True) +class OcdbtTemporaryMetadataContext: + """Context for handling OCDBT temporary metadata. + + OCDBT kvstore configuration supports storing per-process OCDBT metadata + (manifest file and B-tree and version tree nodes) in a separate, local + temporary directory (backed by in-memory file system), which should later be + committed to the persistent metadata directory. This allows to achieve atomic + OCDBT metadata writes - especially for manifest files - without having to rely + on TensorStore transactions. + + Usage (within a single writer process): + 1) create a temporary directory and provide a OcdbtTemporaryMetadataContext + pointing to it to the TensorStore spec construction APIs (ArrayWriteSpec, + build_kvstore_tspec with WRITE mode) alongside the main persistent + directory + 2) write all process-local data to TensorStore + 3) after writing, call `ocdbt_utils.commit_temporary_ocdbt_metadata` to + atomically commit the metadata from the temporary directory to the + persistent directory + + Attributes: + path: The path to the temporary metadata directory. In-memory or local + filesystem are recommended for performance. + """ + path: epath.Path + + def _get_kvstore_for_gcs(ckpt_path: str) -> JsonSpec: """Constructs a TensorStore kvstore spec for a GCS path.""" m = re.fullmatch(_GCS_PATH_RE, ckpt_path, re.DOTALL) @@ -208,12 +240,106 @@ def _normalize_path(path: str) -> str: return os.path.normpath(path).replace('gs:/', 'gs://') +@dataclasses.dataclass(frozen=True) +class _OcdbtKvSpecParameters: + """OCDBT KvStore spec key parameters. + + Attributes: + base_driver_spec: The spec of the underlying (base) kvstore driver, pointing + to the target storage path (or using kvstack driver to support separate + storage of metadata in temporary path and values in the target path) + manifest_spec_override: [Optional] The manifest spec override of the + KvStore. + metadata_prefix_override: [Optional] The metadata prefix override of the + KvStore. If set, `btree_node_data_prefix` and + `version_tree_node_data_prefix` will be set to this value. + value_prefix_override: [Optional] The value prefix override of the KvStore. + If set, `value_data_prefix` will be set to this value. + """ + base_driver_spec: JsonSpec + manifest_spec_override: JsonSpec | str | None = None + metadata_prefix_override: str | None = None + value_prefix_override: str | None = None + + +def _override_ocdbt_kvspec_parameters_for_temporary_metadata( + temporary_metadata_context: OcdbtTemporaryMetadataContext | None, + write_mode: OcdbtWriteMode | None, + current_parameters: _OcdbtKvSpecParameters, +) -> _OcdbtKvSpecParameters: + """Returns KvStore spec parameters with overrides for temporary metadata.""" + if temporary_metadata_context is None: + if write_mode == OcdbtWriteMode.COMMIT_TEMPORARY: + raise ValueError( + 'OCDBT commit mode requires temporary metadata context.' + ) + return current_parameters + + if write_mode == OcdbtWriteMode.MERGE: + raise ValueError( + 'OCDBT merge mode does not support temporary metadata context.' + ) + + manifest_spec = current_parameters.manifest_spec_override + metadata_prefix = current_parameters.metadata_prefix_override + + # Ensure routing of metadata-related files' writes and reads to the temporary + # directory. We achieve this by: + # 1) using the kvstack driver + # 2) when in writing mode, overriding the metadata prefix to match the prefix + # of the layer backed by the temporary directory + # 3) overriding the manifest spec to point to the temporary metadata + # directory (unless in COMMIT_TEMPORARY mode) + # Notes on COMMIT_TEMPORARY mode (used for copying metadata from temporary + # to persistent location): + # 1) we don't set manifest or metadata prefix overrides: this ensures that + # the target kvstore is correctly opened as empty initially, and any + # writes of metadata are now routed to the layer backed by the persistent + # directory + # 2) kvstack driver's implementation of `experimental_copy_range_to` (used + # by `commit_temporary_ocdbt_metadata` to copy metadata to persistent + # location) is very strict about the base_driver_spec of the source + # and destination kvstores, requiring them to be identical. This defines + # how the base_driver_spec is constructed below, to look the same + # regardless of the mode used (read or commit). + if write_mode != OcdbtWriteMode.COMMIT_TEMPORARY: + manifest_spec = ( + f'file://{temporary_metadata_context.path}' + f'/{_OCDBT_TMP_METADATA_PREFIX}' + ) + if write_mode == OcdbtWriteMode.WRITE: + metadata_prefix = _OCDBT_TMP_METADATA_PREFIX + + base_driver_spec = { + 'driver': 'kvstack', + 'layers': [ + # Write to the real persistent checkpoint directory by default. + {'base': current_parameters.base_driver_spec}, + # Per-process metadata is stored in the separate local temporary + # directory. `prefix` ensures that writes and reads of + # metadata-related files are routed to the temporary directory. + { + 'prefix': _OCDBT_TMP_METADATA_PREFIX, + 'base': f'file://{temporary_metadata_context.path}/', + }, + ], + } + + return dataclasses.replace( + current_parameters, + base_driver_spec=base_driver_spec, + manifest_spec_override=manifest_spec, + metadata_prefix_override=metadata_prefix, + ) + + def _build_ocdbt_kvstore_tspec( directory: str, name: str | None = None, *, process_spec: OcdbtProcessSpec | None = None, write_options: OcdbtKvStoreWriteOptions | None = None, + temporary_metadata_context: OcdbtTemporaryMetadataContext | None = None, ) -> JsonSpec: """Constructs a spec for a Tensorstore OCDBT KvStore. @@ -225,6 +351,8 @@ def _build_ocdbt_kvstore_tspec( name). write_options: Options specific to OCDBT KvStore write modes. Should be provided when the kvstore will be used for writing or merging. + temporary_metadata_context: Context for local temporary metadata directory. + See `OcdbtTemporaryMetadataContext` for more details. Returns: A Tensorstore KvStore spec in dictionary form. @@ -242,7 +370,11 @@ def _build_ocdbt_kvstore_tspec( if is_gcs_path: base_driver_spec = _get_kvstore_for_gcs(directory) else: - base_driver_spec = {'driver': DEFAULT_DRIVER, 'path': str(directory)} + trailing_slash = '/' if temporary_metadata_context is not None else '' + base_driver_spec = { + 'driver': DEFAULT_DRIVER, + 'path': str(directory) + trailing_slash, + } # For OCDBT on local filesystems (including GCSFuse), we can safely use # non-atomic writes for data files to avoid expensive renames. However, @@ -259,31 +391,54 @@ def _build_ocdbt_kvstore_tspec( ) resolved_base_spec = base_driver_spec + kvspec_params = _OcdbtKvSpecParameters(base_driver_spec=base_driver_spec) + + if ( + write_options is not None + and write_options.store_ocdbt_metadata_and_values_separately + ): + kvspec_params = dataclasses.replace( + kvspec_params, + metadata_prefix_override=_OCDBT_SPLIT_META_DATA_PREFIX, + value_prefix_override=_OCDBT_SPLIT_VALUE_DATA_PREFIX, + ) + if ( isinstance(resolved_base_spec, dict) and resolved_base_spec.get('driver') == 'file' ): - kv_spec = { - 'driver': 'ocdbt', - 'base': { + kvspec_params = dataclasses.replace( + kvspec_params, + base_driver_spec={ **resolved_base_spec, 'file_io_locking': {'mode': 'non_atomic'}, }, - 'manifest': base_driver_spec, - } - else: - kv_spec = { - 'driver': 'ocdbt', - 'base': base_driver_spec, - } + manifest_spec_override=resolved_base_spec, + ) + + write_mode = None if write_options is None else write_options.mode + kvspec_params = _override_ocdbt_kvspec_parameters_for_temporary_metadata( + temporary_metadata_context=temporary_metadata_context, + write_mode=write_mode, + current_parameters=kvspec_params, + ) + + kv_spec = {'driver': 'ocdbt', 'base': kvspec_params.base_driver_spec} + + if kvspec_params.manifest_spec_override is not None: + kv_spec['manifest'] = kvspec_params.manifest_spec_override + if kvspec_params.metadata_prefix_override is not None: + kv_spec['btree_node_data_prefix'] = kvspec_params.metadata_prefix_override + kv_spec['version_tree_node_data_prefix'] = ( + kvspec_params.metadata_prefix_override + ) + if kvspec_params.value_prefix_override is not None: + kv_spec['value_data_prefix'] = kvspec_params.value_prefix_override if write_options is not None: _add_ocdbt_write_options( kv_spec, target_data_file_size=write_options.target_data_file_size, - store_ocdbt_metadata_and_values_separately=( - write_options.store_ocdbt_metadata_and_values_separately - ), ) if name is not None: @@ -334,6 +489,9 @@ def build_kvstore_tspec( use_ocdbt: bool = True, ocdbt_process_spec: OcdbtProcessSpec | None = None, ocdbt_write_options: OcdbtKvStoreWriteOptions | None = None, + ocdbt_temporary_metadata_context: ( + OcdbtTemporaryMetadataContext | None + ) = None, ) -> JsonSpec: """Constructs a spec for a Tensorstore KvStore. @@ -346,6 +504,8 @@ def build_kvstore_tspec( name). ocdbt_write_options: Options specific to OCDBT KvStore write modes. Should be provided when the kvstore will be used for writing or merging. + ocdbt_temporary_metadata_context: Context for local temporary metadata + directory. See `OcdbtTemporaryMetadataContext` for more details. Returns: A Tensorstore KvStore spec in dictionary form. @@ -356,6 +516,7 @@ def build_kvstore_tspec( name=name, process_spec=ocdbt_process_spec, write_options=ocdbt_write_options, + temporary_metadata_context=ocdbt_temporary_metadata_context, ) return _build_non_ocdbt_kvstore_tspec(directory=directory, name=name) @@ -387,8 +548,6 @@ def _get_backend_ocdbt_target_data_file_size( def _add_ocdbt_write_options( kvstore_tspec: JsonSpec, target_data_file_size: int | None = None, - *, - store_ocdbt_metadata_and_values_separately: bool = False, ) -> None: """Adds write-specific options to a TensorStore OCDBT KVStore spec.""" if target_data_file_size is None: @@ -403,13 +562,6 @@ def _add_ocdbt_write_options( ) kvstore_tspec['target_data_file_size'] = target_data_file_size - if store_ocdbt_metadata_and_values_separately: - kvstore_tspec['value_data_prefix'] = _OCDBT_SPLIT_VALUE_DATA_PREFIX - kvstore_tspec['btree_node_data_prefix'] = _OCDBT_SPLIT_META_DATA_PREFIX - kvstore_tspec['version_tree_node_data_prefix'] = ( - _OCDBT_SPLIT_META_DATA_PREFIX - ) - kvstore_tspec['config'] = { # Store .zarray metadata inline but not large chunks. # If separate storage for OCDBT metadata is enabled, this will mean that @@ -635,6 +787,9 @@ def __init__( replica_separate_folder: bool = False, ext_metadata: ExtMetadata | None = None, store_ocdbt_metadata_and_values_separately: bool = False, + ocdbt_temporary_metadata_context: ( + OcdbtTemporaryMetadataContext | None + ) = None, ): """Builds a TensorStore spec for writing an array.""" # Construct the underlying KvStore spec. @@ -656,6 +811,7 @@ def __init__( store_ocdbt_metadata_and_values_separately ), ), + ocdbt_temporary_metadata_context=ocdbt_temporary_metadata_context, ) # Construct the top-level array spec. tspec = { diff --git a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py index 160598890..9c66325af 100644 --- a/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py +++ b/checkpoint/orbax/checkpoint/_src/serialization/tensorstore_utils_test.py @@ -15,6 +15,7 @@ import functools import math import os +import tempfile import unittest from absl.testing import absltest @@ -1219,5 +1220,177 @@ def test_get_total_bytes_with_real_ops(self, use_compression): self.assertLess(bytes_read, 300000) +class BuildOcdbtKvStoreTspecWithTemporaryMetadataContextTest( + parameterized.TestCase +): + + def setUp(self): + super().setUp() + self.directory = self.create_tempdir().full_path + self.temporary_metadata_path = self.enter_context( + tempfile.TemporaryDirectory() + ) + self.temporary_metadata_context = ts_utils.OcdbtTemporaryMetadataContext( + path=epath.Path(self.temporary_metadata_path), + ) + + def _verify_kvstack_spec( + self, + actual_spec: ts_utils.JsonSpec, + expected_base_path: str, + ) -> None: + self.assertDictEqual( + actual_spec, + { + 'driver': 'kvstack', + 'layers': [ + { + 'base': { + 'driver': ts_utils.DEFAULT_DRIVER, + 'path': expected_base_path, + } + }, + { + 'prefix': 'ocdbt_tmp_meta/', + 'base': f'file://{self.temporary_metadata_path}/', + }, + ], + }, + ) + + @parameterized.product(use_process_spec=(True, False)) + def test_read_mode(self, use_process_spec: bool): + process_spec = None + if use_process_spec: + process_spec = ts_utils.OcdbtProcessSpec(process_id='w13') + + kvstore_tspec = ts_utils.build_kvstore_tspec( + directory=self.directory, + use_ocdbt=True, + ocdbt_process_spec=process_spec, + ocdbt_temporary_metadata_context=self.temporary_metadata_context, + ) + + # Manifest spec is set to the temporary metadata path, but prefixes are no + # longer overridden. + self.assertEqual( + kvstore_tspec['manifest'], + f'file://{self.temporary_metadata_path}/ocdbt_tmp_meta/', + ) + self.assertNotIn('btree_node_data_prefix', kvstore_tspec) + self.assertNotIn('version_tree_node_data_prefix', kvstore_tspec) + self.assertNotIn('value_data_prefix', kvstore_tspec) + + expected_base_path = ( + f'{self.directory}/' + if not use_process_spec + else f'{self.directory}/{str(process_spec)}/' + ) + # kvstack spec is independent of the mode. + self._verify_kvstack_spec(kvstore_tspec['base'], expected_base_path) + + @parameterized.product( + use_process_spec=(True, False), + store_ocdbt_metadata_and_values_separately=(True, False), + ) + def test_write_mode( + self, + use_process_spec: bool, + store_ocdbt_metadata_and_values_separately: bool, + ): + process_spec = None + if use_process_spec: + process_spec = ts_utils.OcdbtProcessSpec(process_id='w13') + + kvstore_tspec = ts_utils.build_kvstore_tspec( + directory=self.directory, + use_ocdbt=True, + ocdbt_process_spec=process_spec, + ocdbt_write_options=ts_utils.OcdbtKvStoreWriteOptions( + mode=ts_utils.OcdbtWriteMode.WRITE, + store_ocdbt_metadata_and_values_separately=( + store_ocdbt_metadata_and_values_separately + ), + ), + ocdbt_temporary_metadata_context=self.temporary_metadata_context, + ) + + # Manifest spec is set to the temporary metadata path, and prefixes are + # overridden to match the routing prefix of kvstack layer configured with + # temporary metadata path. + self.assertEqual( + kvstore_tspec['manifest'], + f'file://{self.temporary_metadata_path}/ocdbt_tmp_meta/', + ) + self.assertEqual( + kvstore_tspec['btree_node_data_prefix'], 'ocdbt_tmp_meta/' + ) + self.assertEqual( + kvstore_tspec['version_tree_node_data_prefix'], 'ocdbt_tmp_meta/' + ) + + if store_ocdbt_metadata_and_values_separately: + self.assertEqual(kvstore_tspec['value_data_prefix'], 'ocdbt_data/') + else: + self.assertNotIn('value_data_prefix', kvstore_tspec) + + expected_base_path = ( + f'{self.directory}/' + if not use_process_spec + else f'{self.directory}/{str(process_spec)}/' + ) + # kvstack spec is independent of the mode and write options. + self._verify_kvstack_spec(kvstore_tspec['base'], expected_base_path) + + @parameterized.product( + use_process_spec=(True, False), + store_ocdbt_metadata_and_values_separately=(True, False), + ) + def test_commit_temporary_metadata_mode( + self, + use_process_spec: bool, + store_ocdbt_metadata_and_values_separately: bool, + ): + process_spec = None + if use_process_spec: + process_spec = ts_utils.OcdbtProcessSpec(process_id='w13') + + kvstore_tspec = ts_utils.build_kvstore_tspec( + directory=self.directory, + use_ocdbt=True, + ocdbt_process_spec=process_spec, + ocdbt_write_options=ts_utils.OcdbtKvStoreWriteOptions( + mode=ts_utils.OcdbtWriteMode.COMMIT_TEMPORARY, + store_ocdbt_metadata_and_values_separately=( + store_ocdbt_metadata_and_values_separately + ), + ), + ocdbt_temporary_metadata_context=self.temporary_metadata_context, + ) + + # Manifest spec is not overridden. + self.assertNotIn('manifest', kvstore_tspec) + + # Prefixes are set according to store_ocdbt_metadata_and_values_separately. + if store_ocdbt_metadata_and_values_separately: + self.assertEqual(kvstore_tspec['value_data_prefix'], 'ocdbt_data/') + self.assertEqual(kvstore_tspec['btree_node_data_prefix'], 'ocdbt_meta/') + self.assertEqual( + kvstore_tspec['version_tree_node_data_prefix'], 'ocdbt_meta/' + ) + else: + self.assertNotIn('value_data_prefix', kvstore_tspec) + self.assertNotIn('btree_node_data_prefix', kvstore_tspec) + self.assertNotIn('version_tree_node_data_prefix', kvstore_tspec) + + expected_base_path = ( + f'{self.directory}/' + if not use_process_spec + else f'{self.directory}/{str(process_spec)}/' + ) + # kvstack spec is independent of the mode and write options. + self._verify_kvstack_spec(kvstore_tspec['base'], expected_base_path) + + if __name__ == '__main__': absltest.main()