[store] add async api (#1265)

* [store] add async api

Signed-off-by: Xuchun Shang <xuchun.shang@gmail.com>
This commit is contained in:
Xuchun Shang 2025-12-25 16:46:14 +08:00 committed by GitHub
parent fdfd74818b
commit f7f65aa140
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
6 changed files with 658 additions and 32 deletions

View File

@ -241,6 +241,17 @@ jobs:
python scripts/test_tensor_api.py -n 1
shell: bash
- name: Run Python Async API Test (CI check)
env:
MOONCAKE_MASTER: "127.0.0.1:50051"
MOONCAKE_TE_META_DATA_SERVER: "http://127.0.0.1:8080/metadata"
MOONCAKE_PROTOCOL: "tcp"
LOCAL_HOSTNAME: "127.0.0.1"
run: |
source test_env/bin/activate
python scripts/test_async_store.py
shell: bash
- name: Run RPC Communicator Bandwidth Test
run: |
source test_env/bin/activate

View File

@ -111,6 +111,8 @@ if (USE_MNNVL)
)
endif()
install(FILES "${CMAKE_CURRENT_SOURCE_DIR}/store/async_store.py" DESTINATION ${PYTHON_SYS_PATH}/${PYTHON_PACKAGE_NAME})
# Install Python scripts from mooncake-wheel/mooncake/ directory
install(FILES
"${CMAKE_CURRENT_SOURCE_DIR}/../mooncake-wheel/mooncake/http_metadata_server.py"

View File

@ -0,0 +1,32 @@
import asyncio
import functools
from mooncake.store import MooncakeDistributedStore
class MooncakeDistributedStoreAsync(MooncakeDistributedStore):
def __getattr__(self, name: str):
if not name.startswith("async_"):
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}'")
sync_method_name = name[6:]
if not hasattr(self, sync_method_name):
raise AttributeError(f"'{type(self).__name__}' object has no attribute '{name}' (nor '{sync_method_name}')")
sync_method = getattr(self, sync_method_name)
if not callable(sync_method):
raise AttributeError(f"'{sync_method_name}' is not callable")
async_method = self._make_async_wrapper(sync_method)
setattr(self, name, async_method)
return async_method
def _make_async_wrapper(self, sync_method):
@functools.wraps(sync_method)
async def wrapper(*args, **kwargs):
loop = asyncio.get_running_loop()
func = functools.partial(sync_method, *args, **kwargs)
return await loop.run_in_executor(None, func)
return wrapper

View File

@ -34,6 +34,8 @@ if [ -f build/mooncake-integration/store.*.so ]; then
cp build/mooncake-store/src/mooncake_master mooncake-wheel/mooncake/
# Copy client binary
cp build/mooncake-store/src/mooncake_client mooncake-wheel/mooncake/
# Copy async_store.py
cp mooncake-integration/store/async_store.py mooncake-wheel/mooncake/async_store.py
else
echo "Skipping store.so (not built - likely WITH_STORE is set to OFF)"
fi

608
scripts/test_async_store.py Normal file
View File

@ -0,0 +1,608 @@
import ctypes
import os
import sys
import time
import argparse
import unittest
import torch
import numpy as np
import asyncio
from mooncake.mooncake_config import MooncakeConfig
from dataclasses import dataclass
try:
from mooncake.async_store import MooncakeDistributedStoreAsync
except ImportError:
print("Warning: Could not import MooncakeDistributedStoreAsync from async_store.py")
MooncakeDistributedStoreAsync = None
# ==========================================
# Global Variables & Configuration
# ==========================================
# Global Store instance
GLOBAL_STORE = None
GLOBAL_CONFIG = None
DEFAULT_MOONCAKE_CONFIG_PATH_ENV = "MOONCAKE_CONFIG_PATH"
DEFAULT_GLOBAL_SEGMENT_SIZE = 16 * 1024 * 1024 * 1024 # 16 GiB
DEFAULT_LOCAL_BUFFER_SIZE = 8 * 1024 * 1024 * 1024 # 8 GB
DEFAULT_MASTER_METRICS_PORT = 9003
DEFAULT_CHECK_SERVER = False
def verify_tensor_equality(original, received, rtol=0, atol=0, verbose=True):
"""
Utility to compare two tensors (CPU/GPU/Numpy).
Identical to the synchronous version.
"""
def to_numpy(x):
if isinstance(x, torch.Tensor):
if x.is_cuda:
x = x.cpu()
return x.detach().numpy()
elif isinstance(x, np.ndarray):
return x
else:
raise TypeError(f"Unsupported tensor type: {type(x)}")
try:
orig_np = to_numpy(original)
recv_np = to_numpy(received)
except Exception as e:
if verbose:
print(f"❌ Error converting tensors: {e}")
return False
if orig_np.shape != recv_np.shape:
if verbose:
print(f"❌ Shape mismatch: original {orig_np.shape} vs received {recv_np.shape}")
return False
if orig_np.dtype != recv_np.dtype:
if verbose:
print(f"❌ Dtype mismatch: original {orig_np.dtype} vs received {recv_np.dtype}")
return False
if np.array_equal(orig_np, recv_np):
return True
else:
# Simple diff reporting
diff_mask = orig_np != recv_np
diff_indices = np.where(diff_mask)
if len(diff_indices[0]) > 0:
first_diff_idx = tuple(idx[0] for idx in diff_indices)
orig_val = orig_np[first_diff_idx]
recv_val = recv_np[first_diff_idx]
if verbose:
print(f"❌ Tensors differ at index {first_diff_idx}")
print(f" Original: {orig_val}")
print(f" Received: {recv_val}")
return False
def parse_global_segment_size(value) -> int:
if isinstance(value, int): return value
if isinstance(value, str):
s = value.strip().lower()
if s.endswith("gb"):
return int(s[:-2].strip()) * 1024**3
return int(s)
return int(value)
def generate_tensors(num_tensors, size_mb):
size_bytes = int(size_mb * 1024 * 1024)
element_size = 4 # float32
num_elements = size_bytes // element_size
dim = int(np.sqrt(num_elements))
dim = (dim // 8) * 8
tensors = [torch.randn(dim, dim, dtype=torch.float32).contiguous() for _ in range(num_tensors)]
keys = [f"test_tensor_{i}_{int(time.time()*1000)}" for i in range(num_tensors)]
return keys, tensors
# ==========================================
# Module Level Setup/Teardown
# ==========================================
def setUpModule():
"""
Initialize the Async wrapper globally.
Note: We invoke setup() synchronously via asyncio.run() to ensure C++ client is ready
before any tests run. This mimics the sync script's behavior.
"""
global GLOBAL_STORE, GLOBAL_CONFIG
if MooncakeDistributedStoreAsync is None:
raise ImportError("Could not find MooncakeDistributedStoreAsync class.")
try:
GLOBAL_CONFIG = MooncakeConfig.load_from_env()
# Increase max_workers to simulate high async concurrency
GLOBAL_STORE = MooncakeDistributedStoreAsync()
print(f"[{os.getpid()}] (Async) Connecting to Mooncake Master at {GLOBAL_CONFIG.master_server_address}...")
def _do_setup():
return GLOBAL_STORE.setup(
GLOBAL_CONFIG.local_hostname,
GLOBAL_CONFIG.metadata_server,
GLOBAL_CONFIG.global_segment_size,
GLOBAL_CONFIG.local_buffer_size,
GLOBAL_CONFIG.protocol,
GLOBAL_CONFIG.device_name,
GLOBAL_CONFIG.master_server_address,
)
# Run setup in a temporary loop just for initialization
rc = _do_setup()
if rc != 0:
raise RuntimeError(f"Failed to setup mooncake store, error code: {rc}")
print("✅ Global Async Store connection established.")
except Exception as e:
print(f"❌ Failed to establish global store connection: {e}")
sys.exit(1)
def tearDownModule():
global GLOBAL_STORE
if GLOBAL_STORE:
print("\nClosing global store connection...")
try:
GLOBAL_STORE.close()
except Exception as e:
print(f"Error during close: {e}")
finally:
GLOBAL_STORE = None
# ==========================================
# Base Test Class (Async)
# ==========================================
class MooncakeAsyncTestBase(unittest.IsolatedAsyncioTestCase):
async def asyncSetUp(self):
"""Executed before each test method."""
if GLOBAL_STORE is None:
self.skipTest("Store not initialized")
self.store = GLOBAL_STORE
self.config = GLOBAL_CONFIG
# Asynchronously clear the store
# Note: Depending on implementation, you might want to call
# await self.store.async_remove_all()
# For now, we assume remove_all wraps C++ synchronous logic nicely in a thread.
await self.store.async_remove_all()
# ==========================================
# Basic Bytes I/O & Auxiliary Tests
# ==========================================
class TestMooncakeBasicFunctional(MooncakeAsyncTestBase):
async def test_01_bytes_put_get(self):
"""Test basic single key-value (bytes) storage."""
key = "test_bytes_single"
value = b"Hello Mooncake Async World!"
# 1. Put Bytes
# Note: put wrapper in Python usually defaults config if None
rc = await self.store.async_put(key, value)
self.assertEqual(rc, 0, f"put failed with rc={rc}")
# 2. Get Bytes
retrieved = await self.store.async_get(key)
self.assertIsInstance(retrieved, bytes)
self.assertEqual(retrieved, value, "Retrieved bytes mismatch")
async def test_02_bytes_batch_io(self):
"""Test batch put and get for bytes."""
count = 10
keys = [f"test_bytes_batch_{i}" for i in range(count)]
# Generate different length values
values = [f"value_{i}_{'x'*i}".encode('utf-8') for i in range(count)]
# 1. Batch Put
rcs = await self.store.async_put_batch(keys, values)
self.assertEqual(rcs, 0, f"batch_put failed with rc={rcs}")
# 2. Batch Get
retrieved_values = await self.store.async_get_batch(keys)
self.assertEqual(len(retrieved_values), count)
for i in range(count):
self.assertEqual(retrieved_values[i], values[i], f"Mismatch at index {i}")
async def test_03_existence_check(self):
"""Test is_exist and batch_is_exist."""
key_exist = "test_exist_yes"
key_missing = "test_exist_no_such_key"
await self.store.async_put(key_exist, b"exist")
# Single check
self.assertTrue(await self.store.async_is_exist(key_exist))
self.assertFalse(await self.store.async_is_exist(key_missing))
# Batch check
keys = [key_exist, key_missing, key_exist]
# C++ batchIsExist returns list of ints: 1=exist, 0=not exist, -1=error
results = await self.store.async_batch_is_exist(keys)
self.assertEqual(len(results), 3)
self.assertEqual(results[0], 1)
self.assertEqual(results[1], 0)
self.assertEqual(results[2], 1)
async def test_04_remove(self):
"""Test remove and remove_all."""
key = "test_remove_key"
await self.store.async_put(key, b"data")
self.assertTrue(await self.store.async_is_exist(key))
await asyncio.sleep(6)
# Remove single
await self.store.async_remove(key)
self.assertFalse(await self.store.async_is_exist(key))
# Test Remove All (implicitly tested in setup, but explicit here)
keys = ["rm_all_1", "rm_all_2"]
await self.store.async_put_batch(keys, [b"1", b"2"])
self.assertTrue(await self.store.async_is_exist(keys[0]))
await asyncio.sleep(6)
await self.store.async_remove_all()
self.assertFalse(await self.store.async_is_exist(keys[0]))
self.assertFalse(await self.store.async_is_exist(keys[1]))
async def test_05_remove_by_regex(self):
"""Test remove_by_regex."""
# Setup keys with specific patterns
prefix_keys = [f"regex_target_{i}" for i in range(5)]
other_keys = [f"regex_safe_{i}" for i in range(5)]
all_keys = prefix_keys + other_keys
all_vals = [b"x"] * len(all_keys)
await self.store.async_put_batch(all_keys, all_vals)
# Verify all exist
results = await self.store.async_batch_is_exist(all_keys)
self.assertTrue(all(r == 1 for r in results))
# Remove by regex pattern "regex_target_.*"
pattern = "regex_target_.*"
await asyncio.sleep(6)
await self.store.async_remove_by_regex(pattern)
# Verify targets are gone
target_res = await self.store.async_batch_is_exist(prefix_keys)
self.assertTrue(all(r == 0 for r in target_res), "Targets should be removed")
# Verify others stay
safe_res = await self.store.async_batch_is_exist(other_keys)
self.assertTrue(all(r == 1 for r in safe_res), "Safe keys should remain")
async def test_07_large_bytes_io(self):
"""Test larger bytes payload (non-tensor path)."""
key = "test_large_bytes"
size = 10 * 1024 * 1024 # 10MB
# Create random bytes
data = os.urandom(size)
await self.store.async_put(key, data)
# Retrieve
retrieved = await self.store.async_get(key)
self.assertEqual(len(retrieved), size)
self.assertEqual(retrieved, data)
# ==========================================
# Tensor Functional Tests
# ==========================================
class TestMooncakeTensorFunctional(MooncakeAsyncTestBase):
async def test_01_basic_put_get(self):
"""Verify basic put and get functionality."""
key = "func_test_single"
tensor = torch.randn(1024, 1024, dtype=torch.float32)
# Await Put
rc = await self.store.async_put_tensor(key, tensor)
self.assertEqual(rc, 0, f"put_tensor failed with rc={rc}")
exists = await self.store.async_is_exist(key)
self.assertTrue(exists, "Key not found after put")
# Await Get
retrieved = await self.store.async_get_tensor(key)
self.assertIsNotNone(retrieved, "Get returned None")
self.assertTrue(torch.equal(tensor, retrieved), "Data mismatch")
async def test_02_tp_single_tensor(self):
tp_size = 4
split_dim = 1
key = "func_test_tp_single"
_, tensors = generate_tensors(1, 16)
target_tensor = tensors[0]
# 1. Put with TP
rc = await self.store.async_put_tensor_with_tp(key, target_tensor, tp_size=tp_size, split_dim=split_dim)
self.assertEqual(rc, 0, "put_tensor_with_tp failed")
# 2. Verify existence of shards
for rank in range(tp_size):
shard_key = f"{key}_tp_{rank}"
exists = await self.store.async_is_exist(shard_key)
self.assertTrue(exists, f"Shard key {shard_key} is missing")
# 3. Get shards and Reconstruct (Concurrently!)
expected_chunks = target_tensor.chunk(tp_size, split_dim)
# Launch all gets in parallel
tasks = []
for rank in range(tp_size):
tasks.append(self.store.async_get_tensor_with_tp(key, tp_rank=rank, tp_size=tp_size, split_dim=split_dim))
slices = await asyncio.gather(*tasks)
for rank, t_slice in enumerate(slices):
self.assertIsNotNone(t_slice, f"Slice for rank {rank} is None")
self.assertTrue(torch.equal(t_slice, expected_chunks[rank]), f"Data mismatch for rank {rank}")
reconstructed = torch.cat(slices, dim=split_dim)
self.assertTrue(torch.equal(reconstructed, target_tensor), "Reconstruction mismatch")
async def test_03_tp_batch(self):
tp_size = 2
split_dim = 0
num_tensors = 4
keys, tensors = generate_tensors(num_tensors, 8)
# 1. Batch Put with TP
results = await self.store.async_batch_put_tensor_with_tp(keys, tensors, tp_size=tp_size, split_dim=split_dim)
self.assertTrue(all(r == 0 for r in results), f"Batch put failed. Results: {results}")
# 2. Batch Get per Rank (Concurrently fetch both ranks)
tasks = []
for rank in range(tp_size):
tasks.append(self.store.async_batch_get_tensor_with_tp(keys, tp_rank=rank, tp_size=tp_size))
all_shards = await asyncio.gather(*tasks) # List of lists
# 3. Verify & Reconstruct
for i in range(num_tensors):
original = tensors[i]
expected_chunks = original.chunk(tp_size, split_dim)
reconstruction_parts = []
for rank in range(tp_size):
shard = all_shards[rank][i]
self.assertTrue(torch.equal(shard, expected_chunks[rank]),
f"Tensor {i} Rank {rank} data mismatch")
reconstruction_parts.append(shard)
recon = torch.cat(reconstruction_parts, dim=split_dim)
self.assertTrue(torch.equal(recon, original), f"Tensor {i} final mismatch")
async def test_04_tp_consistency(self):
input_tensor = torch.arange(12).view(3, 4)
tp_size = 2
await self.store.async_batch_put_tensor_with_tp(['key'], [input_tensor], tp_size=tp_size, split_dim=1)
chunked_tensors = input_tensor.chunk(chunks=2, dim=1)
# Parallel fetch
t0, t1 = await asyncio.gather(
self.store.async_batch_get_tensor_with_tp(['key'], tp_rank=0, tp_size=tp_size),
self.store.async_batch_get_tensor_with_tp(['key'], tp_rank=1, tp_size=tp_size)
)
tmp_tensor_0 = t0[0]
tmp_tensor_1 = t1[0]
self.assertTrue(tmp_tensor_0.sum() == chunked_tensors[0].sum())
self.assertTrue(tmp_tensor_1.sum() == chunked_tensors[1].sum())
buffer_spacing = 1 * 1024 * 1024
buffer_2 = (ctypes.c_ubyte * buffer_spacing)()
buffer_3 = (ctypes.c_ubyte * buffer_spacing)()
buffer_ptr_2 = ctypes.addressof(buffer_2)
buffer_ptr_3 = ctypes.addressof(buffer_3)
# Important: register_buffer needs to be supported in Async wrapper
# Assuming Async wrapper has register_buffer that delegates to C++
res = await self.store.async_register_buffer(buffer_ptr_2, buffer_spacing)
self.assertEqual(res, 0)
res = await self.store.async_register_buffer(buffer_ptr_3, buffer_spacing)
self.assertEqual(res, 0)
# Parallel get_into
t2_task = self.store.async_batch_get_tensor_with_tp_into(
['key'], [buffer_ptr_2], [buffer_spacing], tp_rank=0, tp_size=tp_size)
t3_task = self.store.async_batch_get_tensor_with_tp_into(
['key'], [buffer_ptr_3], [buffer_spacing], tp_rank=1, tp_size=tp_size)
t2_res, t3_res = await asyncio.gather(t2_task, t3_task)
tmp_tensor_2 = t2_res[0]
tmp_tensor_3 = t3_res[0]
self.assertTrue(tmp_tensor_2.sum() == chunked_tensors[0].sum())
self.assertTrue(tmp_tensor_3.sum() == chunked_tensors[1].sum())
await self.store.async_unregister_buffer(buffer_ptr_2)
await self.store.async_unregister_buffer(buffer_ptr_3)
async def test_05_put_get_into(self):
key = "get_into_test"
tensor = torch.randn(1024, 1024, dtype=torch.float32)
total_buffer_size = 64 * 1024 * 1024
large_buffer = (ctypes.c_ubyte * total_buffer_size)()
large_buffer_ptr = ctypes.addressof(large_buffer)
await self.store.async_register_buffer(large_buffer_ptr, total_buffer_size)
rc = await self.store.async_put_tensor(key, tensor)
self.assertEqual(rc, 0)
retrieved = await self.store.async_get_tensor_into(key, large_buffer_ptr, total_buffer_size)
self.assertIsNotNone(retrieved)
self.assertTrue(torch.equal(tensor, retrieved))
await self.store.async_unregister_buffer(large_buffer_ptr)
async def test_06_batch_put_get_into(self):
num_tensors = 4
keys, tensors = generate_tensors(num_tensors, 8)
buffer_spacing = 64 * 1024 * 1024
batch_size = len(keys)
total_buffer_size = buffer_spacing * batch_size
large_buffer = (ctypes.c_ubyte * total_buffer_size)()
large_buffer_ptr = ctypes.addressof(large_buffer)
buffer_ptrs = []
buffer_sizes = []
for i in range(batch_size):
offset = i * buffer_spacing
buffer_ptrs.append(large_buffer_ptr + offset)
buffer_sizes.append(buffer_spacing)
await self.store.async_register_buffer(large_buffer_ptr, total_buffer_size)
results = await self.store.async_batch_put_tensor(keys, tensors)
self.assertTrue(all(r == 0 for r in results))
res = await self.store.async_batch_get_tensor_into(keys, buffer_ptrs, buffer_sizes)
self.assertEqual(len(res), len(tensors))
for j in range(batch_size):
self.assertTrue(verify_tensor_equality(tensors[j], res[j]))
await self.store.async_unregister_buffer(large_buffer_ptr)
async def test_07_put_get_into_with_tp(self):
tp_size = 4
split_dim = 0
key = "get_into_with_tp_test"
tensor = torch.randn(1024, 1024, dtype=torch.float32)
result = await self.store.async_put_tensor_with_tp(key, tensor, tp_size=tp_size, split_dim=split_dim)
self.assertEqual(result, 0)
all_shards = []
registered_buffers = []
# We can launch these concurrently too!
async def _fetch_rank(rank):
buf_size = 64 * 1024 * 1024
l_buf = (ctypes.c_ubyte * buf_size)()
l_ptr = ctypes.addressof(l_buf)
await self.store.async_register_buffer(l_ptr, buf_size)
shard = await self.store.async_get_tensor_with_tp_into(
key, l_ptr, buf_size, tp_rank=rank, tp_size=tp_size
)
return rank, shard, l_buf, l_ptr
tasks = [_fetch_rank(r) for r in range(tp_size)]
results = await asyncio.gather(*tasks)
# Sort by rank to ensure order
results.sort(key=lambda x: x[0])
for _, shard, buf, ptr in results:
all_shards.append(shard)
registered_buffers.append((buf, ptr)) # Keep buf alive
# Validate
original = tensor
expected_chunks = original.chunk(tp_size, split_dim)
reconstruction_parts = []
for rank in range(tp_size):
self.assertTrue(torch.equal(all_shards[rank], expected_chunks[rank]))
reconstruction_parts.append(all_shards[rank])
recon = torch.cat(reconstruction_parts, dim=split_dim)
self.assertTrue(torch.equal(recon, original))
# Cleanup
for _, ptr in registered_buffers:
await self.store.async_unregister_buffer(ptr)
async def test_08_batch_put_get_into_with_tp(self):
# Similar logic to test_07 but for batch
tp_size = 4
split_dim = 0
num_tensors = 4
keys, tensors = generate_tensors(num_tensors, 8)
results = await self.store.async_batch_put_tensor_with_tp(keys, tensors, tp_size=tp_size, split_dim=split_dim)
self.assertTrue(all(r == 0 for r in results))
async def _fetch_batch_rank(rank):
batch_size = len(keys)
spacing = 64 * 1024 * 1024
total_size = spacing * batch_size
l_buf = (ctypes.c_ubyte * total_size)()
l_ptr = ctypes.addressof(l_buf)
ptrs = [l_ptr + i*spacing for i in range(batch_size)]
sizes = [spacing] * batch_size
await self.store.async_register_buffer(l_ptr, total_size)
shards = await self.store.async_batch_get_tensor_with_tp_into(
keys, ptrs, sizes, tp_rank=rank, tp_size=tp_size
)
return rank, shards, l_buf, l_ptr
tasks = [_fetch_batch_rank(r) for r in range(tp_size)]
results = await asyncio.gather(*tasks)
results.sort(key=lambda x: x[0])
all_shards_by_rank = [r[1] for r in results]
registered_buffers = [(r[2], r[3]) for r in results]
# Verify
for i in range(num_tensors):
original = tensors[i]
expected = original.chunk(tp_size, split_dim)
parts = []
for rank in range(tp_size):
shard = all_shards_by_rank[rank][i]
self.assertTrue(torch.equal(shard, expected[rank]))
parts.append(shard)
recon = torch.cat(parts, dim=split_dim)
self.assertTrue(torch.equal(recon, original))
for _, ptr in registered_buffers:
await self.store.async_unregister_buffer(ptr)
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Async Mooncake Store Tests")
parser.add_argument("--mode", type=str, default="all", choices=["all", "Tensorfunc", "Basicfunc"])
parser.add_argument("--threads", type=int, default=8, help="Number of concurrent tasks")
parser.add_argument("--count", type=int, default=800, help="Total items")
parser.add_argument("--size_mb", type=float, default=0.5, help="Tensor MB")
args, unknown = parser.parse_known_args()
suite = unittest.TestSuite()
loader = unittest.TestLoader()
# Note: IsolatedAsyncioTestCase handles the async loop internally for each test
if args.mode in ["all", "Tensorfunc"]:
suite.addTests(loader.loadTestsFromTestCase(TestMooncakeTensorFunctional))
if args.mode in ["all", "Basicfunc"]:
suite.addTests(loader.loadTestsFromTestCase(TestMooncakeBasicFunctional))
runner = unittest.TextTestRunner(verbosity=2)
result = runner.run(suite)
sys.exit(not result.wasSuccessful())

View File

@ -10,6 +10,8 @@ import numpy as np
from dataclasses import dataclass
from mooncake.store import MooncakeDistributedStore
from mooncake.mooncake_config import MooncakeConfig
import concurrent.futures
# ==========================================
@ -86,41 +88,10 @@ def parse_global_segment_size(value) -> int:
return int(s)
return int(value)
@dataclass
class MooncakeStoreConfig:
local_hostname: str
metadata_server: str
global_segment_size: int
local_buffer_size: int
protocol: str
device_name: str
master_server_address: str
master_metrics_port: int
check_server: bool
@staticmethod
def load_from_env() -> "MooncakeStoreConfig":
"""Load configuration from environment variables."""
if not os.getenv("MOONCAKE_MASTER"):
raise ValueError("Environment variable 'MOONCAKE_MASTER' is not set.")
return MooncakeStoreConfig(
local_hostname=os.getenv("LOCAL_HOSTNAME", "localhost"),
metadata_server=os.getenv("MOONCAKE_TE_META_DATA_SERVER", "P2PHANDSHAKE"),
global_segment_size=parse_global_segment_size(
os.getenv("MOONCAKE_GLOBAL_SEGMENT_SIZE", DEFAULT_GLOBAL_SEGMENT_SIZE)
),
local_buffer_size=DEFAULT_LOCAL_BUFFER_SIZE,
protocol=os.getenv("MOONCAKE_PROTOCOL", "tcp"),
device_name=os.getenv("MOONCAKE_DEVICE", ""),
master_server_address=os.getenv("MOONCAKE_MASTER"),
master_metrics_port=int(os.getenv("MOONCAKE_MASTER_METRICS_PORT", DEFAULT_MASTER_METRICS_PORT)),
check_server=bool(os.getenv("MOONCAKE_CHECK_SERVER", DEFAULT_CHECK_SERVER)),
)
def create_store_connection():
"""Create and connect to the Store (called only once by setUpModule)."""
store = MooncakeDistributedStore()
config = MooncakeStoreConfig.load_from_env()
config = MooncakeConfig.load_from_env()
print(f"[{os.getpid()}] Connecting to Mooncake Master at {config.master_server_address} using {config.protocol}...")
rc = store.setup(
config.local_hostname,