Skip to content

Commit afee0b9

Browse files
[CI] 【Hackathon 10th Spring No.30】功能模块 fastdeploy/inter_communicator/engine_worker_queue.py单测补充 (#6102)
* test: add comprehensive tests for EngineWorkerQueue to improve code coverage * style: format tests/inter_communicator/test_e2w_queue.py with black --------- Co-authored-by: CSWYF3634076 <wangyafeng@baidu.com>
1 parent 0306475 commit afee0b9

1 file changed

Lines changed: 226 additions & 0 deletions

File tree

tests/inter_communicator/test_e2w_queue.py

Lines changed: 226 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,13 +14,20 @@
1414
# limitations under the License.
1515
"""
1616

17+
import threading
18+
import time
19+
import types
1720
import unittest
1821

1922
import numpy as np
2023
import paddle
2124

25+
if not hasattr(paddle, "compat"):
26+
paddle.compat = types.SimpleNamespace(enable_torch_proxy=lambda **_: None)
27+
2228
from fastdeploy import envs
2329
from fastdeploy.engine.request import Request
30+
from fastdeploy.inter_communicator.engine_worker_queue import EngineWorkerQueue
2431
from fastdeploy.utils import to_numpy, to_tensor
2532

2633

@@ -30,6 +37,37 @@ def __init__(self, images):
3037

3138

3239
class TestEngineWorkerQueue(unittest.TestCase):
40+
def _build_queue_pair(self):
41+
server = EngineWorkerQueue(address=("127.0.0.1", 0), is_server=True, num_client=1, client_id=0)
42+
client = EngineWorkerQueue(
43+
address=server.address,
44+
is_server=False,
45+
num_client=1,
46+
client_id=0,
47+
)
48+
return server, client
49+
50+
def _cleanup_queue_pair(self, server):
51+
server.cleanup()
52+
53+
def _set_list_after_delay(self, list_proxy, values, delay=0.01):
54+
def updater():
55+
time.sleep(delay)
56+
list_proxy[:] = values
57+
58+
thread = threading.Thread(target=updater)
59+
thread.start()
60+
return thread
61+
62+
def _set_value_after_delay(self, value_proxy, value, delay=0.01):
63+
def updater():
64+
time.sleep(delay)
65+
value_proxy.set(value)
66+
67+
thread = threading.Thread(target=updater)
68+
thread.start()
69+
return thread
70+
3371
def test_to_tensor_success(self):
3472
envs.FD_ENABLE_MAX_PREFILL = 1
3573
# 模拟 numpy 数组输入(使用 paddle 转 numpy)
@@ -172,6 +210,194 @@ def test_features_info_to_numpy(self):
172210
self.assertIsInstance(task.multimodal_inputs["video_features"][0], np.ndarray)
173211
self.assertIsInstance(task.multimodal_inputs["video_features"][1], np.ndarray)
174212

213+
def test_queue_exist_tasks_and_ports(self):
214+
server, client = self._build_queue_pair()
215+
try:
216+
self.assertIsNone(server.exist_tasks_intra_signal)
217+
self.assertFalse(client.exist_tasks())
218+
client.set_exist_tasks(True)
219+
self.assertTrue(client.exist_tasks())
220+
self.assertEqual(server.get_server_port(), server.address[1])
221+
with self.assertRaises(RuntimeError):
222+
client.get_server_port()
223+
finally:
224+
self._cleanup_queue_pair(server)
225+
226+
def test_single_node_signal_updates(self):
227+
server = EngineWorkerQueue(address=("0.0.0.0", 0), is_server=True, num_client=1, client_id=0)
228+
try:
229+
self.assertFalse(server.exist_tasks())
230+
server.set_exist_tasks(True)
231+
self.assertTrue(server.exist_tasks())
232+
server.set_exist_tasks(False)
233+
self.assertFalse(server.exist_tasks())
234+
finally:
235+
server.cleanup()
236+
server.exist_tasks_intra_signal.clear()
237+
238+
def test_put_get_tasks_and_clear_data(self):
239+
envs.FD_ENABLE_MAX_PREFILL = 0
240+
envs.FD_ENABLE_E2W_TENSOR_CONVERT = 0
241+
server, client = self._build_queue_pair()
242+
try:
243+
tasks = ["task-A"]
244+
client.put_tasks(tasks)
245+
self.assertEqual(client.num_tasks(), 1)
246+
fetched, all_read = client.get_tasks()
247+
self.assertTrue(all_read)
248+
self.assertEqual(fetched, [tasks])
249+
self.assertEqual(client.num_tasks(), 0)
250+
client.put_tasks(tasks)
251+
client.clear_data()
252+
self.assertEqual(list(client.client_read_flag), [1])
253+
self.assertEqual(client.num_tasks(), 0)
254+
finally:
255+
self._cleanup_queue_pair(server)
256+
257+
def test_wait_loops_and_tensor_conversion(self):
258+
envs.FD_ENABLE_MAX_PREFILL = 1
259+
envs.FD_ENABLE_E2W_TENSOR_CONVERT = 0
260+
server, client = self._build_queue_pair()
261+
previous_device = paddle.get_device()
262+
paddle.set_device("cpu")
263+
try:
264+
np_images = paddle.randn([1, 3, 4, 4]).numpy()
265+
task = DummyTask(np_images)
266+
tasks = [[task]]
267+
client.client_read_flag[:] = [0]
268+
thread = self._set_list_after_delay(client.client_read_flag, [1])
269+
client.put_tasks(tasks)
270+
thread.join()
271+
self.assertIsInstance(task.multimodal_inputs["images"], paddle.Tensor)
272+
273+
client.client_get_connect_task_flag[:] = [0]
274+
thread = self._set_list_after_delay(client.client_get_connect_task_flag, [1])
275+
client.put_connect_rdma_task({"connect": "wait"})
276+
thread.join()
277+
278+
client.can_put_next_connect_task_response_flag.set(0)
279+
thread = self._set_value_after_delay(client.can_put_next_connect_task_response_flag, 1)
280+
client.put_connect_rdma_task_response({"success": True})
281+
thread.join()
282+
283+
client.connect_rdma_task_responses.append({"success": True})
284+
client.client_get_connect_task_response_flag[:] = [0]
285+
thread = self._set_list_after_delay(client.client_get_connect_task_response_flag, [1])
286+
client.get_connect_rdma_task_response()
287+
thread.join()
288+
289+
client.client_read_info_flag[:] = [0]
290+
thread = self._set_list_after_delay(client.client_read_info_flag, [1])
291+
client.put_cache_info([{"cache": "wait"}])
292+
thread.join()
293+
294+
client.can_put_next_send_cache_finished_flag.set(0)
295+
thread = self._set_value_after_delay(client.can_put_next_send_cache_finished_flag, 1)
296+
client.put_finished_req([["req-wait", {"status": "ok"}]])
297+
thread.join()
298+
299+
client.finished_send_cache_list.append(["req-wait", {"error": "bad"}])
300+
client.client_get_finish_send_cache_flag[:] = [0]
301+
thread = self._set_list_after_delay(client.client_get_finish_send_cache_flag, [1])
302+
client.get_finished_req()
303+
thread.join()
304+
305+
client.can_put_next_add_task_finished_flag.set(0)
306+
thread = self._set_value_after_delay(client.can_put_next_add_task_finished_flag, 1)
307+
client.put_finished_add_cache_task_req(["req-wait"])
308+
thread.join()
309+
310+
client.finished_add_cache_task_list.append(["req-wait"])
311+
client.client_get_finished_add_cache_task_flag[:] = [0]
312+
thread = self._set_list_after_delay(client.client_get_finished_add_cache_task_flag, [1])
313+
client.get_finished_add_cache_task_req()
314+
thread.join()
315+
finally:
316+
paddle.set_device(previous_device)
317+
self._cleanup_queue_pair(server)
318+
319+
def test_connect_rdma_task_flow(self):
320+
server, client = self._build_queue_pair()
321+
try:
322+
client.client_get_connect_task_flag[:] = [1]
323+
client.put_connect_rdma_task({"connect": "ok"})
324+
task, all_read = client.get_connect_rdma_task()
325+
self.assertTrue(all_read)
326+
self.assertEqual(task, {"connect": "ok"})
327+
self.assertEqual(list(client.connect_rdma_tasks), [])
328+
329+
self.assertIsNone(client.get_connect_rdma_task_response())
330+
response = {"success": True}
331+
self.assertTrue(client.put_connect_rdma_task_response(response))
332+
client.connect_rdma_task_responses.append({"success": False})
333+
merged = client.get_connect_rdma_task_response()
334+
self.assertEqual(merged["success"], False)
335+
self.assertEqual(client.can_put_next_connect_task_response_flag.get(), 1)
336+
finally:
337+
self._cleanup_queue_pair(server)
338+
339+
def test_cache_info_and_counts(self):
340+
server, client = self._build_queue_pair()
341+
try:
342+
client.client_read_info_flag[:] = [1]
343+
cache_info = [{"cache": "info"}]
344+
client.put_cache_info(cache_info)
345+
self.assertEqual(client.num_cache_infos(), 1)
346+
self.assertEqual(client.get_cache_info(), cache_info)
347+
self.assertEqual(client.num_cache_infos(), 0)
348+
self.assertEqual(client.get_cache_info(), [])
349+
finally:
350+
self._cleanup_queue_pair(server)
351+
352+
def test_finished_req_flow(self):
353+
server, client = self._build_queue_pair()
354+
try:
355+
send_cache_result = [["req-1", {"status": "ok"}]]
356+
self.assertTrue(client.put_finished_req(send_cache_result))
357+
client.finished_send_cache_list.append(["req-1", {"error": "bad"}])
358+
response = client.get_finished_req()
359+
self.assertEqual(response, [["req-1", {"error": "bad"}]])
360+
self.assertEqual(client.get_finished_req(), [])
361+
self.assertEqual(client.can_put_next_send_cache_finished_flag.get(), 1)
362+
finally:
363+
self._cleanup_queue_pair(server)
364+
365+
def test_finished_add_cache_task_req(self):
366+
server, client = self._build_queue_pair()
367+
try:
368+
req_ids = ["req-2"]
369+
self.assertTrue(client.put_finished_add_cache_task_req(req_ids))
370+
client.finished_add_cache_task_list.append(req_ids)
371+
self.assertEqual(client.get_finished_add_cache_task_req(), req_ids)
372+
self.assertEqual(client.get_finished_add_cache_task_req(), [])
373+
self.assertEqual(client.can_put_next_add_task_finished_flag.get(), 1)
374+
finally:
375+
self._cleanup_queue_pair(server)
376+
377+
def test_disaggregated_queue(self):
378+
server, client = self._build_queue_pair()
379+
try:
380+
self.assertTrue(client.disaggregate_queue_empty())
381+
client.put_disaggregated_tasks({"item": 1})
382+
client.put_disaggregated_tasks({"item": 2})
383+
self.assertFalse(client.disaggregate_queue_empty())
384+
self.assertEqual(client.get_disaggregated_tasks(), [{"item": 1}, {"item": 2}])
385+
self.assertIsNone(client.get_disaggregated_tasks())
386+
finally:
387+
self._cleanup_queue_pair(server)
388+
389+
def test_connect_retry_failure(self):
390+
dummy = EngineWorkerQueue.__new__(EngineWorkerQueue)
391+
392+
class DummyManager:
393+
def connect(self):
394+
raise ConnectionRefusedError("refused")
395+
396+
dummy.manager = DummyManager()
397+
dummy.address = ("127.0.0.1", 9999)
398+
with self.assertRaises(ConnectionError):
399+
dummy._connect_with_retry(max_retries=2, interval=0)
400+
175401

176402
if __name__ == "__main__":
177403
unittest.main()

0 commit comments

Comments
 (0)