diff --git a/scrooge-core/src/main/scala/com/twitter/scrooge/LazyTProtocol.scala b/scrooge-core/src/main/scala/com/twitter/scrooge/LazyTProtocol.scala index 3951150b..e261cfbb 100644 --- a/scrooge-core/src/main/scala/com/twitter/scrooge/LazyTProtocol.scala +++ b/scrooge-core/src/main/scala/com/twitter/scrooge/LazyTProtocol.scala @@ -123,5 +123,4 @@ trait LazyTProtocol extends TProtocol { * Returns: The offset at which the string can be read. */ def offsetSkipBinary(): Int - } diff --git a/scrooge-core/src/main/scala/com/twitter/scrooge/TLazyBinaryProtocol.scala b/scrooge-core/src/main/scala/com/twitter/scrooge/TLazyBinaryProtocol.scala index 33968fcd..3dbcc37a 100644 --- a/scrooge-core/src/main/scala/com/twitter/scrooge/TLazyBinaryProtocol.scala +++ b/scrooge-core/src/main/scala/com/twitter/scrooge/TLazyBinaryProtocol.scala @@ -17,6 +17,7 @@ import org.apache.thrift.protocol._ object TLazyBinaryProtocol { private val AnonymousStruct: TStruct = new TStruct() private val utf8Charset = Charset.forName("UTF-8") + private val NEW_ENUM_TYPE_ID: Byte = -1 } class TLazyBinaryProtocol(transport: TArrayByteTransport) @@ -30,9 +31,10 @@ class TLazyBinaryProtocol(transport: TArrayByteTransport) } override def writeFieldBegin(field: TField): Unit = { + val typeToWrite = if (field.`type` == TType.ENUM) NEW_ENUM_TYPE_ID else field.`type` val buf = transport.getBuffer(3) val offset = transport.writerOffset - buf(offset) = field.`type` + buf(offset) = typeToWrite innerWriteI16(buf, offset + 1, field.id) } @@ -176,7 +178,8 @@ class TLazyBinaryProtocol(transport: TArrayByteTransport) override def readFieldBegin(): TField = { val tpe: Byte = readByte() val id: Short = if (tpe == TType.STOP) 0 else readI16() - new TField("", tpe, id) + val finalType = if (tpe == NEW_ENUM_TYPE_ID) TType.ENUM else tpe + new TField("", finalType, id) } override def readFieldEnd(): Unit = () diff --git a/scrooge-core/src/test/scala/com/twitter/scrooge/internal/TProtocolsTest.scala b/scrooge-core/src/test/scala/com/twitter/scrooge/internal/TProtocolsTest.scala index f33803db..7ab733ed 100644 --- a/scrooge-core/src/test/scala/com/twitter/scrooge/internal/TProtocolsTest.scala +++ b/scrooge-core/src/test/scala/com/twitter/scrooge/internal/TProtocolsTest.scala @@ -1,8 +1,6 @@ package com.twitter.scrooge.internal -import com.twitter.scrooge.TArrayByteTransport -import com.twitter.scrooge.TFieldBlob -import com.twitter.scrooge.ThriftUnion +import com.twitter.scrooge.{TArrayByteTransport, TFieldBlob, TLazyBinaryProtocol, ThriftUnion} import com.twitter.util.mock.Mockito import org.apache.thrift.protocol.TBinaryProtocol import org.apache.thrift.protocol.TCompactProtocol @@ -10,6 +8,7 @@ import org.apache.thrift.protocol.TField import org.apache.thrift.protocol.TProtocolException import org.apache.thrift.protocol.TType import org.apache.thrift.transport.TMemoryBuffer + import scala.collection.immutable import org.scalatest.funsuite.AnyFunSuite @@ -173,4 +172,44 @@ class TProtocolsTest extends AnyFunSuite with Mockito { succeed } + test("readFieldBegin handles new ENUM type identifier") { + val writeBuffer = new TMemoryBuffer(128) + val writeProto = new TBinaryProtocol(writeBuffer) + + val newEnum = -1 + + writeProto.writeByte(newEnum.toByte) // New ENUM type identifier + writeProto.writeI16(1) // Field ID + writeProto.writeI32(2) // Enum value + + val readBuffer = TArrayByteTransport(writeBuffer.getArray) + val readProto = new TLazyBinaryProtocol(readBuffer) + + val field = readProto.readFieldBegin() + assert(field.`type` == TType.ENUM) + assert(field.id == 1) + + val enumValue = readProto.readI32() + assert(enumValue == 2) + } + + test("readFieldBegin handles old ENUM type identifier") { + val writeBuffer = new TMemoryBuffer(128) + val writeProto = new TBinaryProtocol(writeBuffer) + + writeProto.writeByte(TType.ENUM) + writeProto.writeI16(1) + writeProto.writeI32(2) + + val readBuffer = TArrayByteTransport(writeBuffer.getArray) + val readProto = new TLazyBinaryProtocol(readBuffer) + + val field = readProto.readFieldBegin() + assert(field.`type` == TType.ENUM) + assert(field.id == 1) + + val enumValue = readProto.readI32() + assert(enumValue == 2) + } + }