1414# limitations under the License.
1515"""
1616
17+ import threading
18+ import time
19+ import types
1720import unittest
1821
1922import numpy as np
2023import paddle
2124
25+ if not hasattr (paddle , "compat" ):
26+ paddle .compat = types .SimpleNamespace (enable_torch_proxy = lambda ** _ : None )
27+
2228from fastdeploy import envs
2329from fastdeploy .engine .request import Request
30+ from fastdeploy .inter_communicator .engine_worker_queue import EngineWorkerQueue
2431from fastdeploy .utils import to_numpy , to_tensor
2532
2633
@@ -30,6 +37,37 @@ def __init__(self, images):
3037
3138
3239class 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
176402if __name__ == "__main__" :
177403 unittest .main ()
0 commit comments