Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions docs/namespace/README_en.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
95 changes: 55 additions & 40 deletions flexkv/integration/sglang/connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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,
*,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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.

Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
37 changes: 36 additions & 1 deletion tests/test_sglang_store_protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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]
Loading