440 lines
15 KiB
Python
440 lines
15 KiB
Python
# SPDX-License-Identifier: Apache-2.0
|
|
# Standard
|
|
from typing import TYPE_CHECKING, Any
|
|
|
|
# First Party
|
|
from lmcache.logging import init_logger
|
|
from lmcache.v1.cache_controller.controllers.full_sync_tracker import FullSyncTracker
|
|
from lmcache.v1.cache_controller.message import (
|
|
BatchedKVOperationMsg,
|
|
BatchedP2PLookupMsg,
|
|
BatchedP2PLookupRetMsg,
|
|
CheckFinishMsg,
|
|
CheckFinishRetMsg,
|
|
ClearMsg,
|
|
ClearRetMsg,
|
|
CompressMsg,
|
|
CompressRetMsg,
|
|
DecompressMsg,
|
|
DecompressRetMsg,
|
|
FullSyncBatchMsg,
|
|
FullSyncEndMsg,
|
|
FullSyncStartMsg,
|
|
FullSyncStartRetMsg,
|
|
FullSyncStatusMsg,
|
|
FullSyncStatusRetMsg,
|
|
KVOpEvent,
|
|
LookupMsg,
|
|
LookupRetMsg,
|
|
MoveMsg,
|
|
MoveRetMsg,
|
|
OpType,
|
|
PinMsg,
|
|
PinRetMsg,
|
|
)
|
|
from lmcache.v1.cache_controller.observability import PrometheusLogger
|
|
from lmcache.v1.cache_controller.utils import RegistryTree
|
|
from lmcache.v1.token_database import ChunkedTokenDatabase
|
|
|
|
if TYPE_CHECKING:
|
|
# First Party
|
|
from lmcache.v1.cache_controller.controllers import RegistrationController
|
|
|
|
logger = init_logger(__name__)
|
|
|
|
|
|
"""
|
|
The kv controller use `(instance_id, worker_id)` -> [location -> set[chunk_hash]]
|
|
as kv_pool. When the number of instances is small and stable, the time complexity
|
|
of `lookup` in kv controller is O(n). If the number of instance is large or unknown,
|
|
the time complexity will degrade to O(n^2), and the ReverseIndexKVController is a
|
|
better choice.
|
|
"""
|
|
|
|
|
|
class KVController:
|
|
def __init__(
|
|
self,
|
|
registry: RegistryTree,
|
|
full_sync_completion_threshold: float = 0.8,
|
|
full_sync_timeout_s: float = 300.0,
|
|
) -> None:
|
|
# TODO(Jiayi): remove this hardcode
|
|
self.token_database = ChunkedTokenDatabase()
|
|
self.registry = registry
|
|
self.cluster_executor: Any = None
|
|
|
|
# Full sync tracker
|
|
self.full_sync_tracker = FullSyncTracker(
|
|
registry_tree=registry,
|
|
completion_threshold=full_sync_completion_threshold,
|
|
sync_timeout_s=full_sync_timeout_s,
|
|
)
|
|
|
|
def _setup_metrics(self) -> None:
|
|
prometheus_logger = PrometheusLogger.GetInstanceOrNone()
|
|
if prometheus_logger is not None:
|
|
prometheus_logger.kv_pool_keys_count.set_function(
|
|
self.registry.get_total_kv_count
|
|
)
|
|
prometheus_logger.kv_op_seq_discontinuity_count.set_function(
|
|
self.registry.get_seq_discontinuity_count
|
|
)
|
|
# Full sync metrics
|
|
prometheus_logger.full_sync_workers_syncing.set_function(
|
|
self.full_sync_tracker.get_syncing_count
|
|
)
|
|
prometheus_logger.full_sync_workers_completed.set_function(
|
|
self.full_sync_tracker.get_completed_count
|
|
)
|
|
prometheus_logger.full_sync_global_progress.set_function(
|
|
self.full_sync_tracker.get_global_progress
|
|
)
|
|
prometheus_logger.full_sync_missing_batches_total.set_function(
|
|
self.full_sync_tracker.get_total_missing_batches_count
|
|
)
|
|
|
|
def post_init(
|
|
self, reg_controller: "RegistrationController", cluster_executor: Any
|
|
) -> None:
|
|
"""
|
|
Post initialization of the KV controller.
|
|
"""
|
|
self.reg_controller = reg_controller
|
|
self.cluster_executor = cluster_executor
|
|
self._setup_metrics()
|
|
|
|
async def clear(self, msg: ClearMsg) -> ClearRetMsg:
|
|
"""
|
|
Clear kv chunks of instance-worker(s).
|
|
"""
|
|
assert self.cluster_executor is not None
|
|
return await self.cluster_executor.execute("clear", msg)
|
|
|
|
async def pin(self, msg: PinMsg) -> PinRetMsg:
|
|
"""
|
|
Pin kv chunks of instance-worker(s).
|
|
"""
|
|
assert self.cluster_executor is not None
|
|
return await self.cluster_executor.execute("pin", msg)
|
|
|
|
async def compress(self, msg: CompressMsg) -> CompressRetMsg:
|
|
"""
|
|
Compress kv chunks of instance-worker(s).
|
|
"""
|
|
assert self.cluster_executor is not None
|
|
return await self.cluster_executor.execute("compress", msg)
|
|
|
|
async def decompress(self, msg: DecompressMsg) -> DecompressRetMsg:
|
|
"""
|
|
Decompress kv chunks of instance-worker(s).
|
|
"""
|
|
assert self.cluster_executor is not None
|
|
return await self.cluster_executor.execute("decompress", msg)
|
|
|
|
async def move(self, msg: MoveMsg) -> MoveRetMsg:
|
|
"""
|
|
Move kv chunks of instance-worker(s).
|
|
"""
|
|
assert self.cluster_executor is not None
|
|
return await self.cluster_executor.execute("move", msg)
|
|
|
|
async def check_finish(self, msg: CheckFinishMsg) -> CheckFinishRetMsg:
|
|
"""
|
|
Check if an event is finished.
|
|
"""
|
|
assert self.cluster_executor is not None
|
|
return await self.cluster_executor.execute("check_finish", msg)
|
|
|
|
async def handle_batched_kv_operations(self, msg: BatchedKVOperationMsg) -> None:
|
|
"""Handle batched KV operations by forwarding to registry."""
|
|
if not msg.operations:
|
|
return
|
|
|
|
# Check if worker is currently in full sync
|
|
if self.full_sync_tracker.is_worker_syncing(msg.instance_id, msg.worker_id):
|
|
# During full sync, incremental operations should be discarded
|
|
logger.debug(
|
|
"Discarding incremental KV operations during full sync: "
|
|
"instance=%s, worker=%d, sync_id=%s, operation_count=%d",
|
|
msg.instance_id,
|
|
msg.worker_id,
|
|
self.full_sync_tracker.get_sync_id(msg.instance_id, msg.worker_id),
|
|
len(msg.operations),
|
|
)
|
|
return
|
|
|
|
if not self.registry.handle_batched_kv_operations(msg):
|
|
logger.warning(
|
|
"Failed to handle batched KV operations, instance: %s, worker: %d",
|
|
msg.instance_id,
|
|
msg.worker_id,
|
|
)
|
|
|
|
# ============= Full Sync Message Handlers =============
|
|
|
|
async def handle_full_sync_start(
|
|
self, msg: FullSyncStartMsg
|
|
) -> FullSyncStartRetMsg:
|
|
"""
|
|
Handle full sync start request from a worker.
|
|
|
|
This is called when a worker wants to start full sync.
|
|
The controller should:
|
|
1. Clear existing keys for this worker
|
|
2. Mark the worker as syncing (incremental events will be discarded)
|
|
3. Return acceptance
|
|
"""
|
|
instance_id = msg.instance_id
|
|
worker_id = msg.worker_id
|
|
sync_id = msg.sync_id
|
|
report_id = (instance_id, worker_id)
|
|
|
|
# Start sync tracking first (mark worker as SYNCING)
|
|
success = self.full_sync_tracker.start_sync(
|
|
instance_id=instance_id,
|
|
worker_id=worker_id,
|
|
sync_id=sync_id,
|
|
total_keys=msg.total_keys,
|
|
batch_count=msg.batch_count,
|
|
)
|
|
|
|
if not success:
|
|
logger.warning(
|
|
"Failed to start sync for worker %s: sync_id=%s", report_id, sync_id
|
|
)
|
|
return FullSyncStartRetMsg(
|
|
sync_id=sync_id,
|
|
accepted=False,
|
|
error_msg="Failed to start sync: worker already syncing with "
|
|
"different sync_id or worker not found",
|
|
)
|
|
|
|
# Now clear existing keys for this worker/location using efficient batch method
|
|
# This prevents new incremental messages from being processed while we clear
|
|
existing_keys = self.registry.get_worker_kv_keys(
|
|
instance_id, worker_id, msg.location
|
|
)
|
|
if existing_keys:
|
|
old_count = len(existing_keys)
|
|
# Use efficient batch clear method
|
|
cleared = self.registry.clear_worker_kv(
|
|
instance_id, worker_id, msg.location
|
|
)
|
|
if cleared:
|
|
logger.info(
|
|
"Cleared %d existing keys for worker %s location %s "
|
|
"before full sync",
|
|
old_count,
|
|
report_id,
|
|
msg.location,
|
|
)
|
|
else:
|
|
logger.warning(
|
|
"Failed to clear keys for worker %s location %s",
|
|
report_id,
|
|
msg.location,
|
|
)
|
|
|
|
logger.info(
|
|
"Accepted full sync start: worker=%s, sync_id=%s, "
|
|
"total_keys=%d, batch_count=%d",
|
|
report_id,
|
|
sync_id,
|
|
msg.total_keys,
|
|
msg.batch_count,
|
|
)
|
|
return FullSyncStartRetMsg(sync_id=sync_id, accepted=True)
|
|
|
|
async def handle_full_sync_batch(self, msg: FullSyncBatchMsg) -> None:
|
|
"""
|
|
Handle full sync batch message from a worker.
|
|
|
|
This adds the keys from the batch to the registry.
|
|
"""
|
|
instance_id = msg.instance_id
|
|
worker_id = msg.worker_id
|
|
location = msg.location
|
|
sync_id = msg.sync_id
|
|
batch_id = msg.batch_id
|
|
keys = msg.keys
|
|
report_id = (instance_id, worker_id)
|
|
|
|
# Record batch receipt
|
|
if not self.full_sync_tracker.receive_batch(
|
|
instance_id=instance_id,
|
|
worker_id=worker_id,
|
|
sync_id=sync_id,
|
|
batch_id=batch_id,
|
|
keys_count=len(keys),
|
|
):
|
|
logger.warning(
|
|
"Failed to record batch %d for worker %s", batch_id, report_id
|
|
)
|
|
return
|
|
|
|
# Add keys to registry using batched operations
|
|
operations = []
|
|
for seq_num, key in enumerate(keys):
|
|
operations.append(
|
|
KVOpEvent(
|
|
op_type=OpType.ADMIT,
|
|
key=key,
|
|
seq_num=seq_num,
|
|
)
|
|
)
|
|
if operations:
|
|
batch_msg = BatchedKVOperationMsg(
|
|
instance_id=instance_id,
|
|
worker_id=worker_id,
|
|
location=location,
|
|
operations=operations,
|
|
)
|
|
self.registry.handle_batched_kv_operations(batch_msg, is_full_sync=True)
|
|
|
|
current_keys = self.registry.get_worker_kv_keys(
|
|
instance_id, worker_id, location
|
|
)
|
|
logger.debug(
|
|
"Added %d keys from batch %d for worker %s, total now: %d",
|
|
len(keys),
|
|
batch_id,
|
|
report_id,
|
|
len(current_keys),
|
|
)
|
|
|
|
async def handle_full_sync_end(self, msg: FullSyncEndMsg) -> None:
|
|
"""
|
|
Handle full sync end message from a worker.
|
|
|
|
This marks the sync as end-received and records actual total keys.
|
|
"""
|
|
instance_id = msg.instance_id
|
|
worker_id = msg.worker_id
|
|
sync_id = msg.sync_id
|
|
actual_total_keys = msg.actual_total_keys
|
|
report_id = (instance_id, worker_id)
|
|
|
|
success = self.full_sync_tracker.complete_sync(
|
|
instance_id=instance_id,
|
|
worker_id=worker_id,
|
|
sync_id=sync_id,
|
|
actual_total_keys=actual_total_keys,
|
|
)
|
|
|
|
if success:
|
|
# Verify registry has the expected number of keys
|
|
actual_keys_in_pool = len(
|
|
self.registry.get_worker_kv_keys(instance_id, worker_id, msg.location)
|
|
)
|
|
logger.info(
|
|
"Full sync completed for worker %s: sync_id=%s, "
|
|
"reported_keys=%d, keys_in_pool=%d",
|
|
report_id,
|
|
sync_id,
|
|
actual_total_keys,
|
|
actual_keys_in_pool,
|
|
)
|
|
else:
|
|
logger.warning(
|
|
"Failed to complete full sync for worker %s: sync_id=%s",
|
|
report_id,
|
|
sync_id,
|
|
)
|
|
|
|
async def handle_full_sync_status(
|
|
self, msg: FullSyncStatusMsg
|
|
) -> FullSyncStatusRetMsg:
|
|
"""
|
|
Handle full sync status query from a worker.
|
|
|
|
Returns the sync status including any missing batches that need resending.
|
|
"""
|
|
is_complete, global_progress, can_exit_freeze, missing_batches = (
|
|
self.full_sync_tracker.get_sync_status(
|
|
instance_id=msg.instance_id,
|
|
worker_id=msg.worker_id,
|
|
sync_id=msg.sync_id,
|
|
)
|
|
)
|
|
|
|
if missing_batches:
|
|
logger.info(
|
|
"Full sync status query: worker=(%s, %d), sync_id=%s, "
|
|
"is_complete=%s, missing_batches=%s",
|
|
msg.instance_id,
|
|
msg.worker_id,
|
|
msg.sync_id,
|
|
is_complete,
|
|
missing_batches,
|
|
)
|
|
|
|
return FullSyncStatusRetMsg(
|
|
sync_id=msg.sync_id,
|
|
is_complete=is_complete,
|
|
global_progress=global_progress,
|
|
can_exit_freeze=can_exit_freeze,
|
|
missing_batches=missing_batches,
|
|
)
|
|
|
|
# TODO(Jiayi): The current implementation does not handle
|
|
# the case where the prefix chunks are evicted while the
|
|
# suffix chunk is still in the system. LMCache should guarantee
|
|
# this does not happen.
|
|
# TODO(Jiayi): The current implementation does not consider
|
|
# the location of the kv chunks. It simply returns the
|
|
# `instance_id` with longest prefix.
|
|
# TODO(Jiayi): Need to get rid of the hash somehow
|
|
async def lookup(self, msg: LookupMsg) -> LookupRetMsg:
|
|
tokens = msg.tokens
|
|
layout_info = {}
|
|
for start, end, key in self.token_database.process_tokens(
|
|
tokens, make_key=False
|
|
):
|
|
result = self.registry.find_kv(key)
|
|
if result is None:
|
|
break
|
|
matched_instance = result.instance_id
|
|
matched_location = result.location
|
|
layout_info[matched_instance] = (matched_location, end)
|
|
return LookupRetMsg(layout_info=layout_info, event_id=msg.event_id)
|
|
|
|
# TODO: improve the matching logic, return multi results
|
|
async def batched_p2p_lookup(
|
|
self, msg: BatchedP2PLookupMsg
|
|
) -> BatchedP2PLookupRetMsg:
|
|
"""
|
|
Perform batched P2P lookup for multiple keys.
|
|
|
|
:param BatchedP2PLookupMsg msg: The batched P2P lookup message containing keys.
|
|
|
|
:return: A BatchedP2PLookupRetMsg containing the lookup results.
|
|
"""
|
|
hashes = msg.hashes
|
|
if not hashes:
|
|
return BatchedP2PLookupRetMsg(layout_info=[("", "", 0, "")])
|
|
|
|
# Single lookup to get all needed info (optimized path)
|
|
result = self.registry.find_kv_with_worker_info(
|
|
hashes[0], exclude_instance_id=msg.instance_id
|
|
)
|
|
if result is None:
|
|
return BatchedP2PLookupRetMsg(layout_info=[("", "", 0, "")])
|
|
|
|
kv_info, peer_init_url, current_keys = result
|
|
if peer_init_url is None:
|
|
return BatchedP2PLookupRetMsg(layout_info=[("", "", 0, "")])
|
|
|
|
# Count hits efficiently
|
|
num_hit_chunks = 0
|
|
for key in hashes:
|
|
if key not in current_keys:
|
|
break
|
|
num_hit_chunks += 1
|
|
|
|
return BatchedP2PLookupRetMsg(
|
|
layout_info=[
|
|
(kv_info.instance_id, kv_info.location, num_hit_chunks, peer_init_url),
|
|
]
|
|
)
|