Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -576,7 +576,10 @@ class GroupMetadataManager(brokerId: Int,
}
}

private def doLoadGroupsAndOffsets(topicPartition: TopicPartition, onGroupLoaded: GroupMetadata => Unit): Unit = {
// Visible for testing
private[group] def doLoadGroupsAndOffsets(topicPartition: TopicPartition,
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
onGroupLoaded: GroupMetadata => Unit): Unit = {

def logEndOffset: Long = replicaManager.getLogEndOffset(topicPartition).getOrElse(-1L)

replicaManager.getLog(topicPartition) match {
Expand Down Expand Up @@ -651,7 +654,6 @@ class GroupMetadataManager(brokerId: Int,
if (batchBaseOffset.isEmpty)
batchBaseOffset = Some(record.offset)
GroupMetadataManager.readMessageKey(record.key) match {

case offsetKey: OffsetKey =>
if (isTxnOffsetCommit && !pendingOffsets.contains(batch.producerId))
pendingOffsets.put(batch.producerId, mutable.Map[GroupTopicPartition, CommitRecordMetadataAndOffset]())
Expand Down Expand Up @@ -683,8 +685,14 @@ class GroupMetadataManager(brokerId: Int,
removedGroups.add(groupId)
}

case unknownKey =>
throw new IllegalStateException(s"Unexpected message key $unknownKey while loading offsets and group metadata")
case unknownKey: UnknownKey =>
// Unknown versions may exist when a downgraded coordinator is reading records from the log.
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
warn(s"Unknown message key with version ${unknownKey.version}" +
s" while loading offsets and group metadata. Ignoring it. " +
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
s"It could be a left over from an aborted upgrade.")
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated

case unexpectedKey =>
throw new IllegalStateException(s"Unexpected message key $unexpectedKey while loading offsets and group metadata")
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
}
}
}
Expand Down Expand Up @@ -1155,7 +1163,9 @@ object GroupMetadataManager {
// version 2 refers to group metadata
val key = new GroupMetadataKeyData(new ByteBufferAccessor(buffer), version)
GroupMetadataKey(version, key.group)
} else throw new IllegalStateException(s"Unknown group metadata message version: $version")
} else {
UnknownKey(version)
}
}

/**
Expand Down Expand Up @@ -1273,9 +1283,10 @@ object GroupMetadataManager {
throw new KafkaException("Failed to decode message using offset topic decoder (message had a missing key)")
} else {
GroupMetadataManager.readMessageKey(record.key) match {
case offsetKey: OffsetKey => parseOffsets(offsetKey, record.value)
case groupMetadataKey: GroupMetadataKey => parseGroupMetadata(groupMetadataKey, record.value)
case _ => throw new KafkaException("Failed to decode message using offset topic decoder (message had an invalid key)")
case offsetKey: OffsetKey => parseOffsets(offsetKey, record.value)
case groupMetadataKey: GroupMetadataKey => parseGroupMetadata(groupMetadataKey, record.value)
case unknownKey: UnknownKey => (Some(s"UNKNOWN(version=${unknownKey.version})"), None)
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
case _ => throw new KafkaException("Failed to decode message using offset topic decoder (message had an invalid key)")
}
}
}
Expand Down Expand Up @@ -1359,12 +1370,14 @@ trait BaseKey{
}

case class OffsetKey(version: Short, key: GroupTopicPartition) extends BaseKey {

override def toString: String = key.toString
}

case class GroupMetadataKey(version: Short, key: String) extends BaseKey {

override def toString: String = key
}

case class UnknownKey(version: Short) extends BaseKey {
override def key: String = null
override def toString: String = key
}
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@ package kafka.coordinator.transaction
import java.io.PrintStream
import java.nio.ByteBuffer
import java.nio.charset.StandardCharsets

import kafka.internals.generated.{TransactionLogKey, TransactionLogValue}
import org.apache.kafka.clients.consumer.ConsumerRecord
import org.apache.kafka.common.protocol.{ByteBufferAccessor, MessageUtil}
Expand Down Expand Up @@ -98,15 +97,17 @@ object TransactionLog {
*
* @return the key
*/
def readTxnRecordKey(buffer: ByteBuffer): TxnKey = {
def readTxnRecordKey(buffer: ByteBuffer): BaseKey = {
val version = buffer.getShort
if (version >= TransactionLogKey.LOWEST_SUPPORTED_VERSION && version <= TransactionLogKey.HIGHEST_SUPPORTED_VERSION) {
val value = new TransactionLogKey(new ByteBufferAccessor(buffer), version)
TxnKey(
version = version,
transactionalId = value.transactionalId
)
} else throw new IllegalStateException(s"Unknown version $version from the transaction log message")
} else {
UnknownKey(version)
}
}

/**
Expand Down Expand Up @@ -148,17 +149,23 @@ object TransactionLog {
// Formatter for use with tools to read transaction log messages
class TransactionLogMessageFormatter extends MessageFormatter {
def writeTo(consumerRecord: ConsumerRecord[Array[Byte], Array[Byte]], output: PrintStream): Unit = {
Option(consumerRecord.key).map(key => readTxnRecordKey(ByteBuffer.wrap(key))).foreach { txnKey =>
val transactionalId = txnKey.transactionalId
val value = consumerRecord.value
val producerIdMetadata = if (value == null)
None
else
readTxnRecordValue(transactionalId, ByteBuffer.wrap(value))
output.write(transactionalId.getBytes(StandardCharsets.UTF_8))
output.write("::".getBytes(StandardCharsets.UTF_8))
output.write(producerIdMetadata.getOrElse("NULL").toString.getBytes(StandardCharsets.UTF_8))
output.write("\n".getBytes(StandardCharsets.UTF_8))
Option(consumerRecord.key).map(key => readTxnRecordKey(ByteBuffer.wrap(key))).foreach {
case txnKey: TxnKey =>
val transactionalId = txnKey.transactionalId
val value = consumerRecord.value
val producerIdMetadata = if (value == null)
None
else
readTxnRecordValue(transactionalId, ByteBuffer.wrap(value))
output.write(transactionalId.getBytes(StandardCharsets.UTF_8))
output.write("::".getBytes(StandardCharsets.UTF_8))
output.write(producerIdMetadata.getOrElse("NULL").toString.getBytes(StandardCharsets.UTF_8))
output.write("\n".getBytes(StandardCharsets.UTF_8))

case _: UnknownKey => // Only print if this message is a transaction record
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated

case unexpectedKey =>
throw new IllegalStateException(s"Found unexpected key $unexpectedKey while reading transaction log.")
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
Comment thread
dajac marked this conversation as resolved.
Outdated
}
}
}
Expand All @@ -167,25 +174,44 @@ object TransactionLog {
* Exposed for printing records using [[kafka.tools.DumpLogSegments]]
*/
def formatRecordKeyAndValue(record: Record): (Option[String], Option[String]) = {
val txnKey = TransactionLog.readTxnRecordKey(record.key)
val keyString = s"transaction_metadata::transactionalId=${txnKey.transactionalId}"

val valueString = TransactionLog.readTxnRecordValue(txnKey.transactionalId, record.value) match {
case None => "<DELETE>"

case Some(txnMetadata) => s"producerId:${txnMetadata.producerId}," +
s"producerEpoch:${txnMetadata.producerEpoch}," +
s"state=${txnMetadata.state}," +
s"partitions=${txnMetadata.topicPartitions.mkString("[", ",", "]")}," +
s"txnLastUpdateTimestamp=${txnMetadata.txnLastUpdateTimestamp}," +
s"txnTimeoutMs=${txnMetadata.txnTimeoutMs}"
}
TransactionLog.readTxnRecordKey(record.key) match {
case txnKey: TxnKey =>
val keyString = s"transaction_metadata::transactionalId=${txnKey.transactionalId}"

val valueString = TransactionLog.readTxnRecordValue(txnKey.transactionalId, record.value) match {
case None => "<DELETE>"

case Some(txnMetadata) => s"producerId:${txnMetadata.producerId}," +
s"producerEpoch:${txnMetadata.producerEpoch}," +
s"state=${txnMetadata.state}," +
s"partitions=${txnMetadata.topicPartitions.mkString("[", ",", "]")}," +
s"txnLastUpdateTimestamp=${txnMetadata.txnLastUpdateTimestamp}," +
s"txnTimeoutMs=${txnMetadata.txnTimeoutMs}"
}

(Some(keyString), Some(valueString))

case _: UnknownKey =>
(Some("<UNKNOWN>"), Some("<UNKNOWN>"))
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated

(Some(keyString), Some(valueString))
case unexpectedKey =>
throw new IllegalStateException(s"Found unexpected key $unexpectedKey while formatting transaction log.")
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
}
}

}

case class TxnKey(version: Short, transactionalId: String) {
trait BaseKey{
def version: Short
def transactionalId: String
}

case class TxnKey(version: Short, transactionalId: String) extends BaseKey {
override def toString: String = transactionalId
}

case class UnknownKey(version: Short) extends BaseKey {
override def transactionalId: String = null
override def toString: String = transactionalId
}

Original file line number Diff line number Diff line change
Expand Up @@ -467,16 +467,26 @@ class TransactionStateManager(brokerId: Int,
memRecords.batches.forEach { batch =>
for (record <- batch.asScala) {
require(record.hasKey, "Transaction state log's key should not be null")
val txnKey = TransactionLog.readTxnRecordKey(record.key)
// load transaction metadata along with transaction state
val transactionalId = txnKey.transactionalId
TransactionLog.readTxnRecordValue(transactionalId, record.value) match {
case None =>
loadedTransactions.remove(transactionalId)
case Some(txnMetadata) =>
loadedTransactions.put(transactionalId, txnMetadata)
TransactionLog.readTxnRecordKey(record.key) match {
case txnKey: TxnKey =>
// load transaction metadata along with transaction state
val transactionalId = txnKey.transactionalId
TransactionLog.readTxnRecordValue(transactionalId, record.value) match {
case None =>
loadedTransactions.remove(transactionalId)
case Some(txnMetadata) =>
loadedTransactions.put(transactionalId, txnMetadata)
}
currOffset = batch.nextOffset

case unknownKey: UnknownKey =>
warn(s"Unknown message key with version ${unknownKey.version}" +
s" while loading transaction state. Ignoring it. " +
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
s"It could be a left over from an aborted upgrade.")
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated

case unexpectedKey =>
throw new IllegalStateException(s"Found unexpected key $unexpectedKey while reading transaction log.")
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
}
currOffset = batch.nextOffset
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@ import org.apache.kafka.clients.consumer.internals.ConsumerProtocol
import org.apache.kafka.common.{TopicIdPartition, TopicPartition, Uuid}
import org.apache.kafka.common.internals.Topic
import org.apache.kafka.common.metrics.{JmxReporter, KafkaMetricsContext, Metrics => kMetrics}
import org.apache.kafka.common.protocol.Errors
import org.apache.kafka.common.protocol.{Errors, MessageUtil}
import org.apache.kafka.common.record._
import org.apache.kafka.common.requests.OffsetFetchResponse
import org.apache.kafka.common.requests.ProduceResponse.PartitionResponse
Expand Down Expand Up @@ -640,8 +640,14 @@ class GroupMetadataManagerTest {
val offsetCommitRecords = createCommittedOffsetRecords(committedOffsets)
val memberId = "98098230493"
val groupMetadataRecord = buildStableGroupRecordWithMember(generation, protocolType, protocol, memberId)

// Should ignore unknown record
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
val unknownKey = new org.apache.kafka.coordinator.group.generated.GroupMetadataKey()
val unknownMessage = MessageUtil.toVersionPrefixedBytes(Short.MaxValue, unknownKey)
val unknownRecord = new SimpleRecord(unknownMessage, unknownMessage)

val records = MemoryRecords.withRecords(startOffset, CompressionType.NONE,
(offsetCommitRecords ++ Seq(groupMetadataRecord)).toArray: _*)
(offsetCommitRecords ++ Seq(unknownRecord) ++ Seq(groupMetadataRecord)).toArray: _*)
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated

expectGroupMetadataLoad(groupMetadataTopicPartition, startOffset, records)

Expand Down Expand Up @@ -2762,4 +2768,13 @@ class GroupMetadataManagerTest {
assertTrue(partitionLoadTime("partition-load-time-max") >= diff)
assertTrue(partitionLoadTime("partition-load-time-avg") >= diff)
}

@Test
def testIgnoreUnknownMessageKeyVersion(): Unit = {
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
val record = new org.apache.kafka.coordinator.group.generated.GroupMetadataKey()
val unknownRecord = MessageUtil.toVersionPrefixedBytes(Short.MaxValue, record)
val key = GroupMetadataManager.readMessageKey(ByteBuffer.wrap(unknownRecord))
assertEquals(UnknownKey(Short.MaxValue), key)
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -17,12 +17,15 @@
package kafka.coordinator.transaction


import kafka.internals.generated.TransactionLogKey
import kafka.utils.TestUtils
import org.apache.kafka.common.TopicPartition
import org.apache.kafka.common.protocol.MessageUtil
import org.apache.kafka.common.record.{CompressionType, MemoryRecords, SimpleRecord}
import org.junit.jupiter.api.Assertions.{assertEquals, assertThrows}
import org.junit.jupiter.api.Test

import java.nio.ByteBuffer
import scala.jdk.CollectionConverters._

class TransactionLogTest {
Expand Down Expand Up @@ -135,4 +138,12 @@ class TransactionLogTest {
assertEquals(Some("<DELETE>"), valueStringOpt)
}

@Test
def testReadUnknownMessageKeyVersion(): Unit = {
Comment thread
jeffkbkim marked this conversation as resolved.
Outdated
val record = new TransactionLogKey()
val unknownRecord = MessageUtil.toVersionPrefixedBytes(Short.MaxValue, record)
val key = TransactionLog.readTxnRecordKey(ByteBuffer.wrap(unknownRecord))
assertEquals(UnknownKey(Short.MaxValue), key)
}

}
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@
*/
package kafka.coordinator.transaction

import kafka.internals.generated.TransactionLogKey

import java.lang.management.ManagementFactory
import java.nio.ByteBuffer
import java.util.concurrent.CountDownLatch
Expand All @@ -28,7 +30,7 @@ import kafka.zk.KafkaZkClient
import org.apache.kafka.common.TopicPartition
import org.apache.kafka.common.internals.Topic.TRANSACTION_STATE_TOPIC_NAME
import org.apache.kafka.common.metrics.{JmxReporter, KafkaMetricsContext, Metrics}
import org.apache.kafka.common.protocol.Errors
import org.apache.kafka.common.protocol.{Errors, MessageUtil}
import org.apache.kafka.common.record._
import org.apache.kafka.common.requests.ProduceResponse.PartitionResponse
import org.apache.kafka.common.requests.TransactionResult
Expand Down Expand Up @@ -1085,4 +1087,41 @@ class TransactionStateManagerTest {
assertTrue(partitionLoadTime("partition-load-time-max") >= 0)
assertTrue(partitionLoadTime( "partition-load-time-avg") >= 0)
}

Comment thread
jeffkbkim marked this conversation as resolved.

@Test
def testIgnoreUnknownRecordType(): Unit = {
txnMetadata1.state = PrepareCommit
txnMetadata1.addPartitions(Set[TopicPartition](new TopicPartition("topic1", 0),
new TopicPartition("topic1", 1)))

txnRecords += new SimpleRecord(txnMessageKeyBytes1, TransactionLog.valueToBytes(txnMetadata1.prepareNoTransit()))
val startOffset = 0L

val unknownKey = new TransactionLogKey()
val unknownMessage = MessageUtil.toVersionPrefixedBytes(Short.MaxValue, unknownKey)
val unknownRecord = new SimpleRecord(unknownMessage, unknownMessage)

val records = MemoryRecords.withRecords(startOffset, CompressionType.NONE,
(Seq(unknownRecord) ++ txnRecords).toArray: _*)

prepareTxnLog(topicPartition, 0, records)

transactionManager.loadTransactionsForTxnTopicPartition(partitionId, coordinatorEpoch = 1, (_, _, _, _) => ())
assertEquals(0, transactionManager.loadingPartitions.size)
assertTrue(transactionManager.transactionMetadataCache.contains(partitionId))
val txnMetadataPool = transactionManager.transactionMetadataCache(partitionId).metadataPerTransactionalId
assertFalse(txnMetadataPool.isEmpty)
assertTrue(txnMetadataPool.contains(transactionalId1))
val txnMetadata = txnMetadataPool.get(transactionalId1)
assertEquals(txnMetadata1.transactionalId, txnMetadata.transactionalId)
assertEquals(txnMetadata1.producerId, txnMetadata.producerId)
assertEquals(txnMetadata1.lastProducerId, txnMetadata.lastProducerId)
assertEquals(txnMetadata1.producerEpoch, txnMetadata.producerEpoch)
assertEquals(txnMetadata1.lastProducerEpoch, txnMetadata.lastProducerEpoch)
assertEquals(txnMetadata1.txnTimeoutMs, txnMetadata.txnTimeoutMs)
assertEquals(txnMetadata1.state, txnMetadata.state)
assertEquals(txnMetadata1.topicPartitions, txnMetadata.topicPartitions)
assertEquals(1, transactionManager.transactionMetadataCache(partitionId).coordinatorEpoch)
}
}