diff --git a/lightllm/server/multi_level_kv_cache/manager.py b/lightllm/server/multi_level_kv_cache/manager.py index ef5b7369c9..6ce4a6d472 100644 --- a/lightllm/server/multi_level_kv_cache/manager.py +++ b/lightllm/server/multi_level_kv_cache/manager.py @@ -9,7 +9,7 @@ import threading import concurrent.futures import setproctitle -from queue import Queue +from queue import Empty, Queue from typing import List from lightllm.server.core.objs import ShmReqManager, Req, StartArgs from lightllm.server.core.objs.io_objs import GroupReqIndexes @@ -45,6 +45,8 @@ def __init__( # 控制进行 cpu cache 页面匹配的时间,超过时间则不再匹配,直接转发。 self.cpu_cache_time_out = 0.5 self.recv_queue = Queue(maxsize=1024) + # Workers publish completions; only recv_loop accesses the PUSH socket. + self.send_to_router_queue = Queue() self.cpu_cache_thread = threading.Thread(target=self.cpu_cache_hanle_loop, daemon=True) self.cpu_cache_thread.start() @@ -144,7 +146,7 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes # 超时时,放弃进行 cache page 的匹配。 current_time = time.time() if current_time - start_time >= self.cpu_cache_time_out: - self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) + self.send_to_router_queue.put(group_req_indexes) logger.warning( f"cache matching time out {current_time - start_time}s, " f"group_req_id: {group_req_indexes.group_req_id}" @@ -211,14 +213,23 @@ def _handle_group_req_multi_cache_match(self, group_req_indexes: GroupReqIndexes for req in reqs: self.shm_req_manager.put_back_req_obj(req) - self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) + self.send_to_router_queue.put(group_req_indexes) return + def _send_finished_group_reqs(self): + while True: + try: + group_req_indexes = self.send_to_router_queue.get_nowait() + except Empty: + return + self.send_to_router.send_pyobj(group_req_indexes, protocol=pickle.HIGHEST_PROTOCOL) + def recv_loop(self): try: recv_max_count = 128 while True: + self._send_finished_group_reqs() recv_objs = [] try: # 一次最多从 zmq 中取 recv_max_count 个请求,防止 zmq 队列中请求数量过多导致阻塞了主循环。 diff --git a/unit_tests/server/multi_level_kv_cache/test_manager_socket_ownership.py b/unit_tests/server/multi_level_kv_cache/test_manager_socket_ownership.py new file mode 100644 index 0000000000..b376f163c8 --- /dev/null +++ b/unit_tests/server/multi_level_kv_cache/test_manager_socket_ownership.py @@ -0,0 +1,41 @@ +import pickle +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from queue import Queue +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +from lightllm.server.multi_level_kv_cache.manager import MultiLevelKVCacheManager + + +@pytest.mark.parametrize("expired", [False, True]) +def test_concurrent_workers_leave_socket_sends_to_owner_thread(expired): + manager = MultiLevelKVCacheManager.__new__(MultiLevelKVCacheManager) + manager.cpu_cache_time_out = 0.5 + manager.send_to_router_queue = Queue() + manager.send_to_router = Mock() + manager.shm_req_manager = Mock() + groups = [SimpleNamespace(group_req_id=i, shm_req_indexes=[]) for i in range(100)] + start_time = time.time() - 10 if expired else time.time() + + with ThreadPoolExecutor(max_workers=8) as executor: + list(executor.map(lambda group: manager._handle_group_req_multi_cache_match(group, start_time), groups)) + + manager.send_to_router.send_pyobj.assert_not_called() + owner_thread = threading.get_ident() + sent_groups = [] + + def send(group, protocol): + assert threading.get_ident() == owner_thread + assert protocol == pickle.HIGHEST_PROTOCOL + sent_groups.append(group.group_req_id) + + manager.send_to_router.send_pyobj.side_effect = send + manager._send_finished_group_reqs() + assert sorted(sent_groups) == list(range(100)) + assert manager.send_to_router_queue.empty() + manager._send_finished_group_reqs() + assert len(sent_groups) == 100