Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
63 changes: 63 additions & 0 deletions src/tests/test_kvaware_router.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
"""Unit tests for the KV-cache-aware routing logic."""

import asyncio
import time
from types import SimpleNamespace

import pytest
Expand Down Expand Up @@ -220,3 +222,64 @@ async def query_manager(msg):
)

assert selected == url_b


@pytest.mark.asyncio
async def test_kvaware_remote_tokenize_fallback_does_not_block_event_loop(
monkeypatch,
):
"""The remote /tokenize fallback is a blocking HTTP call; it must run off
the event loop so other requests keep being served."""

def no_local_tokenizer(_model_name):
raise OSError("tokenizer not available locally")

def slow_post(*_args, **_kwargs):
time.sleep(0.3)
return SimpleNamespace(json=lambda: {"tokens": [1, 2, 3]})

monkeypatch.setattr(
routing_logic,
"AutoTokenizer",
SimpleNamespace(from_pretrained=no_local_tokenizer),
raising=False,
)
monkeypatch.setattr(routing_logic.requests, "post", slow_post)

url = "http://10.0.0.1:8000"
router = KvawareRouter.__new__(KvawareRouter)
router.tokenizers = {}
router.threshold = 0
router.session_key = "x-session-id"
router.hash_ring = routing_logic.HashRing()
router.instance_id_to_ip = {}
lookup_tokens = []

async def query_manager(msg):
lookup_tokens.append(msg.tokens)
return SimpleNamespace(layout_info={})

router.query_manager = query_manager

ticks = 0

async def ticker():
nonlocal ticks
while True:
await asyncio.sleep(0.01)
ticks += 1

ticker_task = asyncio.create_task(ticker())
selected = await router.route_request(
[EndpointInfo(url)],
{},
{url: SimpleNamespace(qps=0)},
SimpleNamespace(headers={}),
{"prompt": "test"},
)
ticker_task.cancel()

assert selected == url
assert lookup_tokens == [[1, 2, 3]]
# the event loop kept running while the fallback request was in flight
assert ticks >= 10
72 changes: 28 additions & 44 deletions src/vllm_router/routers/routing_logic.py
Original file line number Diff line number Diff line change
Expand Up @@ -367,6 +367,33 @@ def close(self):
pass
self.lmcache_cluster_monitor_task = None

async def tokenize_prompt(
self, endpoints: List[EndpointInfo], request_json: Dict
) -> List[int]:
"""Local-first tokenization with the remote `/tokenize` fallback.

The remote fallback is a blocking HTTP call, so it runs in an
executor rather than on the event loop.
"""
try:
tokenizer = self._get_tokenizer(endpoints)
return tokenizer.encode(request_json.get("prompt", ""))
except Exception:
remote_url = endpoints[0].url + "/tokenize"
headers = {"Content-Type": "application/json"}
data = {
"model": endpoints[0].model_names[0],
"prompt": request_json.get("prompt", ""),
}
loop = asyncio.get_running_loop()
response = await loop.run_in_executor(
None,
lambda: requests.post(
remote_url, headers=headers, json=data, timeout=10
),
)
return response.json()["tokens"]

async def route_request(
self,
endpoints: List[EndpointInfo],
Expand All @@ -391,24 +418,8 @@ async def route_request(
request_json (Dict): The request body (needed for finding the
longest prefix match)
"""
token_ids = None
# Local-first tokenization, fall back to remote "/tokenize" API on failure
# TODO (Yuhan): Handle chat completions
try:
tokenizer = self._get_tokenizer(endpoints)
token_ids = tokenizer.encode(request_json.get("prompt", ""))
except Exception:
# Remote /tokenize fallback (let errors bubble up to keep behavior simple)
remote_url = endpoints[0].url + "/tokenize"
headers = {"Content-Type": "application/json"}
data = {
"model": endpoints[0].model_names[0],
"prompt": request_json.get("prompt", ""),
}
body = requests.post(
remote_url, headers=headers, json=data, timeout=10
).json()
token_ids = body["tokens"]
token_ids = await self.tokenize_prompt(endpoints, request_json)
Comment on lines 426 to +427

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

If endpoints is empty (e.g., when no backends are currently discovered or healthy), calling self.tokenize_prompt(endpoints, request_json) will raise an unhandled IndexError when attempting to access endpoints[0]. To prevent a 500 Internal Server Error, we should defensively check if endpoints is empty and raise a 503 HTTPException, consistent with how LoadAwareRouter handles this scenario.

Suggested change
# TODO (Yuhan): Handle chat completions
try:
tokenizer = self._get_tokenizer(endpoints)
token_ids = tokenizer.encode(request_json.get("prompt", ""))
except Exception:
# Remote /tokenize fallback (let errors bubble up to keep behavior simple)
remote_url = endpoints[0].url + "/tokenize"
headers = {"Content-Type": "application/json"}
data = {
"model": endpoints[0].model_names[0],
"prompt": request_json.get("prompt", ""),
}
body = requests.post(
remote_url, headers=headers, json=data, timeout=10
).json()
token_ids = body["tokens"]
token_ids = await self.tokenize_prompt(endpoints, request_json)
if not endpoints:
raise HTTPException(
status_code=503, detail="No backend endpoints available"
)
# TODO (Yuhan): Handle chat completions
token_ids = await self.tokenize_prompt(endpoints, request_json)


event_id = "Lookup" + str(uuid.uuid4())
msg = LookupMsg(tokens=token_ids, event_id=event_id)
Expand Down Expand Up @@ -662,33 +673,6 @@ async def query_endpoint(endpoint: EndpointInfo) -> None:
await asyncio.gather(*(query_endpoint(e) for e in endpoints))
logger.info(f"Instance id to ip mapping: {self.instance_id_to_ip}")

async def tokenize_prompt(
self, endpoints: List[EndpointInfo], request_json: Dict
) -> List[int]:
"""Local-first tokenization with the remote `/tokenize` fallback.

The remote fallback is a blocking HTTP call, so it runs in an
executor rather than on the event loop.
"""
try:
tokenizer = self._get_tokenizer(endpoints)
return tokenizer.encode(request_json.get("prompt", ""))
except Exception:
remote_url = endpoints[0].url + "/tokenize"
headers = {"Content-Type": "application/json"}
data = {
"model": endpoints[0].model_names[0],
"prompt": request_json.get("prompt", ""),
}
loop = asyncio.get_running_loop()
response = await loop.run_in_executor(
None,
lambda: requests.post(
remote_url, headers=headers, json=data, timeout=10
),
)
return response.json()["tokens"]

def fallback_url(
self,
endpoints: List[EndpointInfo],
Expand Down
Loading