diff --git a/docs/namespace/README_en.md b/docs/namespace/README_en.md index 7d8294ba62..18090e3ee7 100644 --- a/docs/namespace/README_en.md +++ b/docs/namespace/README_en.md @@ -4,6 +4,18 @@ This document describes how to use FlexKV's namespace isolation feature to enabl ## Version Requirements +### SGLang + +The connector accepts `namespace: Optional[List[str]]` on `lookup_kv`, +`store_kv`, and `prefetch_async`. Callers must pass the same identity to all +three operations. `None` preserves the existing unscoped keys. + +The matching SGLang adapter uses one compact JSON namespace component, +`["sglang-cache-v1", extra_key, cache_salt]`, and enables scoped host reuse only +when `supports_cache_namespace` is true. Empty strings and null values are +kept distinct. A chunked-prefetch implementation must also propagate the +namespace into its session/planner before advertising that capability. + ### vLLM - For vLLM, please use the `examples/vllm_adaption/vllm_0_10_1_1-flexkv-connector-namespace.patch` patch to support namespace functionality diff --git a/flexkv/integration/sglang/connector.py b/flexkv/integration/sglang/connector.py index f76508d8aa..b7cfea67c3 100644 --- a/flexkv/integration/sglang/connector.py +++ b/flexkv/integration/sglang/connector.py @@ -33,7 +33,7 @@ import socket import struct import time -from contextlib import nullcontext +from contextlib import nullcontext, suppress from dataclasses import dataclass, replace from typing import Any, Dict, List, Optional, Sequence, Tuple @@ -108,7 +108,7 @@ class FlexKVHostReleaseShim: ``FlexKVConnector.shutdown()`` so unpin stays on that same path. """ - def __init__(self, connector: "FlexKVConnector") -> None: + def __init__(self, connector: FlexKVConnector) -> None: self._connector = connector def destroy(self) -> None: @@ -142,6 +142,17 @@ class FlexKVConnector: * ``reset`` / ``shutdown``. """ + @property + def supports_cache_namespace(self) -> bool: + """All enabled lookup/store/prefetch paths accept ``namespace``. + + A chunked-prefetch connector must opt in separately after forwarding + the namespace through every chunk and its session identity. + """ + return not getattr(self, "_chunked_prefetch", False) or getattr( + self, "_chunked_namespace_supported", False + ) is True + def __init__( self, *, @@ -331,10 +342,8 @@ def __init__( # * atexit runs kv_manager.shutdown() (cudaHostUnregister). # * Stretch the parent's scheduler-exit wait so kill_process_tree # does not SIGKILL mid-unpin (generic env, not FlexKV-named in TM). - try: + with suppress(Exception): signal.signal(signal.SIGINT, signal.SIG_IGN) - except Exception: # noqa: BLE001 - pass # Graceful shutdown timeout hierarchy (see FlexKV config.py): # tokenizer wait (this env) = 1200s # > TM parent-side wait (FLEXKV_TRANSFER_MANAGER_SHUTDOWN_TIMEOUT_S) = 900s @@ -481,6 +490,7 @@ def lookup_kv( token_mask: torch.Tensor, rid: Optional[str] = None, sglang_req_id: Any = _SGLANG_REQ_ID_UNSET, + namespace: Optional[List[str]] = None, ) -> Tuple[int, int]: """Page-aligned prefix lookup against FlexKV. @@ -495,6 +505,8 @@ def lookup_kv( sglang_req_id: business request ID used only for logs. This can be ``None`` when ``rid`` is an internal tracking key. + namespace: cache identity components shared with store and prefetch. + Returns: ``(fkv_task_id, hit_count)``. ``hit_count`` is page-aligned and may be smaller than the raw FlexKV match if the page @@ -515,6 +527,7 @@ def lookup_kv( token_ids=tids_np, token_mask=mask_np, swa_aware=self._swa_kv_pool is not None, + namespace=namespace, ) except Exception as exc: # noqa: BLE001 lookup_error = exc @@ -1061,6 +1074,7 @@ def store_kv( token_ids: List[int], kv_indices: torch.Tensor, sglang_req_id: Any = _SGLANG_REQ_ID_UNSET, + namespace: Optional[List[str]] = None, ) -> int: """Schedule a write back from GPU into FlexKV. @@ -1115,7 +1129,7 @@ def store_kv( try: with self._store_profile_scope("flexkv.connector.store.put_match"): res = self.kv_manager.put_match( - token_ids=token_ids_np, token_mask=None + token_ids=token_ids_np, token_mask=None, namespace=namespace ) except Exception as exc: # noqa: BLE001 match_error = exc @@ -1269,40 +1283,39 @@ def check_completed_stores(self) -> List[str]: completed_rids: List[str] = [] completed_by_rid: Dict[str, Any] = {} - if self._sync_ctx.is_sync_leader and self.kv_manager is not None: - if self._inflight_stores: - fk_to_rid = {v: k for k, v in self._inflight_stores.items()} - try: - completed_dict = ( - self.kv_manager.wait( - list(fk_to_rid.keys()), - timeout=0.0, - completely=True, - ) - or {} - ) - except Exception as exc: # noqa: BLE001 - rid = next(iter(self._inflight_stores)) - context = getattr(self, "_inflight_store_contexts", {}).get(rid) - context = context or self._new_op_context( - "store", rid, task_id=self._inflight_stores[rid] - ) - self._log_cache_op( - context, - "poll", - "failed", - task_id=-1, - flexkv_task_ids=list(fk_to_rid), - error=str(exc), + if self._sync_ctx.is_sync_leader and self.kv_manager is not None and self._inflight_stores: + fk_to_rid = {v: k for k, v in self._inflight_stores.items()} + try: + completed_dict = ( + self.kv_manager.wait( + list(fk_to_rid.keys()), + timeout=0.0, + completely=True, ) - completed_dict = {} - for fk_tid, response in completed_dict.items(): - status = _status_value(response) - if not _is_terminal_status(status): - continue - rid = fk_to_rid[fk_tid] - completed_rids.append(rid) - completed_by_rid[rid] = response + or {} + ) + except Exception as exc: # noqa: BLE001 + rid = next(iter(self._inflight_stores)) + context = getattr(self, "_inflight_store_contexts", {}).get(rid) + context = context or self._new_op_context( + "store", rid, task_id=self._inflight_stores[rid] + ) + self._log_cache_op( + context, + "poll", + "failed", + task_id=-1, + flexkv_task_ids=list(fk_to_rid), + error=str(exc), + ) + completed_dict = {} + for fk_tid, response in completed_dict.items(): + status = _status_value(response) + if not _is_terminal_status(status): + continue + rid = fk_to_rid[fk_tid] + completed_rids.append(rid) + completed_by_rid[rid] = response if self._sync_ctx.needs_sync: completed_rids = self._sync_ctx.scatter( @@ -1389,6 +1402,7 @@ def prefetch_async( rid: str, token_ids: List[int], sglang_req_id: Any = _SGLANG_REQ_ID_UNSET, + namespace: Optional[List[str]] = None, ) -> int: if not self._prefetch_enabled or not rid: return -1 @@ -1399,7 +1413,8 @@ def prefetch_async( if self._sync_ctx.is_sync_leader and self.kv_manager is not None: try: prefetch_result = self.kv_manager.prefetch_async( - token_ids=np.asarray(token_ids, dtype=np.int64) + token_ids=np.asarray(token_ids, dtype=np.int64), + namespace=namespace, ) # KVManager currently returns # ``(task_id, actual_prefetch_tokens)`` even though older diff --git a/tests/test_sglang_store_protocol.py b/tests/test_sglang_store_protocol.py index 72dffc095c..bd52a052b1 100644 --- a/tests/test_sglang_store_protocol.py +++ b/tests/test_sglang_store_protocol.py @@ -8,7 +8,6 @@ from flexkv.common.request import KVResponse, KVResponseStatus from flexkv.integration.sglang.comm import FlexKVScatterChannel from flexkv.integration.sglang.connector import FlexKVConnector -from flexkv.common.request import KVResponseStatus def _follower_connector(payload): @@ -342,3 +341,39 @@ def test_store_reset_waits_for_leader_drain_on_follower(): "blocking": True, } assert connector._inflight_stores == {} + + +def test_namespace_capability_is_explicit_for_all_active_paths(): + connector = FlexKVConnector.__new__(FlexKVConnector) + assert connector.supports_cache_namespace is True + # A chunked-prefetch implementation needs its own full-chain opt-in. + connector._chunked_prefetch = True + connector._chunked_namespace_supported = False + assert connector.supports_cache_namespace is False + connector._chunked_namespace_supported = True + assert connector.supports_cache_namespace is True + + +def test_namespace_is_forwarded_to_lookup_store_and_prefetch(): + connector = _follower_connector(None) + connector._sync_ctx.is_sync_leader = True + connector._sync_ctx.needs_sync = False + connector._sync_ctx.is_pp_active = False + connector._swa_kv_pool = None + connector._pending_lookups = {} + connector._pending_lookup_contexts = {} + connector._prefetch_enabled = True + connector._ongoing_prefetches = {} + connector._prefetch_contexts = {} + connector._prefetch_planned_tokens = {} + connector.kv_manager = MagicMock() + connector.kv_manager.get_match.return_value = (-1, np.zeros(4, dtype=bool)) + connector.kv_manager.put_match.return_value = (-1, np.zeros(4, dtype=bool)) + connector.kv_manager.prefetch_async.return_value = (23, 4) + namespace = ['["sglang-cache-v1","adapter","salt"]'] + connector.lookup_kv([1, 2, 3, 4], torch.ones(4, dtype=torch.bool), rid="r", namespace=namespace) + connector.store_kv("r", [1, 2, 3, 4], torch.arange(4), namespace=namespace) + assert connector.prefetch_async("r", [1, 2, 3, 4], namespace=namespace) == 23 + for method in (connector.kv_manager.get_match, connector.kv_manager.put_match, connector.kv_manager.prefetch_async): + assert method.call_args.kwargs["namespace"] == namespace + assert method.call_args.kwargs["token_ids"].tolist() == [1, 2, 3, 4]