Skip to content

Commit 40e7ce6

Browse files
authored
Reject truncated Thrift containers before allocation (#18631)
* Harden Thrift container deserialization limits * Add generated Thrift request regression test * Validate Thrift reads against remaining frame data
1 parent 5756dc8 commit 40e7ce6

4 files changed

Lines changed: 302 additions & 12 deletions

File tree

‎iotdb-client/service-rpc/src/main/i18n/en/org/apache/iotdb/rpc/i18n/RpcMessages.java‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -35,6 +35,9 @@ public final class RpcMessages {
3535
"Frame size (%d) larger than protect max size (%d)%s!";
3636
public static final String FRAME_ERROR_STRING_LENGTH_EXCEEDED =
3737
"String length (%d) larger than protect max size (%d)%s!";
38+
public static final String
39+
EXCEPTION_REQUIRED_READ_SIZE_ARG_EXCEEDS_REMAINING_FRAME_SIZE_ARG_ARG_9C0541EE =
40+
"Required read size (%d) exceeds remaining frame size (%d)%s!";
3841

3942
// TElasticFramedTransport - SSL
4043
public static final String NON_SSL_TO_SSL_PORT =

‎iotdb-client/service-rpc/src/main/i18n/zh/org/apache/iotdb/rpc/i18n/RpcMessages.java‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,9 @@ public final class RpcMessages {
3232
"帧大小 (%d) 超过保护最大值 (%d)%s!";
3333
public static final String FRAME_ERROR_STRING_LENGTH_EXCEEDED =
3434
"字符串长度 (%d) 超过保护最大值 (%d)%s!";
35+
public static final String
36+
EXCEPTION_REQUIRED_READ_SIZE_ARG_EXCEEDS_REMAINING_FRAME_SIZE_ARG_ARG_9C0541EE =
37+
"请求读取的大小 (%d) 超过当前帧剩余大小 (%d)%s!";
3538

3639
// TElasticFramedTransport - SSL
3740
public static final String NON_SSL_TO_SSL_PORT =

‎iotdb-client/service-rpc/src/main/java/org/apache/iotdb/rpc/TElasticFramedTransport.java‎

Lines changed: 30 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -245,7 +245,9 @@ private enum FrameError {
245245
TLS_REQUEST(RpcMessages.FRAME_ERROR_TLS_REQUEST),
246246
NEGATIVE_FRAME_SIZE(RpcMessages.FRAME_ERROR_NEGATIVE_FRAME_SIZE),
247247
FRAME_SIZE_EXCEEDED(RpcMessages.FRAME_ERROR_FRAME_SIZE_EXCEEDED),
248-
STRING_LENGTH_EXCEEDED(RpcMessages.FRAME_ERROR_STRING_LENGTH_EXCEEDED);
248+
STRING_LENGTH_EXCEEDED(RpcMessages.FRAME_ERROR_STRING_LENGTH_EXCEEDED),
249+
INSUFFICIENT_FRAME_DATA(
250+
RpcMessages.EXCEPTION_REQUIRED_READ_SIZE_ARG_EXCEEDS_REMAINING_FRAME_SIZE_ARG_ARG_9C0541EE);
249251

250252
private final String messageFormat;
251253

@@ -255,7 +257,9 @@ private enum FrameError {
255257

256258
void throwException(long size, String remoteInfo, int maxSize) throws TTransportException {
257259
String message =
258-
(this == FRAME_SIZE_EXCEEDED || this == STRING_LENGTH_EXCEEDED)
260+
(this == FRAME_SIZE_EXCEEDED
261+
|| this == STRING_LENGTH_EXCEEDED
262+
|| this == INSUFFICIENT_FRAME_DATA)
259263
? String.format(messageFormat, size, maxSize, remoteInfo)
260264
: String.format(messageFormat, size, remoteInfo);
261265
throw new TTransportException(TTransportException.CORRUPTED_DATA, message);
@@ -308,18 +312,32 @@ public void updateKnownMessageSize(long size) throws TTransportException {
308312

309313
@Override
310314
public void checkReadBytesAvailable(long numBytes) throws TTransportException {
315+
// RPC messages are flushed as complete frames. Container checks pass their minimum encoded
316+
// size here, before generated code allocates the container. Compare it with actual buffered
317+
// data, not just the configured frame cap. Use readBuffer directly because copyBinary makes
318+
// this transport's getBytesRemainingInBuffer() return -1.
319+
int remaining = readBuffer.getBytesRemainingInBuffer();
320+
FrameError error;
321+
int limit;
311322
if (numBytes >= thriftMaxFrameSize) {
312-
SocketAddress remoteAddress = null;
313-
if (underlying instanceof TSocket) {
314-
remoteAddress = ((TSocket) underlying).getSocket().getRemoteSocketAddress();
315-
}
316-
String remoteInfo =
317-
(remoteAddress == null)
318-
? RpcMessages.EMPTY_MESSAGE
319-
: RpcMessages.REMOTE_ADDRESS_PREFIX + remoteAddress;
320-
close();
321-
FrameError.STRING_LENGTH_EXCEEDED.throwException(numBytes, remoteInfo, thriftMaxFrameSize);
323+
error = FrameError.STRING_LENGTH_EXCEEDED;
324+
limit = thriftMaxFrameSize;
325+
} else if (numBytes > remaining) {
326+
error = FrameError.INSUFFICIENT_FRAME_DATA;
327+
limit = remaining;
328+
} else {
329+
return;
330+
}
331+
SocketAddress remoteAddress = null;
332+
if (underlying instanceof TSocket) {
333+
remoteAddress = ((TSocket) underlying).getSocket().getRemoteSocketAddress();
322334
}
335+
String remoteInfo =
336+
(remoteAddress == null)
337+
? RpcMessages.EMPTY_MESSAGE
338+
: RpcMessages.REMOTE_ADDRESS_PREFIX + remoteAddress;
339+
close();
340+
error.throwException(numBytes, remoteInfo, limit);
323341
}
324342

325343
@Override
Lines changed: 266 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,266 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one
3+
* or more contributor license agreements. See the NOTICE file
4+
* distributed with this work for additional information
5+
* regarding copyright ownership. The ASF licenses this file
6+
* to you under the Apache License, Version 2.0 (the
7+
* "License"); you may not use this file except in compliance
8+
* with the License. You may obtain a copy of the License at
9+
*
10+
* http://www.apache.org/licenses/LICENSE-2.0
11+
*
12+
* Unless required by applicable law or agreed to in writing,
13+
* software distributed under the License is distributed on an
14+
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15+
* KIND, either express or implied. See the License for the
16+
* specific language governing permissions and limitations
17+
* under the License.
18+
*/
19+
20+
package org.apache.iotdb.rpc;
21+
22+
import org.apache.iotdb.rpc.i18n.RpcMessages;
23+
import org.apache.iotdb.service.rpc.thrift.IClientRPCService;
24+
import org.apache.iotdb.service.rpc.thrift.TSInsertRecordReq;
25+
26+
import org.apache.thrift.TException;
27+
import org.apache.thrift.protocol.TBinaryProtocol;
28+
import org.apache.thrift.protocol.TCompactProtocol;
29+
import org.apache.thrift.protocol.TField;
30+
import org.apache.thrift.protocol.TList;
31+
import org.apache.thrift.protocol.TMap;
32+
import org.apache.thrift.protocol.TMessage;
33+
import org.apache.thrift.protocol.TMessageType;
34+
import org.apache.thrift.protocol.TProtocol;
35+
import org.apache.thrift.protocol.TSet;
36+
import org.apache.thrift.protocol.TStruct;
37+
import org.apache.thrift.protocol.TType;
38+
import org.apache.thrift.transport.TMemoryBuffer;
39+
import org.apache.thrift.transport.TMemoryInputTransport;
40+
import org.apache.thrift.transport.TTransport;
41+
import org.apache.thrift.transport.TTransportException;
42+
import org.junit.Test;
43+
import org.junit.runner.RunWith;
44+
import org.junit.runners.Parameterized;
45+
46+
import java.nio.ByteBuffer;
47+
import java.util.ArrayList;
48+
import java.util.Arrays;
49+
import java.util.Collection;
50+
import java.util.Collections;
51+
52+
import static org.junit.Assert.assertEquals;
53+
import static org.junit.Assert.assertNotNull;
54+
import static org.junit.Assert.assertNull;
55+
import static org.junit.Assert.assertThrows;
56+
57+
@RunWith(Parameterized.class)
58+
public class TElasticFramedTransportReadTest {
59+
private static final int MAX_FRAME_SIZE = 1024 * 1024;
60+
private static final int ELEMENT_COUNT = 200_001;
61+
62+
private final boolean compact;
63+
private final boolean copyBinary;
64+
private final boolean snappy;
65+
66+
@Parameterized.Parameters(name = "compact={0}, copyBinary={1}, snappy={2}")
67+
public static Collection<Object[]> parameters() {
68+
Collection<Object[]> parameters = new ArrayList<>();
69+
for (boolean compact : new boolean[] {false, true}) {
70+
for (boolean copyBinary : new boolean[] {false, true}) {
71+
for (boolean snappy : new boolean[] {false, true}) {
72+
parameters.add(new Object[] {compact, copyBinary, snappy});
73+
}
74+
}
75+
}
76+
return parameters;
77+
}
78+
79+
public TElasticFramedTransportReadTest(boolean compact, boolean copyBinary, boolean snappy) {
80+
this.compact = compact;
81+
this.copyBinary = copyBinary;
82+
this.snappy = snappy;
83+
}
84+
85+
@Test
86+
public void testTruncatedRequestRejectedBeforeContainerAllocation() throws Exception {
87+
byte[] frame =
88+
serialize(
89+
p -> {
90+
p.writeMessageBegin(new TMessage("insertRecord", TMessageType.CALL, 1));
91+
p.writeStructBegin(new TStruct("insertRecord_args"));
92+
p.writeFieldBegin(new TField("req", TType.STRUCT, (short) 1));
93+
p.writeStructBegin(new TStruct("TSInsertRecordReq"));
94+
p.writeFieldBegin(new TField("columnCategoryies", TType.LIST, (short) 8));
95+
p.writeListBegin(new TList(TType.BYTE, ELEMENT_COUNT));
96+
// Flush a complete frame that ends at the container header, without any elements.
97+
});
98+
try (TElasticFramedTransport transport = transport(new TMemoryInputTransport(frame))) {
99+
TProtocol protocol = protocol(transport);
100+
assertEquals("insertRecord", protocol.readMessageBegin().name);
101+
IClientRPCService.insertRecord_args args = new IClientRPCService.insertRecord_args();
102+
TTransportException exception =
103+
assertThrows(TTransportException.class, () -> args.read(protocol));
104+
assertNotNull(args.req);
105+
assertNull(args.req.columnCategoryies);
106+
assertInsufficientFrame(exception, ELEMENT_COUNT, 0);
107+
}
108+
}
109+
110+
@Test
111+
public void testLargeCompleteRequestsInSuccessiveFrames() throws Exception {
112+
TSInsertRecordReq large =
113+
new TSInsertRecordReq(1, "root.sg.d", Collections.emptyList(), ByteBuffer.allocate(0), 1)
114+
.setColumnCategoryies(Collections.nCopies(ELEMENT_COUNT, (byte) 0));
115+
TSInsertRecordReq empty =
116+
new TSInsertRecordReq(1, "root.sg.d", Collections.emptyList(), ByteBuffer.allocate(0), 2)
117+
.setColumnCategoryies(Collections.emptyList());
118+
TMemoryBuffer wire = new TMemoryBuffer(128);
119+
try (TElasticFramedTransport output = transport(wire)) {
120+
IClientRPCService.Client client = new IClientRPCService.Client(protocol(output));
121+
client.send_insertRecord(large);
122+
client.send_insertRecord(empty);
123+
}
124+
try (TElasticFramedTransport input =
125+
transport(new TMemoryInputTransport(Arrays.copyOf(wire.getArray(), wire.length())))) {
126+
TProtocol protocol = protocol(input);
127+
for (TSInsertRecordReq expected : new TSInsertRecordReq[] {large, empty}) {
128+
assertEquals("insertRecord", protocol.readMessageBegin().name);
129+
IClientRPCService.insertRecord_args args = new IClientRPCService.insertRecord_args();
130+
args.read(protocol);
131+
protocol.readMessageEnd();
132+
assertEquals(expected, args.req);
133+
}
134+
}
135+
}
136+
137+
@Test
138+
public void testTruncatedContainerHeaders() throws Exception {
139+
for (byte elementType : new byte[] {TType.BYTE, TType.STRUCT, TType.LIST, TType.MAP}) {
140+
for (byte containerType : new byte[] {TType.LIST, TType.SET, TType.MAP}) {
141+
byte[] frame =
142+
serialize(
143+
p -> {
144+
if (containerType == TType.LIST) {
145+
p.writeListBegin(new TList(elementType, ELEMENT_COUNT));
146+
} else if (containerType == TType.SET) {
147+
p.writeSetBegin(new TSet(elementType, ELEMENT_COUNT));
148+
} else {
149+
p.writeMapBegin(new TMap(TType.BYTE, elementType, ELEMENT_COUNT));
150+
}
151+
});
152+
try (TElasticFramedTransport input = transport(new TMemoryInputTransport(frame))) {
153+
TProtocol p = protocol(input);
154+
long minimumBytes =
155+
(long) ELEMENT_COUNT
156+
* (p.getMinSerializedSize(elementType) + (containerType == TType.MAP ? 1 : 0));
157+
TTransportException exception =
158+
assertThrows(
159+
TTransportException.class,
160+
() -> {
161+
if (containerType == TType.LIST) {
162+
p.readListBegin();
163+
} else if (containerType == TType.SET) {
164+
p.readSetBegin();
165+
} else {
166+
p.readMapBegin();
167+
}
168+
});
169+
assertInsufficientFrame(exception, minimumBytes, 0);
170+
}
171+
}
172+
}
173+
}
174+
175+
@Test
176+
public void testCompleteBinaryAtFrameBoundary() throws Exception {
177+
ByteBuffer binary = ByteBuffer.wrap(new byte[ELEMENT_COUNT]);
178+
byte[] frame =
179+
serialize(
180+
p -> {
181+
p.writeString("");
182+
p.writeBinary(binary);
183+
});
184+
try (TElasticFramedTransport input = transport(new TMemoryInputTransport(frame))) {
185+
TProtocol protocol = protocol(input);
186+
assertEquals("", protocol.readString());
187+
assertEquals(binary, protocol.readBinary());
188+
input.checkReadBytesAvailable(0);
189+
}
190+
}
191+
192+
@Test
193+
public void testTruncatedStringAndBinary() throws Exception {
194+
TMemoryBuffer buffer = new TMemoryBuffer(128);
195+
protocol(buffer).writeBinary(ByteBuffer.wrap(new byte[ELEMENT_COUNT]));
196+
int headerLength = buffer.length() - ELEMENT_COUNT;
197+
byte[] frame = serialize(p -> p.getTransport().write(buffer.getArray(), 0, headerLength));
198+
for (boolean binary : new boolean[] {false, true}) {
199+
try (TElasticFramedTransport input = transport(new TMemoryInputTransport(frame))) {
200+
TProtocol p = protocol(input);
201+
TTransportException exception =
202+
assertThrows(
203+
TTransportException.class,
204+
() -> {
205+
if (binary) {
206+
p.readBinary();
207+
} else {
208+
p.readString();
209+
}
210+
});
211+
assertInsufficientFrame(exception, ELEMENT_COUNT, 0);
212+
}
213+
}
214+
}
215+
216+
@Test
217+
public void testExistingMaximumReadSizeProtection() throws Exception {
218+
try (TElasticFramedTransport input = transport(new TMemoryInputTransport(new byte[0]))) {
219+
TTransportException exception =
220+
assertThrows(
221+
TTransportException.class, () -> input.checkReadBytesAvailable(MAX_FRAME_SIZE));
222+
assertEquals(TTransportException.CORRUPTED_DATA, exception.getType());
223+
assertEquals(
224+
String.format(
225+
RpcMessages.FRAME_ERROR_STRING_LENGTH_EXCEEDED, MAX_FRAME_SIZE, MAX_FRAME_SIZE, ""),
226+
exception.getMessage());
227+
}
228+
}
229+
230+
private void assertInsufficientFrame(
231+
TTransportException exception, long required, int remaining) {
232+
assertEquals(TTransportException.CORRUPTED_DATA, exception.getType());
233+
assertEquals(
234+
String.format(
235+
RpcMessages
236+
.EXCEPTION_REQUIRED_READ_SIZE_ARG_EXCEEDS_REMAINING_FRAME_SIZE_ARG_ARG_9C0541EE,
237+
required,
238+
remaining,
239+
""),
240+
exception.getMessage());
241+
}
242+
243+
private TElasticFramedTransport transport(TTransport underlying) throws TTransportException {
244+
return snappy
245+
? new TSnappyElasticFramedTransport(underlying, 128, MAX_FRAME_SIZE, copyBinary)
246+
: new TElasticFramedTransport(underlying, 128, MAX_FRAME_SIZE, copyBinary);
247+
}
248+
249+
private TProtocol protocol(TTransport transport) {
250+
return compact ? new TCompactProtocol(transport) : new TBinaryProtocol(transport);
251+
}
252+
253+
private byte[] serialize(ProtocolWriter writer) throws TException {
254+
TMemoryBuffer wire = new TMemoryBuffer(128);
255+
try (TElasticFramedTransport output = transport(wire)) {
256+
writer.write(protocol(output));
257+
output.flush();
258+
}
259+
return Arrays.copyOf(wire.getArray(), wire.length());
260+
}
261+
262+
@FunctionalInterface
263+
private interface ProtocolWriter {
264+
void write(TProtocol protocol) throws TException;
265+
}
266+
}

0 commit comments

Comments
 (0)