@@ -34,8 +34,8 @@ def setUp(self):
3434 self ._model_name = "simple"
3535 self ._local_host = "localhost"
3636 self ._http_port = 8000
37- self ._max_chunk_count = (
38- 1000000 # defined in server/src/common.h HTTP_MAX_CHUNKS
37+ self ._malicious_chunk_count = (
38+ 1000000 # large enough to cause a stack overflow if using alloca()
3939 )
4040 self ._parse_error = (
4141 "failed to parse the request JSON buffer: Invalid value. at 0"
@@ -80,13 +80,6 @@ def send_chunked_request(
8080 finally :
8181 s .close ()
8282
83- def chunk_exceed_error (self , request_kind : str , chunk_count : int ):
84- # defined in server/src/http_server.cc
85- if request_kind == "input" :
86- return f"Chunks in the { request_kind } buffer exceed the limit { self ._max_chunk_count } , got { chunk_count } chunks"
87- else :
88- return f"Chunks in the { request_kind } request buffer exceed the limit { self ._max_chunk_count } , got { chunk_count } chunks"
89-
9083 def test_infer (self ):
9184 request_header = (
9285 f"POST /v2/models/{ self ._model_name } /infer HTTP/1.1\r \n "
@@ -95,25 +88,15 @@ def test_infer(self):
9588
9689 self .send_chunked_request (
9790 request_header ,
98- self ._max_chunk_count ,
91+ self ._malicious_chunk_count ,
9992 "Raw request must only have 1 input (found 1) to be deduced but got 2 inputs in 'simple' model configuration" ,
10093 )
101- self .send_chunked_request (
102- request_header ,
103- self ._max_chunk_count + 1 ,
104- self .chunk_exceed_error ("input" , self ._max_chunk_count + 1 ),
105- )
10694
10795 def test_registry_index (self ):
10896 request_header = f"POST /v2/repository/index HTTP/1.1\r \n "
10997
11098 self .send_chunked_request (
111- request_header , self ._max_chunk_count , self ._parse_error
112- )
113- self .send_chunked_request (
114- request_header ,
115- self ._max_chunk_count + 1 ,
116- self .chunk_exceed_error ("registry index" , self ._max_chunk_count + 1 ),
99+ request_header , self ._malicious_chunk_count , self ._parse_error
117100 )
118101
119102 def test_model_control (self ):
@@ -123,20 +106,10 @@ def test_model_control(self):
123106 unload_request_header = load_request_header .replace ("/load" , "/unload" )
124107
125108 self .send_chunked_request (
126- load_request_header , self ._max_chunk_count , self ._parse_error
109+ load_request_header , self ._malicious_chunk_count , self ._parse_error
127110 )
128111 self .send_chunked_request (
129- load_request_header ,
130- self ._max_chunk_count + 1 ,
131- self .chunk_exceed_error ("load model" , self ._max_chunk_count + 1 ),
132- )
133- self .send_chunked_request (
134- unload_request_header , self ._max_chunk_count , self ._parse_error
135- )
136- self .send_chunked_request (
137- unload_request_header ,
138- self ._max_chunk_count + 1 ,
139- self .chunk_exceed_error ("unload model" , self ._max_chunk_count + 1 ),
112+ unload_request_header , self ._malicious_chunk_count , self ._parse_error
140113 )
141114
142115 def test_trace (self ):
@@ -145,59 +118,34 @@ def test_trace(self):
145118 )
146119
147120 self .send_chunked_request (
148- request_header , self ._max_chunk_count , self ._parse_error
149- )
150- self .send_chunked_request (
151- request_header ,
152- self ._max_chunk_count + 1 ,
153- self .chunk_exceed_error ("trace" , self ._max_chunk_count + 1 ),
121+ request_header , self ._malicious_chunk_count , self ._parse_error
154122 )
155123
156124 def test_logging (self ):
157125 request_header = f"POST /v2/logging HTTP/1.1\r \n "
158126
159127 self .send_chunked_request (
160- request_header , self ._max_chunk_count , self ._parse_error
161- )
162- self .send_chunked_request (
163- request_header ,
164- self ._max_chunk_count + 1 ,
165- self .chunk_exceed_error ("dynamic logging" , self ._max_chunk_count + 1 ),
128+ request_header , self ._malicious_chunk_count , self ._parse_error
166129 )
167130
168131 def test_system_shm_register (self ):
169132 request_header = f"POST /v2/systemsharedmemory/region/test_system_shm_register/register HTTP/1.1\r \n "
170133
171134 self .send_chunked_request (
172- request_header , self ._max_chunk_count , self ._parse_error
173- )
174- self .send_chunked_request (
175- request_header ,
176- self ._max_chunk_count + 1 ,
177- self .chunk_exceed_error ("register" , self ._max_chunk_count + 1 ),
135+ request_header , self ._malicious_chunk_count , self ._parse_error
178136 )
179137
180138 def test_cuda_shm_register (self ):
181139 request_header = f"POST /v2/cudasharedmemory/region/test_cuda_shm_register/register HTTP/1.1\r \n "
182140
183141 self .send_chunked_request (
184- request_header , self ._max_chunk_count , self ._parse_error
185- )
186- self .send_chunked_request (
187- request_header ,
188- self ._max_chunk_count + 1 ,
189- self .chunk_exceed_error ("register" , self ._max_chunk_count + 1 ),
142+ request_header , self ._malicious_chunk_count , self ._parse_error
190143 )
191144
192145 def test_generate (self ):
193146 request_header = f"POST /v2/models/{ self ._model_name } /generate HTTP/1.1\r \n "
194147 self .send_chunked_request (
195- request_header , self ._max_chunk_count , self ._parse_error
196- )
197- self .send_chunked_request (
198- request_header ,
199- self ._max_chunk_count + 1 ,
200- self .chunk_exceed_error ("generate" , self ._max_chunk_count + 1 ),
148+ request_header , self ._malicious_chunk_count , self ._parse_error
201149 )
202150
203151
0 commit comments