From 4f760d62d9f7697c603feb684028785dd8e5a06f Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 3 Feb 2023 11:57:32 -0800 Subject: [PATCH 01/17] basic batching --- .../requests/AddPartitionsToTxnRequest.java | 135 +++++++++++++++--- .../requests/AddPartitionsToTxnResponse.java | 69 +++++++-- .../message/AddPartitionsToTxnRequest.json | 30 +++- .../message/AddPartitionsToTxnResponse.json | 33 +++-- .../AddPartitionsToTxnRequestTest.java | 69 +++++++-- .../AddPartitionsToTxnResponseTest.java | 47 ++++-- .../transaction/TransactionCoordinator.scala | 64 +++++++++ .../main/scala/kafka/server/KafkaApis.scala | 71 +++++++++ .../AddPartitionsToTxnRequestServerTest.scala | 71 ++++++--- 9 files changed, 507 insertions(+), 82 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 1034c0f7adc55..8ca6acb061d92 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -19,7 +19,16 @@ import org.apache.kafka.common.TopicPartition; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopic; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResult; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResultCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResult; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResultCollection; import org.apache.kafka.common.protocol.ApiKeys; import org.apache.kafka.common.protocol.ByteBufferAccessor; import org.apache.kafka.common.protocol.Errors; @@ -35,21 +44,44 @@ public class AddPartitionsToTxnRequest extends AbstractRequest { private final AddPartitionsToTxnRequestData data; private List cachedPartitions = null; + + private Map> cachedPartitionsByTransaction = null; + + private final short version; public static class Builder extends AbstractRequest.Builder { public final AddPartitionsToTxnRequestData data; + public final boolean isClientRequest; - public Builder(final AddPartitionsToTxnRequestData data) { + public Builder(String transactionalId, + long producerId, + short producerEpoch, + List partitions) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); - this.data = data; + this.isClientRequest = true; + + AddPartitionsToTxnTopicCollection topics = compileTopics(partitions); + + this.data = new AddPartitionsToTxnRequestData() + .setTransactionalId(transactionalId) + .setProducerId(producerId) + .setProducerEpoch(producerEpoch) + .setTopics(topics); } - public Builder(final String transactionalId, - final long producerId, - final short producerEpoch, - final List partitions) { + public Builder(AddPartitionsToTxnTransactionCollection transactions, + boolean verifyOnly) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); + this.isClientRequest = false; + + List transactionsList = new ArrayList<>(); + + this.data = new AddPartitionsToTxnRequestData() + .setTransactions(transactions) + .setVerifyOnly(verifyOnly); + } + private AddPartitionsToTxnTopicCollection compileTopics(final List partitions) { Map> partitionMap = new HashMap<>(); for (TopicPartition topicPartition : partitions) { String topicName = topicPartition.topic(); @@ -66,24 +98,29 @@ public Builder(final String transactionalId, AddPartitionsToTxnTopicCollection topics = new AddPartitionsToTxnTopicCollection(); for (Map.Entry> partitionEntry : partitionMap.entrySet()) { topics.add(new AddPartitionsToTxnTopic() - .setName(partitionEntry.getKey()) - .setPartitions(partitionEntry.getValue())); + .setName(partitionEntry.getKey()) + .setPartitions(partitionEntry.getValue())); } - - this.data = new AddPartitionsToTxnRequestData() - .setTransactionalId(transactionalId) - .setProducerId(producerId) - .setProducerEpoch(producerEpoch) - .setTopics(topics); + return topics; } @Override public AddPartitionsToTxnRequest build(short version) { - return new AddPartitionsToTxnRequest(data, version); + short clampedVersion = (isClientRequest && version > 3) ? 3 : version; + return new AddPartitionsToTxnRequest(data, clampedVersion); } static List getPartitions(AddPartitionsToTxnRequestData data) { List partitions = new ArrayList<>(); + for (AddPartitionsToTxnTransaction transaction : data.transactions()) { + for (AddPartitionsToTxnTopic topicCollection : transaction.topics()) { + for (Integer partition : topicCollection.partitions()) { + partitions.add(new TopicPartition(topicCollection.name(), partition)); + } + } + } + + // Add singleton topics for (AddPartitionsToTxnTopic topicCollection : data.topics()) { for (Integer partition : topicCollection.partitions()) { partitions.add(new TopicPartition(topicCollection.name(), partition)); @@ -101,6 +138,7 @@ public String toString() { public AddPartitionsToTxnRequest(final AddPartitionsToTxnRequestData data, short version) { super(ApiKeys.ADD_PARTITIONS_TO_TXN, version); this.data = data; + this.version = version; } public List partitions() { @@ -110,6 +148,35 @@ public List partitions() { cachedPartitions = Builder.getPartitions(data); return cachedPartitions; } + + public List partitionsForTransaction(String transaction) { + if (cachedPartitionsByTransaction == null) { + cachedPartitionsByTransaction = new HashMap<>(); + } + + return cachedPartitionsByTransaction.computeIfAbsent(transaction, txn -> { + List partitions = new ArrayList<>(); + for (AddPartitionsToTxnTopic topicCollection : data.transactions().find(txn).topics()) { + for (Integer partition : topicCollection.partitions()) { + partitions.add(new TopicPartition(topicCollection.name(), partition)); + } + } + return partitions; + }); + } + + public Map> partitionsByTransaction() { + if (cachedPartitionsByTransaction != null && cachedPartitionsByTransaction.size() == data.transactions().size()) { + return cachedPartitionsByTransaction; + } + + for (AddPartitionsToTxnTransaction transaction : data.transactions()) { + if (cachedPartitionsByTransaction == null || !cachedPartitionsByTransaction.containsKey(transaction.transactionalId())) { + partitionsForTransaction(transaction.transactionalId()); + } + } + return cachedPartitionsByTransaction; + } @Override public AddPartitionsToTxnRequestData data() { @@ -118,11 +185,41 @@ public AddPartitionsToTxnRequestData data() { @Override public AddPartitionsToTxnResponse getErrorResponse(int throttleTimeMs, Throwable e) { - final HashMap errors = new HashMap<>(); - for (TopicPartition partition : partitions()) { - errors.put(partition, Errors.forException(e)); + Errors error = Errors.forException(e); + if (version < 4) { + final HashMap errors = new HashMap<>(); + for (TopicPartition partition : partitions()) { + errors.put(partition, error); + } + return new AddPartitionsToTxnResponse(throttleTimeMs, errors); + } else { + AddPartitionsToTxnResponseData response = new AddPartitionsToTxnResponseData(); + AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); + for (AddPartitionsToTxnTransaction transaction : data().transactions()) { + results.add(errorResponseForTransaction(transaction.transactionalId(), error)); + } + response.setResultsByTransaction(results); + response.setThrottleTimeMs(throttleTimeMs); + return new AddPartitionsToTxnResponse(response); + } + } + + public AddPartitionsToTxnResult errorResponseForTransaction(String transactionalId, Errors e) { + AddPartitionsToTxnResult txnResult = new AddPartitionsToTxnResult().setTransactionalId(transactionalId); + AddPartitionsToTxnTopicResultCollection topicResults = new AddPartitionsToTxnTopicResultCollection(); + for (AddPartitionsToTxnTopic topic : data.transactions().find(transactionalId).topics()) { + AddPartitionsToTxnTopicResult topicResult = new AddPartitionsToTxnTopicResult().setName(topic.name()); + AddPartitionsToTxnPartitionResultCollection partitionResult = new AddPartitionsToTxnPartitionResultCollection(); + for (Integer partition : topic.partitions()) { + partitionResult.add(new AddPartitionsToTxnPartitionResult() + .setPartitionIndex(partition) + .setErrorCode(e.code())); + } + topicResult.setResults(partitionResult); + topicResults.add(topicResult); } - return new AddPartitionsToTxnResponse(throttleTimeMs, errors); + txnResult.setTopicResults(topicResults); + return txnResult; } public static AddPartitionsToTxnRequest parse(ByteBuffer buffer, short version) { diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index 8038f4b8fc66d..9831eabef6989 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -18,6 +18,7 @@ import org.apache.kafka.common.TopicPartition; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResult; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResultCollection; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResult; @@ -27,7 +28,9 @@ import org.apache.kafka.common.protocol.Errors; import java.nio.ByteBuffer; +import java.util.ArrayList; import java.util.HashMap; +import java.util.List; import java.util.Map; /** @@ -49,6 +52,8 @@ public class AddPartitionsToTxnResponse extends AbstractResponse { private final AddPartitionsToTxnResponseData data; private Map cachedErrorsMap = null; + + private Map> cachedAllErrorsMap = null; public AddPartitionsToTxnResponse(AddPartitionsToTxnResponseData data) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); @@ -58,19 +63,25 @@ public AddPartitionsToTxnResponse(AddPartitionsToTxnResponseData data) { public AddPartitionsToTxnResponse(int throttleTimeMs, Map errors) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); + this.data = new AddPartitionsToTxnResponseData() + .setThrottleTimeMs(throttleTimeMs) + .setResults(topicCollectionForErrors(errors)); + } + + private static AddPartitionsToTxnTopicResultCollection topicCollectionForErrors(Map errors) { Map resultMap = new HashMap<>(); - + for (Map.Entry entry : errors.entrySet()) { TopicPartition topicPartition = entry.getKey(); String topicName = topicPartition.topic(); AddPartitionsToTxnPartitionResult partitionResult = - new AddPartitionsToTxnPartitionResult() - .setErrorCode(entry.getValue().code()) - .setPartitionIndex(topicPartition.partition()); + new AddPartitionsToTxnPartitionResult() + .setErrorCode(entry.getValue().code()) + .setPartitionIndex(topicPartition.partition()); AddPartitionsToTxnPartitionResultCollection partitionResultCollection = resultMap.getOrDefault( - topicName, new AddPartitionsToTxnPartitionResultCollection() + topicName, new AddPartitionsToTxnPartitionResultCollection() ); partitionResultCollection.add(partitionResult); @@ -80,13 +91,14 @@ topicName, new AddPartitionsToTxnPartitionResultCollection() AddPartitionsToTxnTopicResultCollection topicCollection = new AddPartitionsToTxnTopicResultCollection(); for (Map.Entry entry : resultMap.entrySet()) { topicCollection.add(new AddPartitionsToTxnTopicResult() - .setName(entry.getKey()) - .setResults(entry.getValue())); + .setName(entry.getKey()) + .setResults(entry.getValue())); } + return topicCollection; + } - this.data = new AddPartitionsToTxnResponseData() - .setThrottleTimeMs(throttleTimeMs) - .setResults(topicCollection); + public static AddPartitionsToTxnResult resultForTransaction(String transactionalId, Map errors) { + return new AddPartitionsToTxnResult().setTransactionalId(transactionalId).setTopicResults(topicCollectionForErrors(errors)); } @Override @@ -115,9 +127,46 @@ public Map errors() { } return cachedErrorsMap; } + + public Map errorsPerTransaction(String transactionalId) { + if (cachedAllErrorsMap == null) { + cachedAllErrorsMap = new HashMap<>(); + } + + return cachedAllErrorsMap.computeIfAbsent(transactionalId, txnId -> { + Map topicResults = new HashMap<>(); + for (AddPartitionsToTxnTopicResult topicResult : data().resultsByTransaction().find(txnId).topicResults()) { + for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.results()) { + topicResults.put( + new TopicPartition(topicResult.name(), partitionResult.partitionIndex()), Errors.forCode(partitionResult.errorCode())); + } + } + return topicResults; + }); + } + + public Map> allErrors() { + if (cachedAllErrorsMap != null && cachedAllErrorsMap.size() == data.resultsByTransaction().size()) { + return cachedAllErrorsMap; + } + + for (AddPartitionsToTxnResult result : this.data.resultsByTransaction()) { + if (cachedAllErrorsMap == null || !cachedAllErrorsMap.containsKey(result.transactionalId())) { + errorsPerTransaction(result.transactionalId()); + } + } + return cachedAllErrorsMap; + } @Override public Map errorCounts() { + if (data.resultsByTransaction().size() > 0) { + List allErrors = new ArrayList<>(); + allErrors().forEach((txnId, errors) -> + allErrors.addAll(errors.values()) + ); + return errorCounts(allErrors); + } return errorCounts(errors().values()); } diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json index 4920da176c723..09a1b427a97bc 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json @@ -23,17 +23,35 @@ // Version 2 adds the support for new error code PRODUCER_FENCED. // // Version 3 enables flexible versions. - "validVersions": "0-3", + // + // Version 4 adds VerifyOnly field to check if partitions are already in transaction and adds support to batch multiple transactions. + "validVersions": "0-4", "flexibleVersions": "3+", "fields": [ - { "name": "TransactionalId", "type": "string", "versions": "0+", "entityType": "transactionalId", + { "name": "VerifyOnly", "type": "bool", "versions": "4+", "default": false, + "about": "Boolean to signify if we want to check if the partition is in the transaction rather than add it." }, + { "name": "Transactions", "type": "[]AddPartitionsToTxnTransaction", "versions": "4+", + "about": "List of transactions to add partitions to.", "fields": [ + { "name": "TransactionalId", "type": "string", "versions": "4+", "mapKey": true, "entityType": "transactionalId", + "about": "The transactional id corresponding to the transaction."}, + { "name": "ProducerId", "type": "int64", "versions": "4+", "entityType": "producerId", + "about": "Current producer id in use by the transactional id." }, + { "name": "ProducerEpoch", "type": "int16", "versions": "4+", + "about": "Current epoch associated with the producer id." }, + { "name": "Topics", "type": "[]AddPartitionsToTxnTopic", "versions": "4+", + "about": "The partitions to add to the transaction." } + ]}, + { "name": "TransactionalId", "type": "string", "versions": "0-3", "entityType": "transactionalId", "about": "The transactional id corresponding to the transaction."}, - { "name": "ProducerId", "type": "int64", "versions": "0+", "entityType": "producerId", + { "name": "ProducerId", "type": "int64", "versions": "0-3", "entityType": "producerId", "about": "Current producer id in use by the transactional id." }, - { "name": "ProducerEpoch", "type": "int16", "versions": "0+", + { "name": "ProducerEpoch", "type": "int16", "versions": "0-3", "about": "Current epoch associated with the producer id." }, - { "name": "Topics", "type": "[]AddPartitionsToTxnTopic", "versions": "0+", - "about": "The partitions to add to the transaction.", "fields": [ + { "name": "Topics", "type": "[]AddPartitionsToTxnTopic", "versions": "0-3", + "about": "The partitions to add to the transaction." } + ], + "commonStructs": [ + { "name": "AddPartitionsToTxnTopic", "versions": "0+", "fields": [ { "name": "Name", "type": "string", "versions": "0+", "mapKey": true, "entityType": "topicName", "about": "The name of the topic." }, { "name": "Partitions", "type": "[]int32", "versions": "0+", diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json index 4241dc77b4a66..ce323155ae0ef 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json @@ -22,22 +22,35 @@ // Version 2 adds the support for new error code PRODUCER_FENCED. // // Version 3 enables flexible versions. - "validVersions": "0-3", + // + // Version 4 adds support to batch multiple transactions. + "validVersions": "0-4", "flexibleVersions": "3+", "fields": [ { "name": "ThrottleTimeMs", "type": "int32", "versions": "0+", "about": "Duration in milliseconds for which the request was throttled due to a quota violation, or zero if the request did not violate any quota." }, - { "name": "Results", "type": "[]AddPartitionsToTxnTopicResult", "versions": "0+", - "about": "The results for each topic.", "fields": [ + { "name": "ResultsByTransaction", "type": "[]AddPartitionsToTxnResult", "versions": "4+", + "about": "Results categorized by transactional ID.", "fields": [ + { "name": "TransactionalId", "type": "string", "versions": "4+", "mapKey": true, "entityType": "transactionalId", + "about": "The transactional id corresponding to the transaction."}, + { "name": "TopicResults", "type": "[]AddPartitionsToTxnTopicResult", "versions": "4+", + "about": "The results for each topic." } + ]}, + { "name": "Results", "type": "[]AddPartitionsToTxnTopicResult", "versions": "0-3", + "about": "The results for each topic." } + ], + "commonStructs": [ + { "name": "AddPartitionsToTxnTopicResult", "versions": "0+", "fields": [ { "name": "Name", "type": "string", "versions": "0+", "mapKey": true, "entityType": "topicName", "about": "The topic name." }, - { "name": "Results", "type": "[]AddPartitionsToTxnPartitionResult", "versions": "0+", - "about": "The results for each partition", "fields": [ - { "name": "PartitionIndex", "type": "int32", "versions": "0+", "mapKey": true, - "about": "The partition indexes." }, - { "name": "ErrorCode", "type": "int16", "versions": "0+", - "about": "The response error code."} - ]} + { "name": "Results", "type": "[]AddPartitionsToTxnPartitionResult", "versions": "0+", + "about": "The results for each partition" } + ]}, + { "name": "AddPartitionsToTxnPartitionResult", "versions": "0+", "fields": [ + { "name": "PartitionIndex", "type": "int32", "versions": "0+", "mapKey": true, + "about": "The partition indexes." }, + { "name": "ErrorCode", "type": "int16", "versions": "0+", + "about": "The response error code." } ]} ] } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java index 04bde4ae61ba4..7d97c946ba15a 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java @@ -17,6 +17,10 @@ package org.apache.kafka.common.requests; import org.apache.kafka.common.TopicPartition; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopic; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection; import org.apache.kafka.common.utils.annotation.ApiKeyVersionsSource; import org.apache.kafka.common.protocol.ApiKeys; import org.apache.kafka.common.protocol.Errors; @@ -39,21 +43,62 @@ public class AddPartitionsToTxnRequestTest { @ParameterizedTest @ApiKeyVersionsSource(apiKey = ApiKeys.ADD_PARTITIONS_TO_TXN) public void testConstructor(short version) { - List partitions = new ArrayList<>(); - partitions.add(new TopicPartition("topic", 0)); - partitions.add(new TopicPartition("topic", 1)); + TopicPartition tp0 = new TopicPartition("topic", 0); + TopicPartition tp1 = new TopicPartition("topic", 1); - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId, producerId, producerEpoch, partitions); - AddPartitionsToTxnRequest request = builder.build(version); + if (version < 4) { + List partitions = new ArrayList<>(); + partitions.add(tp0); + partitions.add(tp1); - assertEquals(transactionalId, request.data().transactionalId()); - assertEquals(producerId, request.data().producerId()); - assertEquals(producerEpoch, request.data().producerEpoch()); - assertEquals(partitions, request.partitions()); + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId, producerId, producerEpoch, partitions); + AddPartitionsToTxnRequest request = builder.build(version); - AddPartitionsToTxnResponse response = request.getErrorResponse(throttleTimeMs, Errors.UNKNOWN_TOPIC_OR_PARTITION.exception()); + assertEquals(transactionalId, request.data().transactionalId()); + assertEquals(producerId, request.data().producerId()); + assertEquals(producerEpoch, request.data().producerEpoch()); + assertEquals(partitions, request.partitions()); - assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); - assertEquals(throttleTimeMs, response.throttleTimeMs()); + AddPartitionsToTxnResponse response = request.getErrorResponse(throttleTimeMs, Errors.UNKNOWN_TOPIC_OR_PARTITION.exception()); + + assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); + assertEquals(throttleTimeMs, response.throttleTimeMs()); + } else { + String transaction1 = "transaction1"; + String transaction2 = "transaction2"; + + AddPartitionsToTxnTopicCollection topics0 = new AddPartitionsToTxnTopicCollection(); + topics0.add(new AddPartitionsToTxnTopic() + .setName(tp0.topic()) + .setPartitions(Collections.singletonList(tp0.partition()))); + AddPartitionsToTxnTopicCollection topics1 = new AddPartitionsToTxnTopicCollection(); + topics1.add(new AddPartitionsToTxnTopic() + .setName(tp1.topic()) + .setPartitions(Collections.singletonList(tp1.partition()))); + + AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection(); + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transaction1) + .setProducerId(producerId) + .setProducerEpoch(producerEpoch) + .setTopics(topics0)); + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transaction2) + .setProducerId(producerId + 1) + .setProducerEpoch((short) (producerEpoch + 1)) + .setTopics(topics1)); + + boolean verifyOnly = true; + + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions, verifyOnly); + AddPartitionsToTxnRequest request = builder.build(version); + + AddPartitionsToTxnTransaction reqTxn1 = request.data().transactions().find(transaction1); + AddPartitionsToTxnTransaction reqTxn2 = request.data().transactions().find(transaction2); + + assertEquals(verifyOnly, request.data().verifyOnly()); + assertEquals(transactions.find(transaction1), reqTxn1); + assertEquals(transactions.find(transaction2), reqTxn2); + } } } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index 5b67bd47a01f6..08865c33e2d97 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -18,6 +18,8 @@ import org.apache.kafka.common.TopicPartition; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResult; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResult; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResultCollection; @@ -26,6 +28,7 @@ import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import java.util.Collections; import java.util.HashMap; import java.util.Map; @@ -60,6 +63,7 @@ public void setUp() { @Test public void testConstructorWithErrorResponse() { + // This test only applies to versions 0-3. AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(throttleTimeMs, errorsMap); assertEquals(expectedErrorCounts, response.errorCounts()); @@ -84,16 +88,41 @@ public void testParse() { topicCollection.add(topicResult); - AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() - .setResults(topicCollection) - .setThrottleTimeMs(throttleTimeMs); - AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); - for (short version : ApiKeys.ADD_PARTITIONS_TO_TXN.allVersions()) { - AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); - assertEquals(expectedErrorCounts, parsedResponse.errorCounts()); - assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); - assertEquals(version >= 1, parsedResponse.shouldClientThrottle(version)); + + if (version < 4) { + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() + .setResults(topicCollection) + .setThrottleTimeMs(throttleTimeMs); + AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); + + AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); + assertEquals(expectedErrorCounts, parsedResponse.errorCounts()); + assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); + assertEquals(version >= 1, parsedResponse.shouldClientThrottle(version)); + } else { + AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); + results.add(new AddPartitionsToTxnResult().setTransactionalId("txn1").setTopicResults(topicCollection)); + + // Create another transaction with new name and errorOne for a single partition. + Map txnTwoExpectedErrors = Collections.singletonMap(tp2, errorOne); + results.add(AddPartitionsToTxnResponse.resultForTransaction("txn2", txnTwoExpectedErrors)); + + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() + .setResultsByTransaction(results) + .setThrottleTimeMs(throttleTimeMs); + AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); + + Map newExpectedErrorCounts = new HashMap<>(); + newExpectedErrorCounts.put(errorOne, 2); + newExpectedErrorCounts.put(errorTwo, 1); + + AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); + assertEquals(txnTwoExpectedErrors, parsedResponse.errorsPerTransaction("txn2")); + assertEquals(newExpectedErrorCounts, parsedResponse.errorCounts()); + assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); + assertEquals(true, parsedResponse.shouldClientThrottle(version)); + } } } } diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index 1ec906cd22379..7655fd3210b1b 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -16,6 +16,7 @@ */ package kafka.coordinator.transaction +import java.util import java.util.Properties import java.util.concurrent.atomic.AtomicBoolean import kafka.server.{KafkaConfig, MetadataCache, ReplicaManager, RequestLocal} @@ -23,12 +24,14 @@ import kafka.utils.Logging import org.apache.kafka.common.TopicPartition import org.apache.kafka.common.internals.Topic import org.apache.kafka.common.message.{DescribeTransactionsResponseData, ListTransactionsResponseData} +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection import org.apache.kafka.common.metrics.Metrics import org.apache.kafka.common.protocol.Errors import org.apache.kafka.common.record.RecordBatch import org.apache.kafka.common.requests.TransactionResult import org.apache.kafka.common.utils.{LogContext, ProducerIdAndEpoch, Time} import org.apache.kafka.server.util.Scheduler +import scala.jdk.CollectionConverters._ object TransactionCoordinator { @@ -92,6 +95,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, type InitProducerIdCallback = InitProducerIdResult => Unit type AddPartitionsCallback = Errors => Unit + type BatchedAddPartitionsCallback = (String, Errors) => Unit type EndTxnCallback = Errors => Unit type ApiResult[T] = Either[Errors, T] @@ -368,6 +372,66 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } } } + + def handleBatchedAddPartitionsToTransaction(transactions: AddPartitionsToTxnTransactionCollection, + partitionsMap: util.Map[String, util.List[TopicPartition]], + responseCallback: BatchedAddPartitionsCallback, + requestLocal: RequestLocal = RequestLocal.NoCaching): Unit = { + transactions.forEach(transaction => { + val transactionalId = transaction.transactionalId() + val producerId = transaction.producerId() + val producerEpoch = transaction.producerEpoch() + val partitions = partitionsMap.get(transactionalId).asScala.toSet + + if (transactionalId == null || transactionalId.isEmpty) { + debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") + responseCallback(transactionalId, Errors.INVALID_REQUEST) + } else { + // try to update the transaction metadata and append the updated metadata to txn log; + // if there is no such metadata treat it as invalid producerId mapping error. + val result: ApiResult[(Int, TxnTransitMetadata)] = txnManager.getTransactionState(transactionalId).flatMap { + case None => Left(Errors.INVALID_PRODUCER_ID_MAPPING) + + case Some(epochAndMetadata) => + val coordinatorEpoch = epochAndMetadata.coordinatorEpoch + val txnMetadata = epochAndMetadata.transactionMetadata + + // generate the new transaction metadata with added partitions + txnMetadata.inLock { + if (txnMetadata.producerId != producerId) { + Left(Errors.INVALID_PRODUCER_ID_MAPPING) + } else if (txnMetadata.producerEpoch != producerEpoch) { + Left(Errors.PRODUCER_FENCED) + } else if (txnMetadata.pendingTransitionInProgress) { + // return a retriable exception to let the client backoff and retry + Left(Errors.CONCURRENT_TRANSACTIONS) + } else if (txnMetadata.state == PrepareCommit || txnMetadata.state == PrepareAbort) { + Left(Errors.CONCURRENT_TRANSACTIONS) + } else if (txnMetadata.state == Ongoing && partitions.subsetOf(txnMetadata.topicPartitions)) { + // this is an optimization: if the partitions are already in the metadata reply OK immediately + Left(Errors.NONE) + } else { + Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions, time.milliseconds())) + } + } + } + + def perTransactionResponseCallback(error: Errors): Unit = { + responseCallback(transactionalId, error) + } + + result match { + case Left(err) => + debug(s"Returning $err error code to client for $transactionalId's AddPartitions request") + responseCallback(transactionalId, err) + + case Right((coordinatorEpoch, newMetadata)) => + txnManager.appendTransactionToLog(transactionalId, coordinatorEpoch, newMetadata, + perTransactionResponseCallback, requestLocal = requestLocal) + } + } + }) + } /** * Load state from the given partition and begin handling requests for groups which map to this partition. diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 8666b28513ba6..6d237eb12c56f 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -34,6 +34,8 @@ import org.apache.kafka.common.config.ConfigResource import org.apache.kafka.common.errors._ import org.apache.kafka.common.internals.Topic.{GROUP_METADATA_TOPIC_NAME, TRANSACTION_STATE_TOPIC_NAME, isInternal} import org.apache.kafka.common.internals.{FatalExitError, Topic} +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection import org.apache.kafka.common.message.AlterConfigsResponseData.AlterConfigsResourceResponse import org.apache.kafka.common.message.AlterPartitionReassignmentsResponseData.{ReassignablePartitionResponse, ReassignableTopicResponse} import org.apache.kafka.common.message.CreatePartitionsResponseData.CreatePartitionsTopicResult @@ -2386,6 +2388,14 @@ class KafkaApis(val requestChannel: RequestChannel, def handleAddPartitionToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { ensureInterBrokerVersion(IBP_0_11_0_IV0) + if (request.context.apiVersion() < 4) { + handleAddPartitionToTxnRequestV3AndBelow(request, requestLocal) + } else { + handleAddPartitionsToTxnRequestV4(request, requestLocal) + } + } + + def handleAddPartitionToTxnRequestV3AndBelow(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { val addPartitionsToTxnRequest = request.body[AddPartitionsToTxnRequest] val transactionalId = addPartitionsToTxnRequest.data.transactionalId val partitionsToAdd = addPartitionsToTxnRequest.partitions.asScala @@ -2446,6 +2456,67 @@ class KafkaApis(val requestChannel: RequestChannel, } } } + + def handleAddPartitionsToTxnRequestV4(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { + val lock = new Object + val addPartitionsToTxnRequest = request.body[AddPartitionsToTxnRequest] + val responses = new AddPartitionsToTxnResultCollection() + val partitionsByTransaction = addPartitionsToTxnRequest.partitionsByTransaction() + val validTransactions = new AddPartitionsToTxnTransactionCollection() + + addPartitionsToTxnRequest.data().transactions().forEach( transaction => { + val transactionalId = transaction.transactionalId() + val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala + + if (!authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) + responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) + else { + val unauthorizedTopicErrors = mutable.Map[TopicPartition, Errors]() + val nonExistingTopicErrors = mutable.Map[TopicPartition, Errors]() + val authorizedPartitions = mutable.Set[TopicPartition]() + + val authorizedTopics = authHelper.filterByAuthorized(request.context, WRITE, TOPIC, + partitionsToAdd.filterNot(tp => Topic.isInternal(tp.topic)))(_.topic) + for (topicPartition <- partitionsToAdd) { + if (!authorizedTopics.contains(topicPartition.topic)) + unauthorizedTopicErrors += topicPartition -> Errors.TOPIC_AUTHORIZATION_FAILED + else if (!metadataCache.contains(topicPartition)) + nonExistingTopicErrors += topicPartition -> Errors.UNKNOWN_TOPIC_OR_PARTITION + else + authorizedPartitions.add(topicPartition) + } + + if (unauthorizedTopicErrors.nonEmpty || nonExistingTopicErrors.nonEmpty) { + // Any failed partition check causes the entire transaction to fail. We send the appropriate error codes for the + // partitions which failed, and an 'OPERATION_NOT_ATTEMPTED' error code for the partitions which succeeded + // the authorization check to indicate that they were not added to the transaction. + val partitionErrors = unauthorizedTopicErrors ++ nonExistingTopicErrors ++ + authorizedPartitions.map(_ -> Errors.OPERATION_NOT_ATTEMPTED) + responses.add(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitionErrors.asJava)) + } else { + validTransactions.add(transaction) + } + } + }) + if (responses.size() == addPartitionsToTxnRequest.data().transactions().size()) { + requestHelper.sendResponseMaybeThrottle(request, createResponse) + } + + def createResponse(requestThrottleMs: Int): AbstractResponse = { + new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setThrottleTimeMs(requestThrottleMs).setResultsByTransaction(responses)) + } + + def sendResponseCallback(transactionalId: String, error: Errors): Unit = { + + lock synchronized { + responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, error)) + if (responses.size() == addPartitionsToTxnRequest.data().transactions().size()) { + requestHelper.sendResponseMaybeThrottle(request, createResponse) + } + } + } + txnCoordinator.handleBatchedAddPartitionsToTransaction(validTransactions, partitionsByTransaction, sendResponseCallback, requestLocal) + } def handleAddOffsetsToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { ensureInterBrokerVersion(IBP_0_11_0_IV0) diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index 74320e62b49a1..1dd5de41989ea 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -19,15 +19,21 @@ package kafka.server import kafka.utils.TestInfoUtils -import java.util.Properties +import java.util.{Collections, Properties} +import java.util.stream.{Stream => JStream} import org.apache.kafka.common.TopicPartition -import org.apache.kafka.common.protocol.Errors +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopic +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection +import org.apache.kafka.common.protocol.{ApiKeys, Errors} import org.apache.kafka.common.requests.{AddPartitionsToTxnRequest, AddPartitionsToTxnResponse} import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.{BeforeEach, TestInfo} import org.junit.jupiter.params.ParameterizedTest -import org.junit.jupiter.params.provider.ValueSource +import org.junit.jupiter.params.provider.{Arguments, MethodSource} +import scala.collection.mutable import scala.jdk.CollectionConverters._ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { @@ -44,8 +50,8 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { } @ParameterizedTest(name = TestInfoUtils.TestWithParameterizedQuorumName) - @ValueSource(strings = Array("zk", "kraft")) - def shouldReceiveOperationNotAttemptedWhenOtherPartitionHasError(quorum: String): Unit = { + @MethodSource(value = Array("parameters")) + def shouldReceiveOperationNotAttemptedWhenOtherPartitionHasError(quorum: String, version: Short): Unit = { // The basic idea is that we have one unknown topic and one created topic. We should get the 'UNKNOWN_TOPIC_OR_PARTITION' // error for the unknown topic and the 'OPERATION_NOT_ATTEMPTED' error for the known and authorized topic. val nonExistentTopic = new TopicPartition("unknownTopic", 0) @@ -55,22 +61,55 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val producerId = 1000L val producerEpoch: Short = 0 - val request = new AddPartitionsToTxnRequest.Builder( - transactionalId, - producerId, - producerEpoch, - List(createdTopicPartition, nonExistentTopic).asJava) - .build() + val request = + if (version < 4) { + new AddPartitionsToTxnRequest.Builder( + transactionalId, + producerId, + producerEpoch, + List(createdTopicPartition, nonExistentTopic).asJava) + .build() + } else { + val topics = new AddPartitionsToTxnTopicCollection() + topics.add(new AddPartitionsToTxnTopic() + .setName(createdTopicPartition.topic()) + .setPartitions(Collections.singletonList(createdTopicPartition.partition()))) + topics.add(new AddPartitionsToTxnTopic() + .setName(nonExistentTopic.topic()) + .setPartitions(Collections.singletonList(nonExistentTopic.partition()))) + + val transactions = new AddPartitionsToTxnTransactionCollection() + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId) + .setProducerId(producerId) + .setProducerEpoch(producerEpoch) + .setTopics(topics)) + new AddPartitionsToTxnRequest.Builder(transactions, false).build() + } val leaderId = brokers.head.config.brokerId val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(leaderId)) + + val errors = if (version < 4) response.errors else response.errorsPerTransaction(transactionalId) + + assertEquals(2, errors.size) - assertEquals(2, response.errors.size) + assertTrue(errors.containsKey(createdTopicPartition)) + assertEquals(Errors.OPERATION_NOT_ATTEMPTED, errors.get(createdTopicPartition)) - assertTrue(response.errors.containsKey(createdTopicPartition)) - assertEquals(Errors.OPERATION_NOT_ATTEMPTED, response.errors.get(createdTopicPartition)) + assertTrue(errors.containsKey(nonExistentTopic)) + assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION, errors.get(nonExistentTopic)) + } +} - assertTrue(response.errors.containsKey(nonExistentTopic)) - assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION, response.errors.get(nonExistentTopic)) +object AddPartitionsToTxnRequestServerTest { + def parameters: JStream[Arguments] = { + val arguments = mutable.ListBuffer[Arguments]() + ApiKeys.ADD_PARTITIONS_TO_TXN.allVersions().forEach( version => + Array("kraft", "zk").foreach( quorum => + arguments += Arguments.of(quorum, version) + ) + ) + arguments.asJava.stream() } } From 70d11baa372e546cd4e60701a522f100ff94956d Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 3 Feb 2023 17:05:38 -0800 Subject: [PATCH 02/17] cleaner batching --- .../requests/AddPartitionsToTxnRequest.java | 12 +- .../transaction/TransactionCoordinator.scala | 6 +- .../main/scala/kafka/server/KafkaApis.scala | 147 ++++++------------ .../unit/kafka/server/KafkaApisTest.scala | 8 +- 4 files changed, 67 insertions(+), 106 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 8ca6acb061d92..b30be2f8781ab 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -74,8 +74,6 @@ public Builder(AddPartitionsToTxnTransactionCollection transactions, super(ApiKeys.ADD_PARTITIONS_TO_TXN); this.isClientRequest = false; - List transactionsList = new ArrayList<>(); - this.data = new AddPartitionsToTxnRequestData() .setTransactions(transactions) .setVerifyOnly(verifyOnly); @@ -177,6 +175,16 @@ public Map> partitionsByTransaction() { } return cachedPartitionsByTransaction; } + + public AddPartitionsToTxnTransactionCollection singletonTransaction() { + AddPartitionsToTxnTransactionCollection singleTxn = new AddPartitionsToTxnTransactionCollection(); + singleTxn.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(data.transactionalId()) + .setProducerId(data.producerId()) + .setProducerEpoch(data.producerEpoch()) + .setTopics(data.topics())); + return singleTxn; + } @Override public AddPartitionsToTxnRequestData data() { diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index 7655fd3210b1b..7db2b116381de 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -16,7 +16,6 @@ */ package kafka.coordinator.transaction -import java.util import java.util.Properties import java.util.concurrent.atomic.AtomicBoolean import kafka.server.{KafkaConfig, MetadataCache, ReplicaManager, RequestLocal} @@ -24,14 +23,12 @@ import kafka.utils.Logging import org.apache.kafka.common.TopicPartition import org.apache.kafka.common.internals.Topic import org.apache.kafka.common.message.{DescribeTransactionsResponseData, ListTransactionsResponseData} -import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection import org.apache.kafka.common.metrics.Metrics import org.apache.kafka.common.protocol.Errors import org.apache.kafka.common.record.RecordBatch import org.apache.kafka.common.requests.TransactionResult import org.apache.kafka.common.utils.{LogContext, ProducerIdAndEpoch, Time} import org.apache.kafka.server.util.Scheduler -import scala.jdk.CollectionConverters._ object TransactionCoordinator { @@ -372,7 +369,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } } } - + /* def handleBatchedAddPartitionsToTransaction(transactions: AddPartitionsToTxnTransactionCollection, partitionsMap: util.Map[String, util.List[TopicPartition]], responseCallback: BatchedAddPartitionsCallback, @@ -432,6 +429,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } }) } + */ /** * Load state from the given partition and begin handling requests for groups which map to this partition. diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 6d237eb12c56f..af6dad784d368 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -34,7 +34,6 @@ import org.apache.kafka.common.config.ConfigResource import org.apache.kafka.common.errors._ import org.apache.kafka.common.internals.Topic.{GROUP_METADATA_TOPIC_NAME, TRANSACTION_STATE_TOPIC_NAME, isInternal} import org.apache.kafka.common.internals.{FatalExitError, Topic} -import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection import org.apache.kafka.common.message.AlterConfigsResponseData.AlterConfigsResourceResponse import org.apache.kafka.common.message.AlterPartitionReassignmentsResponseData.{ReassignablePartitionResponse, ReassignableTopicResponse} @@ -201,7 +200,7 @@ class KafkaApis(val requestChannel: RequestChannel, case ApiKeys.DELETE_RECORDS => handleDeleteRecordsRequest(request) case ApiKeys.INIT_PRODUCER_ID => handleInitProducerIdRequest(request, requestLocal) case ApiKeys.OFFSET_FOR_LEADER_EPOCH => handleOffsetForLeaderEpochRequest(request) - case ApiKeys.ADD_PARTITIONS_TO_TXN => handleAddPartitionToTxnRequest(request, requestLocal) + case ApiKeys.ADD_PARTITIONS_TO_TXN => handleAddPartitionsToTxnRequest(request, requestLocal) case ApiKeys.ADD_OFFSETS_TO_TXN => handleAddOffsetsToTxnRequest(request, requestLocal) case ApiKeys.END_TXN => handleEndTxnRequest(request, requestLocal) case ApiKeys.WRITE_TXN_MARKERS => handleWriteTxnMarkersRequest(request, requestLocal) @@ -2385,88 +2384,34 @@ class KafkaApis(val requestChannel: RequestChannel, if (config.interBrokerProtocolVersion.isLessThan(version)) throw new UnsupportedVersionException(s"inter.broker.protocol.version: ${config.interBrokerProtocolVersion.version} is less than the required version: ${version.version}") } - - def handleAddPartitionToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { - ensureInterBrokerVersion(IBP_0_11_0_IV0) - if (request.context.apiVersion() < 4) { - handleAddPartitionToTxnRequestV3AndBelow(request, requestLocal) - } else { - handleAddPartitionsToTxnRequestV4(request, requestLocal) - } - } - - def handleAddPartitionToTxnRequestV3AndBelow(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { + def handleAddPartitionsToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { + val lock = new Object val addPartitionsToTxnRequest = request.body[AddPartitionsToTxnRequest] - val transactionalId = addPartitionsToTxnRequest.data.transactionalId - val partitionsToAdd = addPartitionsToTxnRequest.partitions.asScala - if (!authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) - requestHelper.sendResponseMaybeThrottle(request, requestThrottleMs => - addPartitionsToTxnRequest.getErrorResponse(requestThrottleMs, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED.exception)) - else { - val unauthorizedTopicErrors = mutable.Map[TopicPartition, Errors]() - val nonExistingTopicErrors = mutable.Map[TopicPartition, Errors]() - val authorizedPartitions = mutable.Set[TopicPartition]() - - val authorizedTopics = authHelper.filterByAuthorized(request.context, WRITE, TOPIC, - partitionsToAdd.filterNot(tp => Topic.isInternal(tp.topic)))(_.topic) - for (topicPartition <- partitionsToAdd) { - if (!authorizedTopics.contains(topicPartition.topic)) - unauthorizedTopicErrors += topicPartition -> Errors.TOPIC_AUTHORIZATION_FAILED - else if (!metadataCache.contains(topicPartition)) - nonExistingTopicErrors += topicPartition -> Errors.UNKNOWN_TOPIC_OR_PARTITION - else - authorizedPartitions.add(topicPartition) - } + val version = addPartitionsToTxnRequest.version + val responses = new AddPartitionsToTxnResultCollection() + val partitionsByTransaction = addPartitionsToTxnRequest.partitionsByTransaction() - if (unauthorizedTopicErrors.nonEmpty || nonExistingTopicErrors.nonEmpty) { - // Any failed partition check causes the entire request to fail. We send the appropriate error codes for the - // partitions which failed, and an 'OPERATION_NOT_ATTEMPTED' error code for the partitions which succeeded - // the authorization check to indicate that they were not added to the transaction. - val partitionErrors = unauthorizedTopicErrors ++ nonExistingTopicErrors ++ - authorizedPartitions.map(_ -> Errors.OPERATION_NOT_ATTEMPTED) - requestHelper.sendResponseMaybeThrottle(request, requestThrottleMs => - new AddPartitionsToTxnResponse(requestThrottleMs, partitionErrors.asJava)) + // V4 requests introduced batches of transactions. We need all transactions to be handled before sending the + // response so there are a few differences in handling errors and sending responses. + def createResponse(requestThrottleMs: Int): AbstractResponse = { + if (version < 4) { + val data = new AddPartitionsToTxnResponseData() + responses.forEach(result => { + data.setResults(result.topicResults()) + data.setThrottleTimeMs(requestThrottleMs) + }) + new AddPartitionsToTxnResponse(data) } else { - def sendResponseCallback(error: Errors): Unit = { - def createResponse(requestThrottleMs: Int): AbstractResponse = { - val finalError = - if (addPartitionsToTxnRequest.version < 2 && error == Errors.PRODUCER_FENCED) { - // For older clients, they could not understand the new PRODUCER_FENCED error code, - // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. - Errors.INVALID_PRODUCER_EPOCH - } else { - error - } - - val responseBody: AddPartitionsToTxnResponse = new AddPartitionsToTxnResponse(requestThrottleMs, - partitionsToAdd.map{tp => (tp, finalError)}.toMap.asJava) - trace(s"Completed $transactionalId's AddPartitionsToTxnRequest with partitions $partitionsToAdd: errors: $error from client ${request.header.clientId}") - responseBody - } - - requestHelper.sendResponseMaybeThrottle(request, createResponse) - } - - txnCoordinator.handleAddPartitionsToTransaction(transactionalId, - addPartitionsToTxnRequest.data.producerId, - addPartitionsToTxnRequest.data.producerEpoch, - authorizedPartitions, - sendResponseCallback, - requestLocal) + new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setThrottleTimeMs(requestThrottleMs).setResultsByTransaction(responses)) } } - } - - def handleAddPartitionsToTxnRequestV4(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { - val lock = new Object - val addPartitionsToTxnRequest = request.body[AddPartitionsToTxnRequest] - val responses = new AddPartitionsToTxnResultCollection() - val partitionsByTransaction = addPartitionsToTxnRequest.partitionsByTransaction() - val validTransactions = new AddPartitionsToTxnTransactionCollection() + + val txns = if (version < 4) addPartitionsToTxnRequest.singletonTransaction() else addPartitionsToTxnRequest.data.transactions + def allResponsesPresent: Boolean = responses.size() == txns.size() - addPartitionsToTxnRequest.data().transactions().forEach( transaction => { + txns.forEach( transaction => { val transactionalId = transaction.transactionalId() - val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala + val partitionsToAdd = if (version < 4) addPartitionsToTxnRequest.partitions.asScala else partitionsByTransaction.get(transactionalId).asScala if (!authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) @@ -2494,28 +2439,38 @@ class KafkaApis(val requestChannel: RequestChannel, authorizedPartitions.map(_ -> Errors.OPERATION_NOT_ATTEMPTED) responses.add(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitionErrors.asJava)) } else { - validTransactions.add(transaction) - } - } - }) - if (responses.size() == addPartitionsToTxnRequest.data().transactions().size()) { - requestHelper.sendResponseMaybeThrottle(request, createResponse) - } - - def createResponse(requestThrottleMs: Int): AbstractResponse = { - new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setThrottleTimeMs(requestThrottleMs).setResultsByTransaction(responses)) - } - - def sendResponseCallback(transactionalId: String, error: Errors): Unit = { + def sendResponseCallback(error: Errors): Unit = { + val finalError = { + if (version < 2 && error == Errors.PRODUCER_FENCED) { + // For older clients, they could not understand the new PRODUCER_FENCED error code, + // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. + Errors.INVALID_PRODUCER_EPOCH + } else { + error + } + } + lock synchronized { + responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) + } + + if (allResponsesPresent) { + requestHelper.sendResponseMaybeThrottle(request, createResponse) + } + } - lock synchronized { - responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, error)) - if (responses.size() == addPartitionsToTxnRequest.data().transactions().size()) { - requestHelper.sendResponseMaybeThrottle(request, createResponse) + txnCoordinator.handleAddPartitionsToTransaction(transactionalId, + transaction.producerId, + transaction.producerEpoch, + authorizedPartitions, + sendResponseCallback, + requestLocal) } + + // If all transactions are present send response. + if (allResponsesPresent) + requestHelper.sendResponseMaybeThrottle(request, createResponse) } - } - txnCoordinator.handleBatchedAddPartitionsToTransaction(validTransactions, partitionsByTransaction, sendResponseCallback, requestLocal) + }) } def handleAddOffsetsToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index 58c9956792616..b0760b1a9b78f 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -2020,7 +2020,7 @@ class KafkaApisTest { ArgumentMatchers.eq(requestLocal) )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) - createKafkaApis().handleAddPartitionToTxnRequest(request, requestLocal) + createKafkaApis().handleAddPartitionsToTxnRequest(request, requestLocal) verify(requestChannel).sendResponse( ArgumentMatchers.eq(request), @@ -2158,7 +2158,7 @@ class KafkaApisTest { when(clientRequestQuotaManager.maybeRecordAndGetThrottleTimeMs(any[RequestChannel.Request](), any[Long])).thenReturn(0) - createKafkaApis().handleAddPartitionToTxnRequest(request, RequestLocal.withThreadConfinedCaching) + createKafkaApis().handleAddPartitionsToTxnRequest(request, RequestLocal.withThreadConfinedCaching) val response = verifyNoThrottling[AddPartitionsToTxnResponse](request) assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION, response.errors().get(invalidTopicPartition)) @@ -2177,13 +2177,13 @@ class KafkaApisTest { @Test def shouldThrowUnsupportedVersionExceptionOnHandleAddPartitionsToTxnRequestWhenInterBrokerProtocolNotSupported(): Unit = { assertThrows(classOf[UnsupportedVersionException], - () => createKafkaApis(IBP_0_10_2_IV0).handleAddPartitionToTxnRequest(null, RequestLocal.withThreadConfinedCaching)) + () => createKafkaApis(IBP_0_10_2_IV0).handleAddPartitionsToTxnRequest(null, RequestLocal.withThreadConfinedCaching)) } @Test def shouldThrowUnsupportedVersionExceptionOnHandleTxnOffsetCommitRequestWhenInterBrokerProtocolNotSupported(): Unit = { assertThrows(classOf[UnsupportedVersionException], - () => createKafkaApis(IBP_0_10_2_IV0).handleAddPartitionToTxnRequest(null, RequestLocal.withThreadConfinedCaching)) + () => createKafkaApis(IBP_0_10_2_IV0).handleAddPartitionsToTxnRequest(null, RequestLocal.withThreadConfinedCaching)) } @Test From a8d6c006eb9c614a6aefa217c441c8919046a084 Mon Sep 17 00:00:00 2001 From: Justine Date: Mon, 6 Feb 2023 11:43:48 -0800 Subject: [PATCH 03/17] bug fixes --- .../requests/AddPartitionsToTxnRequest.java | 4 +++ .../main/scala/kafka/server/KafkaApis.scala | 33 ++++++++++--------- 2 files changed, 22 insertions(+), 15 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index b30be2f8781ab..610c6c7e5c454 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -176,6 +176,10 @@ public Map> partitionsByTransaction() { return cachedPartitionsByTransaction; } + public AddPartitionsToTxnRequest normalizeRequest() { + return new AddPartitionsToTxnRequest(new AddPartitionsToTxnRequestData().setTransactions(singletonTransaction()), version); + } + public AddPartitionsToTxnTransactionCollection singletonTransaction() { AddPartitionsToTxnTransactionCollection singleTxn = new AddPartitionsToTxnTransactionCollection(); singleTxn.add(new AddPartitionsToTxnTransaction() diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index af6dad784d368..d4c9acb2bc390 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2385,8 +2385,9 @@ class KafkaApis(val requestChannel: RequestChannel, throw new UnsupportedVersionException(s"inter.broker.protocol.version: ${config.interBrokerProtocolVersion.version} is less than the required version: ${version.version}") } def handleAddPartitionsToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { + ensureInterBrokerVersion(IBP_0_11_0_IV0) val lock = new Object - val addPartitionsToTxnRequest = request.body[AddPartitionsToTxnRequest] + val addPartitionsToTxnRequest = if (request.context.apiVersion() < 4) request.body[AddPartitionsToTxnRequest].normalizeRequest() else request.body[AddPartitionsToTxnRequest] val version = addPartitionsToTxnRequest.version val responses = new AddPartitionsToTxnResultCollection() val partitionsByTransaction = addPartitionsToTxnRequest.partitionsByTransaction() @@ -2395,6 +2396,7 @@ class KafkaApis(val requestChannel: RequestChannel, // response so there are a few differences in handling errors and sending responses. def createResponse(requestThrottleMs: Int): AbstractResponse = { if (version < 4) { + // There will only be one response in data. Add it to the response data object. val data = new AddPartitionsToTxnResponseData() responses.forEach(result => { data.setResults(result.topicResults()) @@ -2406,16 +2408,23 @@ class KafkaApis(val requestChannel: RequestChannel, } } - val txns = if (version < 4) addPartitionsToTxnRequest.singletonTransaction() else addPartitionsToTxnRequest.data.transactions - def allResponsesPresent: Boolean = responses.size() == txns.size() - + val txns = addPartitionsToTxnRequest.data.transactions + def maybeSendResponse(): Unit = { + lock synchronized { + if (responses.size() == txns.size()) { + requestHelper.sendResponseMaybeThrottle(request, createResponse) + } + } + } + txns.forEach( transaction => { val transactionalId = transaction.transactionalId() - val partitionsToAdd = if (version < 4) addPartitionsToTxnRequest.partitions.asScala else partitionsByTransaction.get(transactionalId).asScala + val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala - if (!authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) + if (!authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) { responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) - else { + maybeSendResponse() + } else { val unauthorizedTopicErrors = mutable.Map[TopicPartition, Errors]() val nonExistingTopicErrors = mutable.Map[TopicPartition, Errors]() val authorizedPartitions = mutable.Set[TopicPartition]() @@ -2438,6 +2447,7 @@ class KafkaApis(val requestChannel: RequestChannel, val partitionErrors = unauthorizedTopicErrors ++ nonExistingTopicErrors ++ authorizedPartitions.map(_ -> Errors.OPERATION_NOT_ATTEMPTED) responses.add(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitionErrors.asJava)) + maybeSendResponse() } else { def sendResponseCallback(error: Errors): Unit = { val finalError = { @@ -2452,10 +2462,7 @@ class KafkaApis(val requestChannel: RequestChannel, lock synchronized { responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) } - - if (allResponsesPresent) { - requestHelper.sendResponseMaybeThrottle(request, createResponse) - } + maybeSendResponse() } txnCoordinator.handleAddPartitionsToTransaction(transactionalId, @@ -2465,10 +2472,6 @@ class KafkaApis(val requestChannel: RequestChannel, sendResponseCallback, requestLocal) } - - // If all transactions are present send response. - if (allResponsesPresent) - requestHelper.sendResponseMaybeThrottle(request, createResponse) } }) } From 47f0fff9dd3f5545aa90e71fda9bba65ae3ea20b Mon Sep 17 00:00:00 2001 From: Justine Date: Mon, 6 Feb 2023 16:39:20 -0800 Subject: [PATCH 04/17] add verifyOnly and auth changes --- .../transaction/TransactionCoordinator.scala | 70 ++----------------- .../main/scala/kafka/server/KafkaApis.scala | 7 +- ...ransactionCoordinatorConcurrencyTest.scala | 3 +- .../TransactionCoordinatorTest.scala | 39 ++++++++--- .../unit/kafka/server/KafkaApisTest.scala | 2 + 5 files changed, 45 insertions(+), 76 deletions(-) diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index 7db2b116381de..31b64d7b2fdaf 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -92,7 +92,6 @@ class TransactionCoordinator(txnConfig: TransactionConfig, type InitProducerIdCallback = InitProducerIdResult => Unit type AddPartitionsCallback = Errors => Unit - type BatchedAddPartitionsCallback = (String, Errors) => Unit type EndTxnCallback = Errors => Unit type ApiResult[T] = Either[Errors, T] @@ -323,6 +322,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, producerId: Long, producerEpoch: Short, partitions: collection.Set[TopicPartition], + verifyOnly: Boolean, responseCallback: AddPartitionsCallback, requestLocal: RequestLocal = RequestLocal.NoCaching): Unit = { if (transactionalId == null || transactionalId.isEmpty) { @@ -353,7 +353,12 @@ class TransactionCoordinator(txnConfig: TransactionConfig, // this is an optimization: if the partitions are already in the metadata reply OK immediately Left(Errors.NONE) } else { - Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds())) + // If verifyOnly, we should have returned in the step above. If we didn't the partitions are not present in the transaction. + if (verifyOnly) { + Left(Errors.INVALID_TXN_STATE) + } else { + Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds())) + } } } } @@ -369,67 +374,6 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } } } - /* - def handleBatchedAddPartitionsToTransaction(transactions: AddPartitionsToTxnTransactionCollection, - partitionsMap: util.Map[String, util.List[TopicPartition]], - responseCallback: BatchedAddPartitionsCallback, - requestLocal: RequestLocal = RequestLocal.NoCaching): Unit = { - transactions.forEach(transaction => { - val transactionalId = transaction.transactionalId() - val producerId = transaction.producerId() - val producerEpoch = transaction.producerEpoch() - val partitions = partitionsMap.get(transactionalId).asScala.toSet - - if (transactionalId == null || transactionalId.isEmpty) { - debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") - responseCallback(transactionalId, Errors.INVALID_REQUEST) - } else { - // try to update the transaction metadata and append the updated metadata to txn log; - // if there is no such metadata treat it as invalid producerId mapping error. - val result: ApiResult[(Int, TxnTransitMetadata)] = txnManager.getTransactionState(transactionalId).flatMap { - case None => Left(Errors.INVALID_PRODUCER_ID_MAPPING) - - case Some(epochAndMetadata) => - val coordinatorEpoch = epochAndMetadata.coordinatorEpoch - val txnMetadata = epochAndMetadata.transactionMetadata - - // generate the new transaction metadata with added partitions - txnMetadata.inLock { - if (txnMetadata.producerId != producerId) { - Left(Errors.INVALID_PRODUCER_ID_MAPPING) - } else if (txnMetadata.producerEpoch != producerEpoch) { - Left(Errors.PRODUCER_FENCED) - } else if (txnMetadata.pendingTransitionInProgress) { - // return a retriable exception to let the client backoff and retry - Left(Errors.CONCURRENT_TRANSACTIONS) - } else if (txnMetadata.state == PrepareCommit || txnMetadata.state == PrepareAbort) { - Left(Errors.CONCURRENT_TRANSACTIONS) - } else if (txnMetadata.state == Ongoing && partitions.subsetOf(txnMetadata.topicPartitions)) { - // this is an optimization: if the partitions are already in the metadata reply OK immediately - Left(Errors.NONE) - } else { - Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions, time.milliseconds())) - } - } - } - - def perTransactionResponseCallback(error: Errors): Unit = { - responseCallback(transactionalId, error) - } - - result match { - case Left(err) => - debug(s"Returning $err error code to client for $transactionalId's AddPartitions request") - responseCallback(transactionalId, err) - - case Right((coordinatorEpoch, newMetadata)) => - txnManager.appendTransactionToLog(transactionalId, coordinatorEpoch, newMetadata, - perTransactionResponseCallback, requestLocal = requestLocal) - } - } - }) - } - */ /** * Load state from the given partition and begin handling requests for groups which map to this partition. diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index d4c9acb2bc390..73a91398596e2 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2391,6 +2391,9 @@ class KafkaApis(val requestChannel: RequestChannel, val version = addPartitionsToTxnRequest.version val responses = new AddPartitionsToTxnResultCollection() val partitionsByTransaction = addPartitionsToTxnRequest.partitionsByTransaction() + + // Newer versions of the request should only come from other brokers. + if (version >= 4) authHelper.authorizeClusterOperation(request, CLUSTER_ACTION) // V4 requests introduced batches of transactions. We need all transactions to be handled before sending the // response so there are a few differences in handling errors and sending responses. @@ -2421,7 +2424,7 @@ class KafkaApis(val requestChannel: RequestChannel, val transactionalId = transaction.transactionalId() val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala - if (!authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) { + if (version < 4 && !authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) { responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) maybeSendResponse() } else { @@ -2469,6 +2472,7 @@ class KafkaApis(val requestChannel: RequestChannel, transaction.producerId, transaction.producerEpoch, authorizedPartitions, + addPartitionsToTxnRequest.data.verifyOnly, sendResponseCallback, requestLocal) } @@ -2521,6 +2525,7 @@ class KafkaApis(val requestChannel: RequestChannel, addOffsetsToTxnRequest.data.producerId, addOffsetsToTxnRequest.data.producerEpoch, Set(offsetTopicPartition), + false, sendResponseCallback, requestLocal) } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala index 02852d6b94337..e1c910d0b5cbe 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala @@ -520,7 +520,8 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren transactionCoordinator.handleAddPartitionsToTransaction(txn.transactionalId, txnMetadata.producerId, txnMetadata.producerEpoch, - partitions, + partitions, + false, resultCallback, RequestLocal.withThreadConfinedCaching) replicaManager.tryCompleteActions() diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index 1c8e0fcdc1b90..cc873c99a5520 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -24,6 +24,8 @@ import org.apache.kafka.common.utils.{LogContext, MockTime, ProducerIdAndEpoch} import org.apache.kafka.server.util.MockScheduler import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.Test +import org.junit.jupiter.params.ParameterizedTest +import org.junit.jupiter.params.provider.ValueSource import org.mockito.{ArgumentCaptor, ArgumentMatchers} import org.mockito.ArgumentMatchers.{any, anyInt} import org.mockito.Mockito.{mock, times, verify, when} @@ -200,19 +202,19 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(None)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback) assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsEmpty(): Unit = { - coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, false, errorsCallback) assertEquals(Errors.INVALID_REQUEST, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsNull(): Unit = { - coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, false, errorsCallback) assertEquals(Errors.INVALID_REQUEST, error) } @@ -221,7 +223,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.NOT_COORDINATOR)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback) assertEquals(Errors.NOT_COORDINATOR, error) } @@ -230,7 +232,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.COORDINATOR_LOAD_IN_PROGRESS)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback) assertEquals(Errors.COORDINATOR_LOAD_IN_PROGRESS, error) } @@ -249,7 +251,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, state, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback) assertEquals(Errors.CONCURRENT_TRANSACTIONS, error) } @@ -259,7 +261,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 10, 9, 0, PrepareCommit, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback) assertEquals(Errors.PRODUCER_FENCED, error) } @@ -290,7 +292,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, txnMetadata)))) - coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, false, errorsCallback) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) verify(transactionManager).appendTransactionToLog( @@ -302,15 +304,30 @@ class TransactionCoordinatorTest { any() ) } + + @ParameterizedTest + @ValueSource(booleans = Array(true, false)) + def shouldRespondWithErrorsNoneOnAddPartitionWhenNoErrorsAndPartitionsTheSame(verifyOnly: Boolean): Unit = { + when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) + .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, + new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, verifyOnly, errorsCallback) + assertEquals(Errors.NONE, error) + verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) + } + @Test - def shouldRespondWithErrorsNoneOnAddPartitionWhenNoErrorsAndPartitionsTheSame(): Unit = { + def shouldRespondWithInvalidTxnStateWhenVerifyOnlyAndPartitionNotPresent(): Unit = { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, errorsCallback) - assertEquals(Errors.NONE, error) + + val extraPartitions = partitions ++ Set(new TopicPartition("topic2", 0)) + + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, extraPartitions, true, errorsCallback) + assertEquals(Errors.INVALID_TXN_STATE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index b0760b1a9b78f..1ef7cc4446dc0 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -1962,6 +1962,7 @@ class KafkaApisTest { ArgumentMatchers.eq(producerId), ArgumentMatchers.eq(epoch), ArgumentMatchers.eq(Set(new TopicPartition(Topic.GROUP_METADATA_TOPIC_NAME, partition))), + ArgumentMatchers.eq(false), responseCallback.capture(), ArgumentMatchers.eq(requestLocal) )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) @@ -2016,6 +2017,7 @@ class KafkaApisTest { ArgumentMatchers.eq(producerId), ArgumentMatchers.eq(epoch), ArgumentMatchers.eq(Set(topicPartition)), + ArgumentMatchers.eq(false), responseCallback.capture(), ArgumentMatchers.eq(requestLocal) )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) From 4c4d99edb12e2b8798ecacfa68fa11a0c904a146 Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 10 Feb 2023 09:44:43 -0800 Subject: [PATCH 05/17] more tests --- .../requests/AddPartitionsToTxnRequest.java | 18 +-- .../requests/AddPartitionsToTxnResponse.java | 2 + .../AddPartitionsToTxnRequestTest.java | 134 ++++++++++++------ .../AddPartitionsToTxnResponseTest.java | 24 +++- .../AddPartitionsToTxnRequestServerTest.scala | 68 ++++++++- 5 files changed, 186 insertions(+), 60 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 610c6c7e5c454..624b9ecb8b857 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -53,6 +53,7 @@ public static class Builder extends AbstractRequest.Builder getPartitions(AddPartitionsToTxnRequestData data) { List partitions = new ArrayList<>(); - for (AddPartitionsToTxnTransaction transaction : data.transactions()) { - for (AddPartitionsToTxnTopic topicCollection : transaction.topics()) { - for (Integer partition : topicCollection.partitions()) { - partitions.add(new TopicPartition(topicCollection.name(), partition)); - } - } - } - // Add singleton topics for (AddPartitionsToTxnTopic topicCollection : data.topics()) { for (Integer partition : topicCollection.partitions()) { partitions.add(new TopicPartition(topicCollection.name(), partition)); @@ -138,7 +132,8 @@ public AddPartitionsToTxnRequest(final AddPartitionsToTxnRequestData data, short this.data = data; this.version = version; } - + + // Only used for versions < 4 public List partitions() { if (cachedPartitions != null) { return cachedPartitions; @@ -147,7 +142,7 @@ public List partitions() { return cachedPartitions; } - public List partitionsForTransaction(String transaction) { + private List partitionsForTransaction(String transaction) { if (cachedPartitionsByTransaction == null) { cachedPartitionsByTransaction = new HashMap<>(); } @@ -176,11 +171,12 @@ public Map> partitionsByTransaction() { return cachedPartitionsByTransaction; } + // Takes a version 3 or below request and returns a v4+ singleton (one transaction ID) request. public AddPartitionsToTxnRequest normalizeRequest() { return new AddPartitionsToTxnRequest(new AddPartitionsToTxnRequestData().setTransactions(singletonTransaction()), version); } - public AddPartitionsToTxnTransactionCollection singletonTransaction() { + private AddPartitionsToTxnTransactionCollection singletonTransaction() { AddPartitionsToTxnTransactionCollection singleTxn = new AddPartitionsToTxnTransactionCollection(); singleTxn.add(new AddPartitionsToTxnTransaction() .setTransactionalId(data.transactionalId()) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index 9831eabef6989..d2e42f744ddbd 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -60,6 +60,7 @@ public AddPartitionsToTxnResponse(AddPartitionsToTxnResponseData data) { this.data = data; } + // Only used for versions < 4 public AddPartitionsToTxnResponse(int throttleTimeMs, Map errors) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); @@ -111,6 +112,7 @@ public void maybeSetThrottleTimeMs(int throttleTimeMs) { data.setThrottleTimeMs(throttleTimeMs); } + // Only used for versions < 4 public Map errors() { if (cachedErrorsMap != null) { return cachedErrorsMap; diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java index 7d97c946ba15a..bf1046abe85ee 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java @@ -21,84 +21,134 @@ import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; import org.apache.kafka.common.utils.annotation.ApiKeyVersionsSource; import org.apache.kafka.common.protocol.ApiKeys; import org.apache.kafka.common.protocol.Errors; import java.util.ArrayList; - import java.util.Collections; +import java.util.HashMap; import java.util.List; +import java.util.Map; + +import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; import static org.junit.jupiter.api.Assertions.assertEquals; public class AddPartitionsToTxnRequestTest { - - private static String transactionalId = "transactionalId"; + private final String transactionalId1 = "transaction1"; + private final String transactionalId2 = "transaction2"; private static int producerId = 10; private static short producerEpoch = 1; private static int throttleTimeMs = 10; + private static TopicPartition tp0 = new TopicPartition("topic", 0); + private static TopicPartition tp1 = new TopicPartition("topic", 1); @ParameterizedTest @ApiKeyVersionsSource(apiKey = ApiKeys.ADD_PARTITIONS_TO_TXN) public void testConstructor(short version) { - TopicPartition tp0 = new TopicPartition("topic", 0); - TopicPartition tp1 = new TopicPartition("topic", 1); + + AddPartitionsToTxnRequest request; if (version < 4) { List partitions = new ArrayList<>(); partitions.add(tp0); partitions.add(tp1); - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId, producerId, producerEpoch, partitions); - AddPartitionsToTxnRequest request = builder.build(version); + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId1, producerId, producerEpoch, partitions); + request = builder.build(version); - assertEquals(transactionalId, request.data().transactionalId()); + assertEquals(transactionalId1, request.data().transactionalId()); assertEquals(producerId, request.data().producerId()); assertEquals(producerEpoch, request.data().producerEpoch()); assertEquals(partitions, request.partitions()); - - AddPartitionsToTxnResponse response = request.getErrorResponse(throttleTimeMs, Errors.UNKNOWN_TOPIC_OR_PARTITION.exception()); - - assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); - assertEquals(throttleTimeMs, response.throttleTimeMs()); } else { - String transaction1 = "transaction1"; - String transaction2 = "transaction2"; - - AddPartitionsToTxnTopicCollection topics0 = new AddPartitionsToTxnTopicCollection(); - topics0.add(new AddPartitionsToTxnTopic() - .setName(tp0.topic()) - .setPartitions(Collections.singletonList(tp0.partition()))); - AddPartitionsToTxnTopicCollection topics1 = new AddPartitionsToTxnTopicCollection(); - topics1.add(new AddPartitionsToTxnTopic() - .setName(tp1.topic()) - .setPartitions(Collections.singletonList(tp1.partition()))); - - AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection(); - transactions.add(new AddPartitionsToTxnTransaction() - .setTransactionalId(transaction1) - .setProducerId(producerId) - .setProducerEpoch(producerEpoch) - .setTopics(topics0)); - transactions.add(new AddPartitionsToTxnTransaction() - .setTransactionalId(transaction2) - .setProducerId(producerId + 1) - .setProducerEpoch((short) (producerEpoch + 1)) - .setTopics(topics1)); - + AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); boolean verifyOnly = true; AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions, verifyOnly); - AddPartitionsToTxnRequest request = builder.build(version); + request = builder.build(version); - AddPartitionsToTxnTransaction reqTxn1 = request.data().transactions().find(transaction1); - AddPartitionsToTxnTransaction reqTxn2 = request.data().transactions().find(transaction2); + AddPartitionsToTxnTransaction reqTxn1 = request.data().transactions().find(transactionalId1); + AddPartitionsToTxnTransaction reqTxn2 = request.data().transactions().find(transactionalId2); assertEquals(verifyOnly, request.data().verifyOnly()); - assertEquals(transactions.find(transaction1), reqTxn1); - assertEquals(transactions.find(transaction2), reqTxn2); + assertEquals(transactions.find(transactionalId1), reqTxn1); + assertEquals(transactions.find(transactionalId2), reqTxn2); } + AddPartitionsToTxnResponse response = request.getErrorResponse(throttleTimeMs, Errors.UNKNOWN_TOPIC_OR_PARTITION.exception()); + + assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); + assertEquals(throttleTimeMs, response.throttleTimeMs()); + } + + @Test + public void testBatchedRequests() { + AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); + boolean verifyOnly = true; + + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions, verifyOnly); + AddPartitionsToTxnRequest request = builder.build(ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion()); + + Map> expectedMap = new HashMap<>(); + expectedMap.put(transactionalId1, Collections.singletonList(tp0)); + expectedMap.put(transactionalId2, Collections.singletonList(tp1)); + + assertEquals(expectedMap, request.partitionsByTransaction()); + + AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection(); + + results.add(request.errorResponseForTransaction(transactionalId1, Errors.UNKNOWN_TOPIC_OR_PARTITION)); + results.add(request.errorResponseForTransaction(transactionalId2, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)); + + AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData() + .setResultsByTransaction(results) + .setThrottleTimeMs(throttleTimeMs)); + + assertEquals(Collections.singletonMap(tp0, Errors.UNKNOWN_TOPIC_OR_PARTITION), response.errorsPerTransaction(transactionalId1)); + assertEquals(Collections.singletonMap(tp1, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED), response.errorsPerTransaction(transactionalId2)); + } + + @Test + public void testNormalizeRequest() { + List partitions = new ArrayList<>(); + partitions.add(tp0); + partitions.add(tp1); + + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId1, producerId, producerEpoch, partitions); + AddPartitionsToTxnRequest request = builder.build((short) 3); + + AddPartitionsToTxnRequest singleton = request.normalizeRequest(); + assertEquals(partitions, singleton.partitionsByTransaction().get(transactionalId1)); + + AddPartitionsToTxnTransaction transaction = singleton.data().transactions().find(transactionalId1); + assertEquals(producerId, transaction.producerId()); + assertEquals(producerEpoch, transaction.producerEpoch()); + } + + private AddPartitionsToTxnTransactionCollection createTwoTransactionCollection() { + AddPartitionsToTxnTopicCollection topics0 = new AddPartitionsToTxnTopicCollection(); + topics0.add(new AddPartitionsToTxnTopic() + .setName(tp0.topic()) + .setPartitions(Collections.singletonList(tp0.partition()))); + AddPartitionsToTxnTopicCollection topics1 = new AddPartitionsToTxnTopicCollection(); + topics1.add(new AddPartitionsToTxnTopic() + .setName(tp1.topic()) + .setPartitions(Collections.singletonList(tp1.partition()))); + + AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection(); + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId1) + .setProducerId(producerId) + .setProducerEpoch(producerEpoch) + .setTopics(topics0)); + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId2) + .setProducerId(producerId + 1) + .setProducerEpoch((short) (producerEpoch + 1)) + .setTopics(topics1)); + return transactions; } } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index 08865c33e2d97..f856732e47632 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -33,18 +33,17 @@ import java.util.Map; import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; public class AddPartitionsToTxnResponseTest { protected final int throttleTimeMs = 10; - protected final String topicOne = "topic1"; protected final int partitionOne = 1; protected final Errors errorOne = Errors.COORDINATOR_NOT_AVAILABLE; protected final Errors errorTwo = Errors.NOT_COORDINATOR; protected final String topicTwo = "topic2"; protected final int partitionTwo = 2; - protected TopicPartition tp1 = new TopicPartition(topicOne, partitionOne); protected TopicPartition tp2 = new TopicPartition(topicTwo, partitionTwo); protected Map expectedErrorCounts; @@ -72,7 +71,6 @@ public void testConstructorWithErrorResponse() { @Test public void testParse() { - AddPartitionsToTxnTopicResultCollection topicCollection = new AddPartitionsToTxnTopicResultCollection(); AddPartitionsToTxnTopicResult topicResult = new AddPartitionsToTxnTopicResult(); @@ -121,8 +119,26 @@ public void testParse() { assertEquals(txnTwoExpectedErrors, parsedResponse.errorsPerTransaction("txn2")); assertEquals(newExpectedErrorCounts, parsedResponse.errorCounts()); assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); - assertEquals(true, parsedResponse.shouldClientThrottle(version)); + assertTrue(parsedResponse.shouldClientThrottle(version)); } } } + + @Test + public void testBatchedErrors() { + Map txn1Errors = Collections.singletonMap(tp1, errorOne); + Map txn2Errors = Collections.singletonMap(tp1, errorOne); + + AddPartitionsToTxnResult transaction1 = AddPartitionsToTxnResponse.resultForTransaction("txn1", txn1Errors); + AddPartitionsToTxnResult transaction2 = AddPartitionsToTxnResponse.resultForTransaction("txn2", txn2Errors); + + AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); + results.add(transaction1); + results.add(transaction2); + + AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setResultsByTransaction(results)); + + assertEquals(txn1Errors, response.errorsPerTransaction("txn1")); + assertEquals(txn2Errors, response.errorsPerTransaction("txn2")); + } } diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index 1dd5de41989ea..62e95dd9d7862 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -17,7 +17,7 @@ package kafka.server -import kafka.utils.TestInfoUtils +import kafka.utils.{TestInfoUtils, TestUtils} import java.util.{Collections, Properties} import java.util.stream.{Stream => JStream} @@ -26,10 +26,12 @@ import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitio import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection +import org.apache.kafka.common.message.{FindCoordinatorRequestData, InitProducerIdRequestData} import org.apache.kafka.common.protocol.{ApiKeys, Errors} -import org.apache.kafka.common.requests.{AddPartitionsToTxnRequest, AddPartitionsToTxnResponse} +import org.apache.kafka.common.requests.FindCoordinatorRequest.CoordinatorType +import org.apache.kafka.common.requests.{AddPartitionsToTxnRequest, AddPartitionsToTxnResponse, FindCoordinatorRequest, FindCoordinatorResponse, InitProducerIdRequest, InitProducerIdResponse} import org.junit.jupiter.api.Assertions._ -import org.junit.jupiter.api.{BeforeEach, TestInfo} +import org.junit.jupiter.api.{BeforeEach, Test, TestInfo} import org.junit.jupiter.params.ParameterizedTest import org.junit.jupiter.params.provider.{Arguments, MethodSource} @@ -100,6 +102,66 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { assertTrue(errors.containsKey(nonExistentTopic)) assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION, errors.get(nonExistentTopic)) } + + @Test + def testOneSuccessOneErrorInBatchedRequest(): Unit = { + val transactionalId1 = "foobar" + + val findCoordinatorRequest = new FindCoordinatorRequest.Builder(new FindCoordinatorRequestData().setKey(transactionalId1).setKeyType(CoordinatorType.TRANSACTION.id)).build() + // First find coordinator request creates the state topic, then wait for transactional topics to be created. + connectAndReceive[FindCoordinatorResponse](findCoordinatorRequest, brokerSocketServer(brokers.head.config.brokerId)) + TestUtils.waitForAllPartitionsMetadata(brokers, "__transaction_state", 50) + val findCoordinatorResponse = connectAndReceive[FindCoordinatorResponse](findCoordinatorRequest, brokerSocketServer(brokers.head.config.brokerId)) + val coordinatorId = findCoordinatorResponse.data().coordinators().get(0).nodeId() + + val initPidRequest = new InitProducerIdRequest.Builder(new InitProducerIdRequestData().setTransactionalId(transactionalId1).setTransactionTimeoutMs(10000)).build() + val initPidResponse = connectAndReceive[InitProducerIdResponse](initPidRequest, brokerSocketServer(coordinatorId)) + + val producerId1 = initPidResponse.data().producerId() + val producerEpoch1 = initPidResponse.data().producerEpoch() + + val transactionalId2 = "barfoo" // "barfoo" maps to the same transaction coordinator + val producerId2 = 1000L + val producerEpoch2: Short = 0 + + val tp0 = new TopicPartition(topic1, 0) + + val txn1Topics = new AddPartitionsToTxnTopicCollection() + txn1Topics.add(new AddPartitionsToTxnTopic() + .setName(tp0.topic()) + .setPartitions(Collections.singletonList(tp0.partition()))) + + val txn2Topics = new AddPartitionsToTxnTopicCollection() + txn2Topics.add(new AddPartitionsToTxnTopic() + .setName(tp0.topic()) + .setPartitions(Collections.singletonList(tp0.partition()))) + + val transactions = new AddPartitionsToTxnTransactionCollection() + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId1) + .setProducerId(producerId1) + .setProducerEpoch(producerEpoch1) + .setTopics(txn1Topics)) + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId2) + .setProducerId(producerId2) + .setProducerEpoch(producerEpoch2) + .setTopics(txn2Topics)) + + val request = new AddPartitionsToTxnRequest.Builder(transactions, false).build() + + val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(coordinatorId)) + + val errors = response.allErrors() + + assertTrue(errors.containsKey(transactionalId1)) + assertTrue(errors.get(transactionalId1).containsKey(tp0)) + assertEquals(Errors.NONE, errors.get(transactionalId1).get(tp0)) + + assertTrue(errors.containsKey(transactionalId2)) + assertTrue(errors.get(transactionalId1).containsKey(tp0)) + assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, errors.get(transactionalId2).get(tp0)) + } } object AddPartitionsToTxnRequestServerTest { From 4ad3df0cb91556ae5927a6e9238b8dd5edfcea4c Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 10 Feb 2023 10:59:49 -0800 Subject: [PATCH 06/17] fixes --- .../AddPartitionsToTxnResponseTest.java | 6 ++++-- .../main/scala/kafka/server/KafkaApis.scala | 5 +++-- .../TransactionCoordinatorTest.scala | 20 +++++++++++++------ 3 files changed, 21 insertions(+), 10 deletions(-) diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index f856732e47632..cb131a5642849 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -38,14 +38,16 @@ public class AddPartitionsToTxnResponseTest { protected final int throttleTimeMs = 10; + protected final String topicOne = "topic1"; protected final int partitionOne = 1; protected final Errors errorOne = Errors.COORDINATOR_NOT_AVAILABLE; protected final Errors errorTwo = Errors.NOT_COORDINATOR; protected final String topicTwo = "topic2"; protected final int partitionTwo = 2; - protected TopicPartition tp1 = new TopicPartition(topicOne, partitionOne); - protected TopicPartition tp2 = new TopicPartition(topicTwo, partitionTwo); + protected final TopicPartition tp1 = new TopicPartition(topicOne, partitionOne); + protected final TopicPartition tp2 = new TopicPartition(topicTwo, partitionTwo); + protected Map expectedErrorCounts; protected Map errorsMap; diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 73a91398596e2..43049fb2319d6 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2424,6 +2424,7 @@ class KafkaApis(val requestChannel: RequestChannel, val transactionalId = transaction.transactionalId() val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala + // Versions < 4 come from clients and must be authorized to write for the given transaction and for the given topics. if (version < 4 && !authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) { responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) maybeSendResponse() @@ -2432,8 +2433,8 @@ class KafkaApis(val requestChannel: RequestChannel, val nonExistingTopicErrors = mutable.Map[TopicPartition, Errors]() val authorizedPartitions = mutable.Set[TopicPartition]() - val authorizedTopics = authHelper.filterByAuthorized(request.context, WRITE, TOPIC, - partitionsToAdd.filterNot(tp => Topic.isInternal(tp.topic)))(_.topic) + val authorizedTopics = if (version < 4) authHelper.filterByAuthorized(request.context, WRITE, TOPIC, + partitionsToAdd.filterNot(tp => Topic.isInternal(tp.topic)))(_.topic) else partitionsToAdd.map(_.topic).toSet for (topicPartition <- partitionsToAdd) { if (!authorizedTopics.contains(topicPartition.topic)) unauthorizedTopicErrors += topicPartition -> Errors.TOPIC_AUTHORIZATION_FAILED diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index cc873c99a5520..0050c84989b22 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -24,8 +24,6 @@ import org.apache.kafka.common.utils.{LogContext, MockTime, ProducerIdAndEpoch} import org.apache.kafka.server.util.MockScheduler import org.junit.jupiter.api.Assertions._ import org.junit.jupiter.api.Test -import org.junit.jupiter.params.ParameterizedTest -import org.junit.jupiter.params.provider.ValueSource import org.mockito.{ArgumentCaptor, ArgumentMatchers} import org.mockito.ArgumentMatchers.{any, anyInt} import org.mockito.Mockito.{mock, times, verify, when} @@ -305,14 +303,24 @@ class TransactionCoordinatorTest { ) } - @ParameterizedTest - @ValueSource(booleans = Array(true, false)) - def shouldRespondWithErrorsNoneOnAddPartitionWhenNoErrorsAndPartitionsTheSame(verifyOnly: Boolean): Unit = { + @Test + def shouldRespondWithErrorsNoneOnAddPartitionWhenNoErrorsAndPartitionsTheSame(): Unit = { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, verifyOnly, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback) + assertEquals(Errors.NONE, error) + verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) + } + + @Test + def shouldRespondWithErrorsNoneOnAddPartitionWhenOngoingVerifyOnlyAndPartitionsTheSame(): Unit = { + when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) + .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, + new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Ongoing, partitions, 0, 0))))) + + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, true, errorsCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } From 3f0c89c7d20c908aee40a6b0cb3304d649a6810e Mon Sep 17 00:00:00 2001 From: Justine Date: Tue, 21 Feb 2023 13:46:55 -0800 Subject: [PATCH 07/17] update request json --- .../common/requests/AddPartitionsToTxnRequest.java | 8 +++----- .../common/requests/AddPartitionsToTxnResponse.java | 10 +++++----- .../common/message/AddPartitionsToTxnRequest.json | 8 ++++---- .../common/message/AddPartitionsToTxnResponse.json | 8 +++++--- .../common/requests/AddPartitionsToTxnRequestTest.java | 8 ++++---- .../requests/AddPartitionsToTxnResponseTest.java | 6 +++--- core/src/main/scala/kafka/server/KafkaApis.scala | 4 ++-- .../server/AddPartitionsToTxnRequestServerTest.scala | 7 +++++-- 8 files changed, 31 insertions(+), 28 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 624b9ecb8b857..a46686fd9370b 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -70,14 +70,12 @@ public Builder(String transactionalId, .setTopics(topics); } - public Builder(AddPartitionsToTxnTransactionCollection transactions, - boolean verifyOnly) { + public Builder(AddPartitionsToTxnTransactionCollection transactions) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); this.isClientRequest = false; this.data = new AddPartitionsToTxnRequestData() - .setTransactions(transactions) - .setVerifyOnly(verifyOnly); + .setTransactions(transactions); } private AddPartitionsToTxnTopicCollection compileTopics(final List partitions) { @@ -221,7 +219,7 @@ public AddPartitionsToTxnResult errorResponseForTransaction(String transactional for (Integer partition : topic.partitions()) { partitionResult.add(new AddPartitionsToTxnPartitionResult() .setPartitionIndex(partition) - .setErrorCode(e.code())); + .setPartitionErrorCode(e.code())); } topicResult.setResults(partitionResult); topicResults.add(topicResult); diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index d2e42f744ddbd..c9f1b971fffe3 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -66,7 +66,7 @@ public AddPartitionsToTxnResponse(int throttleTimeMs, Map errors) { @@ -78,7 +78,7 @@ private static AddPartitionsToTxnTopicResultCollection topicCollectionForErrors( AddPartitionsToTxnPartitionResult partitionResult = new AddPartitionsToTxnPartitionResult() - .setErrorCode(entry.getValue().code()) + .setPartitionErrorCode(entry.getValue().code()) .setPartitionIndex(topicPartition.partition()); AddPartitionsToTxnPartitionResultCollection partitionResultCollection = resultMap.getOrDefault( @@ -120,11 +120,11 @@ public Map errors() { cachedErrorsMap = new HashMap<>(); - for (AddPartitionsToTxnTopicResult topicResult : this.data.results()) { + for (AddPartitionsToTxnTopicResult topicResult : this.data.resultsByTopicV3AndBelow()) { for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.results()) { cachedErrorsMap.put(new TopicPartition( topicResult.name(), partitionResult.partitionIndex()), - Errors.forCode(partitionResult.errorCode())); + Errors.forCode(partitionResult.partitionErrorCode())); } } return cachedErrorsMap; @@ -140,7 +140,7 @@ public Map errorsPerTransaction(String transactionalId) for (AddPartitionsToTxnTopicResult topicResult : data().resultsByTransaction().find(txnId).topicResults()) { for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.results()) { topicResults.put( - new TopicPartition(topicResult.name(), partitionResult.partitionIndex()), Errors.forCode(partitionResult.errorCode())); + new TopicPartition(topicResult.name(), partitionResult.partitionIndex()), Errors.forCode(partitionResult.partitionErrorCode())); } } return topicResults; diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json index 09a1b427a97bc..0c030ee457293 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json @@ -28,21 +28,21 @@ "validVersions": "0-4", "flexibleVersions": "3+", "fields": [ - { "name": "VerifyOnly", "type": "bool", "versions": "4+", "default": false, - "about": "Boolean to signify if we want to check if the partition is in the transaction rather than add it." }, { "name": "Transactions", "type": "[]AddPartitionsToTxnTransaction", "versions": "4+", "about": "List of transactions to add partitions to.", "fields": [ { "name": "TransactionalId", "type": "string", "versions": "4+", "mapKey": true, "entityType": "transactionalId", - "about": "The transactional id corresponding to the transaction."}, + "about": "The transactional id corresponding to the transaction." }, { "name": "ProducerId", "type": "int64", "versions": "4+", "entityType": "producerId", "about": "Current producer id in use by the transactional id." }, { "name": "ProducerEpoch", "type": "int16", "versions": "4+", "about": "Current epoch associated with the producer id." }, + { "name": "VerifyOnly", "type": "bool", "versions": "4+", "default": false, + "about": "Boolean to signify if we want to check if the partition is in the transaction rather than add it." }, { "name": "Topics", "type": "[]AddPartitionsToTxnTopic", "versions": "4+", "about": "The partitions to add to the transaction." } ]}, { "name": "TransactionalId", "type": "string", "versions": "0-3", "entityType": "transactionalId", - "about": "The transactional id corresponding to the transaction."}, + "about": "The transactional id corresponding to the transaction." }, { "name": "ProducerId", "type": "int64", "versions": "0-3", "entityType": "producerId", "about": "Current producer id in use by the transactional id." }, { "name": "ProducerEpoch", "type": "int16", "versions": "0-3", diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json index ce323155ae0ef..f53d8dcc5ec61 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json @@ -29,14 +29,16 @@ "fields": [ { "name": "ThrottleTimeMs", "type": "int32", "versions": "0+", "about": "Duration in milliseconds for which the request was throttled due to a quota violation, or zero if the request did not violate any quota." }, + { "name": "ErrorCode", "type": "int16", "versions": "4+", + "about": "The response top level error code." }, { "name": "ResultsByTransaction", "type": "[]AddPartitionsToTxnResult", "versions": "4+", "about": "Results categorized by transactional ID.", "fields": [ { "name": "TransactionalId", "type": "string", "versions": "4+", "mapKey": true, "entityType": "transactionalId", - "about": "The transactional id corresponding to the transaction."}, + "about": "The transactional id corresponding to the transaction." }, { "name": "TopicResults", "type": "[]AddPartitionsToTxnTopicResult", "versions": "4+", "about": "The results for each topic." } ]}, - { "name": "Results", "type": "[]AddPartitionsToTxnTopicResult", "versions": "0-3", + { "name": "ResultsByTopicV3AndBelow", "type": "[]AddPartitionsToTxnTopicResult", "versions": "0-3", "about": "The results for each topic." } ], "commonStructs": [ @@ -49,7 +51,7 @@ { "name": "AddPartitionsToTxnPartitionResult", "versions": "0+", "fields": [ { "name": "PartitionIndex", "type": "int32", "versions": "0+", "mapKey": true, "about": "The partition indexes." }, - { "name": "ErrorCode", "type": "int16", "versions": "0+", + { "name": "PartitionErrorCode", "type": "int16", "versions": "0+", "about": "The response error code." } ]} ] diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java index bf1046abe85ee..8b9f6b0421f1a 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java @@ -66,15 +66,13 @@ public void testConstructor(short version) { assertEquals(partitions, request.partitions()); } else { AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); - boolean verifyOnly = true; - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions, verifyOnly); + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions); request = builder.build(version); AddPartitionsToTxnTransaction reqTxn1 = request.data().transactions().find(transactionalId1); AddPartitionsToTxnTransaction reqTxn2 = request.data().transactions().find(transactionalId2); - assertEquals(verifyOnly, request.data().verifyOnly()); assertEquals(transactions.find(transactionalId1), reqTxn1); assertEquals(transactions.find(transactionalId2), reqTxn2); } @@ -89,7 +87,7 @@ public void testBatchedRequests() { AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); boolean verifyOnly = true; - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions, verifyOnly); + AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions); AddPartitionsToTxnRequest request = builder.build(ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion()); Map> expectedMap = new HashMap<>(); @@ -143,11 +141,13 @@ private AddPartitionsToTxnTransactionCollection createTwoTransactionCollection() .setTransactionalId(transactionalId1) .setProducerId(producerId) .setProducerEpoch(producerEpoch) + .setVerifyOnly(true) .setTopics(topics0)); transactions.add(new AddPartitionsToTxnTransaction() .setTransactionalId(transactionalId2) .setProducerId(producerId + 1) .setProducerEpoch((short) (producerEpoch + 1)) + .setVerifyOnly(false) .setTopics(topics1)); return transactions; } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index cb131a5642849..3188a80c6b74e 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -79,11 +79,11 @@ public void testParse() { topicResult.setName(topicOne); topicResult.results().add(new AddPartitionsToTxnPartitionResult() - .setErrorCode(errorOne.code()) + .setPartitionErrorCode(errorOne.code()) .setPartitionIndex(partitionOne)); topicResult.results().add(new AddPartitionsToTxnPartitionResult() - .setErrorCode(errorTwo.code()) + .setPartitionErrorCode(errorTwo.code()) .setPartitionIndex(partitionTwo)); topicCollection.add(topicResult); @@ -92,7 +92,7 @@ public void testParse() { if (version < 4) { AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() - .setResults(topicCollection) + .setResultsByTopicV3AndBelow(topicCollection) .setThrottleTimeMs(throttleTimeMs); AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 43049fb2319d6..25878f8c3bedf 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2402,7 +2402,7 @@ class KafkaApis(val requestChannel: RequestChannel, // There will only be one response in data. Add it to the response data object. val data = new AddPartitionsToTxnResponseData() responses.forEach(result => { - data.setResults(result.topicResults()) + data.setResultsByTopicV3AndBelow(result.topicResults()) data.setThrottleTimeMs(requestThrottleMs) }) new AddPartitionsToTxnResponse(data) @@ -2473,7 +2473,7 @@ class KafkaApis(val requestChannel: RequestChannel, transaction.producerId, transaction.producerEpoch, authorizedPartitions, - addPartitionsToTxnRequest.data.verifyOnly, + transaction.verifyOnly, sendResponseCallback, requestLocal) } diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index 62e95dd9d7862..c20847b820d0b 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -85,8 +85,9 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { .setTransactionalId(transactionalId) .setProducerId(producerId) .setProducerEpoch(producerEpoch) + .setVerifyOnly(false) .setTopics(topics)) - new AddPartitionsToTxnRequest.Builder(transactions, false).build() + new AddPartitionsToTxnRequest.Builder(transactions).build() } val leaderId = brokers.head.config.brokerId @@ -141,14 +142,16 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { .setTransactionalId(transactionalId1) .setProducerId(producerId1) .setProducerEpoch(producerEpoch1) + .setVerifyOnly(false) .setTopics(txn1Topics)) transactions.add(new AddPartitionsToTxnTransaction() .setTransactionalId(transactionalId2) .setProducerId(producerId2) .setProducerEpoch(producerEpoch2) + .setVerifyOnly(false) .setTopics(txn2Topics)) - val request = new AddPartitionsToTxnRequest.Builder(transactions, false).build() + val request = new AddPartitionsToTxnRequest.Builder(transactions).build() val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(coordinatorId)) From 791d56bd84a28e546ba278182c6b1c397d58de94 Mon Sep 17 00:00:00 2001 From: Justine Date: Tue, 21 Feb 2023 15:05:23 -0800 Subject: [PATCH 08/17] new builders --- .../internals/TransactionManager.java | 2 +- .../requests/AddPartitionsToTxnRequest.java | 88 ++++++++++--------- .../message/AddPartitionsToTxnRequest.json | 8 +- .../internals/TransactionManagerTest.java | 6 +- .../kafka/common/message/MessageTest.java | 8 +- .../AddPartitionsToTxnRequestTest.java | 14 +-- .../common/requests/RequestResponseTest.java | 21 ++++- .../kafka/api/AuthorizerIntegrationTest.scala | 2 +- .../AddPartitionsToTxnRequestServerTest.scala | 6 +- .../unit/kafka/server/KafkaApisTest.scala | 4 +- .../unit/kafka/server/RequestQuotaTest.scala | 2 +- 11 files changed, 90 insertions(+), 71 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java b/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java index de5a6ced41c85..70657743f2581 100644 --- a/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java +++ b/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java @@ -1052,7 +1052,7 @@ private TxnRequestHandler addPartitionsToTransactionHandler() { pendingPartitionsInTransaction.addAll(newPartitionsInTransaction); newPartitionsInTransaction.clear(); AddPartitionsToTxnRequest.Builder builder = - new AddPartitionsToTxnRequest.Builder(transactionalId, + AddPartitionsToTxnRequest.Builder.forClient(transactionalId, producerIdAndEpoch.producerId, producerIdAndEpoch.epoch, new ArrayList<>(pendingPartitionsInTransaction)); diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index a46686fd9370b..f015949c15b9a 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -51,34 +51,37 @@ public class AddPartitionsToTxnRequest extends AbstractRequest { public static class Builder extends AbstractRequest.Builder { public final AddPartitionsToTxnRequestData data; - public final boolean isClientRequest; - - // Only used for versions < 4 - public Builder(String transactionalId, - long producerId, - short producerEpoch, - List partitions) { - super(ApiKeys.ADD_PARTITIONS_TO_TXN); - this.isClientRequest = true; - - AddPartitionsToTxnTopicCollection topics = compileTopics(partitions); - - this.data = new AddPartitionsToTxnRequestData() - .setTransactionalId(transactionalId) - .setProducerId(producerId) - .setProducerEpoch(producerEpoch) - .setTopics(topics); + + public static Builder forClient(String transactionalId, + long producerId, + short producerEpoch, + List partitions) { + + AddPartitionsToTxnTopicCollection topics = buildTxnTopicCollection(partitions); + + return new Builder(ApiKeys.ADD_PARTITIONS_TO_TXN.oldestVersion(), + (short) 3, + new AddPartitionsToTxnRequestData() + .setV3AndBelowTransactionalId(transactionalId) + .setV3AndBelowProducerId(producerId) + .setV3AndBelowProducerEpoch(producerEpoch) + .setV3AndBelowTopics(topics)); } + + public static Builder forBroker(AddPartitionsToTxnTransactionCollection transactions) { + return new Builder((short) 4, + ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion(), + new AddPartitionsToTxnRequestData() + .setTransactions(transactions)); + } + + public Builder(short minVersion, short maxVersion, AddPartitionsToTxnRequestData data) { + super(ApiKeys.ADD_PARTITIONS_TO_TXN, minVersion, maxVersion); - public Builder(AddPartitionsToTxnTransactionCollection transactions) { - super(ApiKeys.ADD_PARTITIONS_TO_TXN); - this.isClientRequest = false; - - this.data = new AddPartitionsToTxnRequestData() - .setTransactions(transactions); + this.data = data; } - private AddPartitionsToTxnTopicCollection compileTopics(final List partitions) { + private static AddPartitionsToTxnTopicCollection buildTxnTopicCollection(final List partitions) { Map> partitionMap = new HashMap<>(); for (TopicPartition topicPartition : partitions) { String topicName = topicPartition.topic(); @@ -103,20 +106,7 @@ private AddPartitionsToTxnTopicCollection compileTopics(final List 3) ? 3 : version; - return new AddPartitionsToTxnRequest(data, clampedVersion); - } - - // Only used for versions < 4 - static List getPartitions(AddPartitionsToTxnRequestData data) { - List partitions = new ArrayList<>(); - - for (AddPartitionsToTxnTopic topicCollection : data.topics()) { - for (Integer partition : topicCollection.partitions()) { - partitions.add(new TopicPartition(topicCollection.name(), partition)); - } - } - return partitions; + return new AddPartitionsToTxnRequest(data, version); } @Override @@ -136,9 +126,21 @@ public List partitions() { if (cachedPartitions != null) { return cachedPartitions; } - cachedPartitions = Builder.getPartitions(data); + cachedPartitions = getPartitions(data); return cachedPartitions; } + + // Only used for versions < 4 + static List getPartitions(AddPartitionsToTxnRequestData data) { + List partitions = new ArrayList<>(); + + for (AddPartitionsToTxnTopic topicCollection : data.v3AndBelowTopics()) { + for (Integer partition : topicCollection.partitions()) { + partitions.add(new TopicPartition(topicCollection.name(), partition)); + } + } + return partitions; + } private List partitionsForTransaction(String transaction) { if (cachedPartitionsByTransaction == null) { @@ -177,10 +179,10 @@ public AddPartitionsToTxnRequest normalizeRequest() { private AddPartitionsToTxnTransactionCollection singletonTransaction() { AddPartitionsToTxnTransactionCollection singleTxn = new AddPartitionsToTxnTransactionCollection(); singleTxn.add(new AddPartitionsToTxnTransaction() - .setTransactionalId(data.transactionalId()) - .setProducerId(data.producerId()) - .setProducerEpoch(data.producerEpoch()) - .setTopics(data.topics())); + .setTransactionalId(data.v3AndBelowTransactionalId()) + .setProducerId(data.v3AndBelowProducerId()) + .setProducerEpoch(data.v3AndBelowProducerEpoch()) + .setTopics(data.v3AndBelowTopics())); return singleTxn; } diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json index 0c030ee457293..fbd34f73ed80c 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json @@ -41,13 +41,13 @@ { "name": "Topics", "type": "[]AddPartitionsToTxnTopic", "versions": "4+", "about": "The partitions to add to the transaction." } ]}, - { "name": "TransactionalId", "type": "string", "versions": "0-3", "entityType": "transactionalId", + { "name": "V3AndBelowTransactionalId", "type": "string", "versions": "0-3", "entityType": "transactionalId", "about": "The transactional id corresponding to the transaction." }, - { "name": "ProducerId", "type": "int64", "versions": "0-3", "entityType": "producerId", + { "name": "V3AndBelowProducerId", "type": "int64", "versions": "0-3", "entityType": "producerId", "about": "Current producer id in use by the transactional id." }, - { "name": "ProducerEpoch", "type": "int16", "versions": "0-3", + { "name": "V3AndBelowProducerEpoch", "type": "int16", "versions": "0-3", "about": "Current epoch associated with the producer id." }, - { "name": "Topics", "type": "[]AddPartitionsToTxnTopic", "versions": "0-3", + { "name": "V3AndBelowTopics", "type": "[]AddPartitionsToTxnTopic", "versions": "0-3", "about": "The partitions to add to the transaction." } ], "commonStructs": [ diff --git a/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java b/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java index ce9b80522076e..8d0349a8f5aea 100644 --- a/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java +++ b/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java @@ -3536,10 +3536,10 @@ private MockClient.RequestMatcher addPartitionsRequestMatcher(final TopicPartiti final short epoch, final long producerId) { return body -> { AddPartitionsToTxnRequest addPartitionsToTxnRequest = (AddPartitionsToTxnRequest) body; - assertEquals(producerId, addPartitionsToTxnRequest.data().producerId()); - assertEquals(epoch, addPartitionsToTxnRequest.data().producerEpoch()); + assertEquals(producerId, addPartitionsToTxnRequest.data().v3AndBelowProducerId()); + assertEquals(epoch, addPartitionsToTxnRequest.data().v3AndBelowProducerEpoch()); assertEquals(singletonList(topicPartition), addPartitionsToTxnRequest.partitions()); - assertEquals(transactionalId, addPartitionsToTxnRequest.data().transactionalId()); + assertEquals(transactionalId, addPartitionsToTxnRequest.data().v3AndBelowTransactionalId()); return true; }; } diff --git a/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java b/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java index 3fcd0071c49a2..1f6ad3ed1133a 100644 --- a/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java +++ b/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java @@ -97,10 +97,10 @@ public void testAddOffsetsToTxnVersions() throws Exception { @Test public void testAddPartitionsToTxnVersions() throws Exception { testAllMessageRoundTrips(new AddPartitionsToTxnRequestData(). - setTransactionalId("blah"). - setProducerId(0xbadcafebadcafeL). - setProducerEpoch((short) 30000). - setTopics(new AddPartitionsToTxnTopicCollection(singletonList( + setV3AndBelowTransactionalId("blah"). + setV3AndBelowProducerId(0xbadcafebadcafeL). + setV3AndBelowProducerEpoch((short) 30000). + setV3AndBelowTopics(new AddPartitionsToTxnTopicCollection(singletonList( new AddPartitionsToTxnTopic(). setName("Topic"). setPartitions(singletonList(1))).iterator()))); diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java index 8b9f6b0421f1a..59ec7ff9bb795 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java @@ -57,17 +57,17 @@ public void testConstructor(short version) { partitions.add(tp0); partitions.add(tp1); - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId1, producerId, producerEpoch, partitions); + AddPartitionsToTxnRequest.Builder builder = AddPartitionsToTxnRequest.Builder.forClient(transactionalId1, producerId, producerEpoch, partitions); request = builder.build(version); - assertEquals(transactionalId1, request.data().transactionalId()); - assertEquals(producerId, request.data().producerId()); - assertEquals(producerEpoch, request.data().producerEpoch()); + assertEquals(transactionalId1, request.data().v3AndBelowTransactionalId()); + assertEquals(producerId, request.data().v3AndBelowProducerId()); + assertEquals(producerEpoch, request.data().v3AndBelowProducerEpoch()); assertEquals(partitions, request.partitions()); } else { AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions); + AddPartitionsToTxnRequest.Builder builder = AddPartitionsToTxnRequest.Builder.forBroker(transactions); request = builder.build(version); AddPartitionsToTxnTransaction reqTxn1 = request.data().transactions().find(transactionalId1); @@ -87,7 +87,7 @@ public void testBatchedRequests() { AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); boolean verifyOnly = true; - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactions); + AddPartitionsToTxnRequest.Builder builder = AddPartitionsToTxnRequest.Builder.forBroker(transactions); AddPartitionsToTxnRequest request = builder.build(ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion()); Map> expectedMap = new HashMap<>(); @@ -115,7 +115,7 @@ public void testNormalizeRequest() { partitions.add(tp0); partitions.add(tp1); - AddPartitionsToTxnRequest.Builder builder = new AddPartitionsToTxnRequest.Builder(transactionalId1, producerId, producerEpoch, partitions); + AddPartitionsToTxnRequest.Builder builder = AddPartitionsToTxnRequest.Builder.forClient(transactionalId1, producerId, producerEpoch, partitions); AddPartitionsToTxnRequest request = builder.build((short) 3); AddPartitionsToTxnRequest singleton = request.normalizeRequest(); diff --git a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java index b4212f10cb128..a87a5d72a1e3b 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java @@ -37,6 +37,10 @@ import org.apache.kafka.common.errors.UnsupportedVersionException; import org.apache.kafka.common.message.AddOffsetsToTxnRequestData; import org.apache.kafka.common.message.AddOffsetsToTxnResponseData; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopic; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection; import org.apache.kafka.common.message.AllocateProducerIdsRequestData; import org.apache.kafka.common.message.AllocateProducerIdsResponseData; import org.apache.kafka.common.message.AlterClientQuotasResponseData; @@ -2554,8 +2558,21 @@ private OffsetsForLeaderEpochResponse createLeaderEpochResponse() { } private AddPartitionsToTxnRequest createAddPartitionsToTxnRequest(short version) { - return new AddPartitionsToTxnRequest.Builder("tid", 21L, (short) 42, - singletonList(new TopicPartition("topic", 73))).build(version); + if (version < 3) { + return AddPartitionsToTxnRequest.Builder.forClient("tid", 21L, (short) 42, + singletonList(new TopicPartition("topic", 73))).build(version); + } else { + AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection(); + AddPartitionsToTxnTopicCollection topics = new AddPartitionsToTxnTopicCollection(); + topics.add(new AddPartitionsToTxnTopic().setName("topic").setPartitions(Collections.singletonList(73))); + transactions.add(new AddPartitionsToTxnTransaction() + .setTransactionalId("tid") + .setProducerId(21L) + .setProducerEpoch((short) 42) + .setVerifyOnly(false) + .setTopics(topics)); + return AddPartitionsToTxnRequest.Builder.forBroker(transactions).build(version); + } } private AddPartitionsToTxnResponse createAddPartitionsToTxnResponse() { diff --git a/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala b/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala index 3b4f893eb552c..fca36644a3154 100644 --- a/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala +++ b/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala @@ -672,7 +672,7 @@ class AuthorizerIntegrationTest extends BaseRequestTest { private def describeLogDirsRequest = new DescribeLogDirsRequest.Builder(new DescribeLogDirsRequestData().setTopics(new DescribeLogDirsRequestData.DescribableLogDirTopicCollection(Collections.singleton( new DescribeLogDirsRequestData.DescribableLogDirTopic().setTopic(tp.topic).setPartitions(Collections.singletonList(tp.partition))).iterator()))).build() - private def addPartitionsToTxnRequest = new AddPartitionsToTxnRequest.Builder(transactionalId, 1, 1, Collections.singletonList(tp)).build() + private def addPartitionsToTxnRequest = AddPartitionsToTxnRequest.Builder.forClient(transactionalId, 1, 1, Collections.singletonList(tp)).build() private def addOffsetsToTxnRequest = new AddOffsetsToTxnRequest.Builder( new AddOffsetsToTxnRequestData() diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index c20847b820d0b..c04bc2b4a656f 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -65,7 +65,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val request = if (version < 4) { - new AddPartitionsToTxnRequest.Builder( + AddPartitionsToTxnRequest.Builder.forClient( transactionalId, producerId, producerEpoch, @@ -87,7 +87,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { .setProducerEpoch(producerEpoch) .setVerifyOnly(false) .setTopics(topics)) - new AddPartitionsToTxnRequest.Builder(transactions).build() + AddPartitionsToTxnRequest.Builder.forBroker(transactions).build() } val leaderId = brokers.head.config.brokerId @@ -151,7 +151,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { .setVerifyOnly(false) .setTopics(txn2Topics)) - val request = new AddPartitionsToTxnRequest.Builder(transactions).build() + val request = AddPartitionsToTxnRequest.Builder.forBroker(transactions).build() val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(coordinatorId)) diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index 1ef7cc4446dc0..43559705c488d 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -2003,7 +2003,7 @@ class KafkaApisTest { val partition = 1 val topicPartition = new TopicPartition(topic, partition) - val addPartitionsToTxnRequest = new AddPartitionsToTxnRequest.Builder( + val addPartitionsToTxnRequest = AddPartitionsToTxnRequest.Builder.forClient( transactionalId, producerId, epoch, @@ -2153,7 +2153,7 @@ class KafkaApisTest { reset(replicaManager, clientRequestQuotaManager, requestChannel) val invalidTopicPartition = new TopicPartition(topic, invalidPartitionId) - val addPartitionsToTxnRequest = new AddPartitionsToTxnRequest.Builder( + val addPartitionsToTxnRequest = AddPartitionsToTxnRequest.Builder.forClient( "txnlId", 15L, 0.toShort, List(invalidTopicPartition).asJava ).build() val request = buildRequest(addPartitionsToTxnRequest) diff --git a/core/src/test/scala/unit/kafka/server/RequestQuotaTest.scala b/core/src/test/scala/unit/kafka/server/RequestQuotaTest.scala index 766861d0a3476..d45b658ab2d44 100644 --- a/core/src/test/scala/unit/kafka/server/RequestQuotaTest.scala +++ b/core/src/test/scala/unit/kafka/server/RequestQuotaTest.scala @@ -426,7 +426,7 @@ class RequestQuotaTest extends BaseRequestTest { OffsetsForLeaderEpochRequest.Builder.forConsumer(epochs) case ApiKeys.ADD_PARTITIONS_TO_TXN => - new AddPartitionsToTxnRequest.Builder("test-transactional-id", 1, 0, List(tp).asJava) + AddPartitionsToTxnRequest.Builder.forClient("test-transactional-id", 1, 0, List(tp).asJava) case ApiKeys.ADD_OFFSETS_TO_TXN => new AddOffsetsToTxnRequest.Builder(new AddOffsetsToTxnRequestData() From f8dbaa1fe45cdd2933c542c734db02e0fc8c5f0c Mon Sep 17 00:00:00 2001 From: Justine Date: Wed, 22 Feb 2023 13:46:12 -0800 Subject: [PATCH 09/17] Fix up requests/responses --- .../internals/TransactionManager.java | 2 +- .../requests/AddPartitionsToTxnRequest.java | 133 +++++++----------- .../requests/AddPartitionsToTxnResponse.java | 116 ++++++--------- .../message/AddPartitionsToTxnResponse.json | 2 +- .../producer/internals/SenderTest.java | 26 ++-- .../internals/TransactionManagerTest.java | 28 +++- .../AddPartitionsToTxnRequestTest.java | 11 +- .../AddPartitionsToTxnResponseTest.java | 91 ++++++------ .../common/requests/RequestResponseTest.java | 6 +- .../main/scala/kafka/server/KafkaApis.scala | 19 ++- .../kafka/api/AuthorizerIntegrationTest.scala | 2 +- .../TransactionCoordinatorTest.scala | 2 +- .../AddPartitionsToTxnRequestServerTest.scala | 8 +- 13 files changed, 205 insertions(+), 241 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java b/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java index 70657743f2581..a41792ada07e0 100644 --- a/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java +++ b/clients/src/main/java/org/apache/kafka/clients/producer/internals/TransactionManager.java @@ -1328,7 +1328,7 @@ Priority priority() { @Override public void handleResponse(AbstractResponse response) { AddPartitionsToTxnResponse addPartitionsToTxnResponse = (AddPartitionsToTxnResponse) response; - Map errors = addPartitionsToTxnResponse.errors(); + Map errors = addPartitionsToTxnResponse.errors().get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID); boolean hasPartitionErrors = false; Set unauthorizedTopics = new HashSet<>(); retryBackoffMs = TransactionManager.this.retryBackoffMs; diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index f015949c15b9a..2d86eceec5290 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -43,10 +43,6 @@ public class AddPartitionsToTxnRequest extends AbstractRequest { private final AddPartitionsToTxnRequestData data; - private List cachedPartitions = null; - - private Map> cachedPartitionsByTransaction = null; - private final short version; public static class Builder extends AbstractRequest.Builder { @@ -69,10 +65,9 @@ public static Builder forClient(String transactionalId, } public static Builder forBroker(AddPartitionsToTxnTransactionCollection transactions) { - return new Builder((short) 4, - ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion(), - new AddPartitionsToTxnRequestData() - .setTransactions(transactions)); + return new Builder((short) 4, ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion(), + new AddPartitionsToTxnRequestData() + .setTransactions(transactions)); } public Builder(short minVersion, short maxVersion, AddPartitionsToTxnRequestData data) { @@ -98,8 +93,8 @@ private static AddPartitionsToTxnTopicCollection buildTxnTopicCollection(final L AddPartitionsToTxnTopicCollection topics = new AddPartitionsToTxnTopicCollection(); for (Map.Entry> partitionEntry : partitionMap.entrySet()) { topics.add(new AddPartitionsToTxnTopic() - .setName(partitionEntry.getKey()) - .setPartitions(partitionEntry.getValue())); + .setName(partitionEntry.getKey()) + .setPartitions(partitionEntry.getValue())); } return topics; } @@ -120,114 +115,86 @@ public AddPartitionsToTxnRequest(final AddPartitionsToTxnRequestData data, short this.data = data; this.version = version; } - - // Only used for versions < 4 - public List partitions() { - if (cachedPartitions != null) { - return cachedPartitions; + + @Override + public AddPartitionsToTxnRequestData data() { + return data; + } + + @Override + public AddPartitionsToTxnResponse getErrorResponse(int throttleTimeMs, Throwable e) { + Errors error = Errors.forException(e); + AddPartitionsToTxnResponseData response = new AddPartitionsToTxnResponseData(); + if (version < 4) { + response.setResultsByTopicV3AndBelow(errorResponseForTopics(data.v3AndBelowTopics(), error)); + } else { + AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); + for (AddPartitionsToTxnTransaction transaction : data().transactions()) { + results.add(errorResponseForTransaction(transaction.transactionalId(), error)); + } + response.setResultsByTransaction(results); + response.setErrorCode(error.code()); } - cachedPartitions = getPartitions(data); - return cachedPartitions; + response.setThrottleTimeMs(throttleTimeMs); + return new AddPartitionsToTxnResponse(response); } - // Only used for versions < 4 - static List getPartitions(AddPartitionsToTxnRequestData data) { + public static List getPartitions(AddPartitionsToTxnTopicCollection topics) { List partitions = new ArrayList<>(); - for (AddPartitionsToTxnTopic topicCollection : data.v3AndBelowTopics()) { + for (AddPartitionsToTxnTopic topicCollection : topics) { for (Integer partition : topicCollection.partitions()) { partitions.add(new TopicPartition(topicCollection.name(), partition)); } } return partitions; } - - private List partitionsForTransaction(String transaction) { - if (cachedPartitionsByTransaction == null) { - cachedPartitionsByTransaction = new HashMap<>(); - } - - return cachedPartitionsByTransaction.computeIfAbsent(transaction, txn -> { - List partitions = new ArrayList<>(); - for (AddPartitionsToTxnTopic topicCollection : data.transactions().find(txn).topics()) { - for (Integer partition : topicCollection.partitions()) { - partitions.add(new TopicPartition(topicCollection.name(), partition)); - } - } - return partitions; - }); - } - + public Map> partitionsByTransaction() { - if (cachedPartitionsByTransaction != null && cachedPartitionsByTransaction.size() == data.transactions().size()) { - return cachedPartitionsByTransaction; - } - + Map> partitionsByTransaction = new HashMap<>(); for (AddPartitionsToTxnTransaction transaction : data.transactions()) { - if (cachedPartitionsByTransaction == null || !cachedPartitionsByTransaction.containsKey(transaction.transactionalId())) { - partitionsForTransaction(transaction.transactionalId()); - } + List partitions = getPartitions(transaction.topics()); + partitionsByTransaction.put(transaction.transactionalId(), partitions); } - return cachedPartitionsByTransaction; + return partitionsByTransaction; } - + // Takes a version 3 or below request and returns a v4+ singleton (one transaction ID) request. public AddPartitionsToTxnRequest normalizeRequest() { return new AddPartitionsToTxnRequest(new AddPartitionsToTxnRequestData().setTransactions(singletonTransaction()), version); } - + private AddPartitionsToTxnTransactionCollection singletonTransaction() { AddPartitionsToTxnTransactionCollection singleTxn = new AddPartitionsToTxnTransactionCollection(); singleTxn.add(new AddPartitionsToTxnTransaction() - .setTransactionalId(data.v3AndBelowTransactionalId()) - .setProducerId(data.v3AndBelowProducerId()) - .setProducerEpoch(data.v3AndBelowProducerEpoch()) - .setTopics(data.v3AndBelowTopics())); + .setTransactionalId(data.v3AndBelowTransactionalId()) + .setProducerId(data.v3AndBelowProducerId()) + .setProducerEpoch(data.v3AndBelowProducerEpoch()) + .setTopics(data.v3AndBelowTopics())); return singleTxn; } - - @Override - public AddPartitionsToTxnRequestData data() { - return data; - } - - @Override - public AddPartitionsToTxnResponse getErrorResponse(int throttleTimeMs, Throwable e) { - Errors error = Errors.forException(e); - if (version < 4) { - final HashMap errors = new HashMap<>(); - for (TopicPartition partition : partitions()) { - errors.put(partition, error); - } - return new AddPartitionsToTxnResponse(throttleTimeMs, errors); - } else { - AddPartitionsToTxnResponseData response = new AddPartitionsToTxnResponseData(); - AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); - for (AddPartitionsToTxnTransaction transaction : data().transactions()) { - results.add(errorResponseForTransaction(transaction.transactionalId(), error)); - } - response.setResultsByTransaction(results); - response.setThrottleTimeMs(throttleTimeMs); - return new AddPartitionsToTxnResponse(response); - } - } public AddPartitionsToTxnResult errorResponseForTransaction(String transactionalId, Errors e) { AddPartitionsToTxnResult txnResult = new AddPartitionsToTxnResult().setTransactionalId(transactionalId); + AddPartitionsToTxnTopicResultCollection topicResults = errorResponseForTopics(data.transactions().find(transactionalId).topics(), e); + txnResult.setTopicResults(topicResults); + return txnResult; + } + + private AddPartitionsToTxnTopicResultCollection errorResponseForTopics(AddPartitionsToTxnTopicCollection topics, Errors e) { AddPartitionsToTxnTopicResultCollection topicResults = new AddPartitionsToTxnTopicResultCollection(); - for (AddPartitionsToTxnTopic topic : data.transactions().find(transactionalId).topics()) { + for (AddPartitionsToTxnTopic topic : topics) { AddPartitionsToTxnTopicResult topicResult = new AddPartitionsToTxnTopicResult().setName(topic.name()); AddPartitionsToTxnPartitionResultCollection partitionResult = new AddPartitionsToTxnPartitionResultCollection(); for (Integer partition : topic.partitions()) { partitionResult.add(new AddPartitionsToTxnPartitionResult() - .setPartitionIndex(partition) - .setPartitionErrorCode(e.code())); + .setPartitionIndex(partition) + .setPartitionErrorCode(e.code())); } - topicResult.setResults(partitionResult); + topicResult.setResultsByPartition(partitionResult); topicResults.add(topicResult); } - txnResult.setTopicResults(topicResults); - return txnResult; + return topicResults; } public static AddPartitionsToTxnRequest parse(ByteBuffer buffer, short version) { diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index c9f1b971fffe3..f305cb805bbee 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -50,36 +50,48 @@ public class AddPartitionsToTxnResponse extends AbstractResponse { private final AddPartitionsToTxnResponseData data; - - private Map cachedErrorsMap = null; - private Map> cachedAllErrorsMap = null; + public static final String V3_AND_BELOW_TXN_ID = ""; public AddPartitionsToTxnResponse(AddPartitionsToTxnResponseData data) { super(ApiKeys.ADD_PARTITIONS_TO_TXN); this.data = data; } - // Only used for versions < 4 - public AddPartitionsToTxnResponse(int throttleTimeMs, Map errors) { - super(ApiKeys.ADD_PARTITIONS_TO_TXN); + @Override + public int throttleTimeMs() { + return data.throttleTimeMs(); + } - this.data = new AddPartitionsToTxnResponseData() - .setThrottleTimeMs(throttleTimeMs) - .setResultsByTopicV3AndBelow(topicCollectionForErrors(errors)); + @Override + public void maybeSetThrottleTimeMs(int throttleTimeMs) { + data.setThrottleTimeMs(throttleTimeMs); } - + + public Map> errors() { + Map> errorsMap = new HashMap<>(); + + errorsMap.put(V3_AND_BELOW_TXN_ID, errorsForTransaction(this.data.resultsByTopicV3AndBelow())); + + for (AddPartitionsToTxnResult result : this.data.resultsByTransaction()) { + String transactionalId = result.transactionalId(); + errorsMap.put(transactionalId, errorsForTransaction(data().resultsByTransaction().find(transactionalId).topicResults())); + } + + return errorsMap; + } + private static AddPartitionsToTxnTopicResultCollection topicCollectionForErrors(Map errors) { Map resultMap = new HashMap<>(); - + for (Map.Entry entry : errors.entrySet()) { TopicPartition topicPartition = entry.getKey(); String topicName = topicPartition.topic(); AddPartitionsToTxnPartitionResult partitionResult = new AddPartitionsToTxnPartitionResult() - .setPartitionErrorCode(entry.getValue().code()) - .setPartitionIndex(topicPartition.partition()); + .setPartitionErrorCode(entry.getValue().code()) + .setPartitionIndex(topicPartition.partition()); AddPartitionsToTxnPartitionResultCollection partitionResultCollection = resultMap.getOrDefault( topicName, new AddPartitionsToTxnPartitionResultCollection() @@ -92,8 +104,8 @@ topicName, new AddPartitionsToTxnPartitionResultCollection() AddPartitionsToTxnTopicResultCollection topicCollection = new AddPartitionsToTxnTopicResultCollection(); for (Map.Entry entry : resultMap.entrySet()) { topicCollection.add(new AddPartitionsToTxnTopicResult() - .setName(entry.getKey()) - .setResults(entry.getValue())); + .setName(entry.getKey()) + .setResultsByPartition(entry.getValue())); } return topicCollection; } @@ -102,74 +114,28 @@ public static AddPartitionsToTxnResult resultForTransaction(String transactional return new AddPartitionsToTxnResult().setTransactionalId(transactionalId).setTopicResults(topicCollectionForErrors(errors)); } - @Override - public int throttleTimeMs() { - return data.throttleTimeMs(); - } - - @Override - public void maybeSetThrottleTimeMs(int throttleTimeMs) { - data.setThrottleTimeMs(throttleTimeMs); - } - - // Only used for versions < 4 - public Map errors() { - if (cachedErrorsMap != null) { - return cachedErrorsMap; - } - - cachedErrorsMap = new HashMap<>(); - - for (AddPartitionsToTxnTopicResult topicResult : this.data.resultsByTopicV3AndBelow()) { - for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.results()) { - cachedErrorsMap.put(new TopicPartition( - topicResult.name(), partitionResult.partitionIndex()), - Errors.forCode(partitionResult.partitionErrorCode())); - } - } - return cachedErrorsMap; - } - - public Map errorsPerTransaction(String transactionalId) { - if (cachedAllErrorsMap == null) { - cachedAllErrorsMap = new HashMap<>(); - } - - return cachedAllErrorsMap.computeIfAbsent(transactionalId, txnId -> { - Map topicResults = new HashMap<>(); - for (AddPartitionsToTxnTopicResult topicResult : data().resultsByTransaction().find(txnId).topicResults()) { - for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.results()) { - topicResults.put( - new TopicPartition(topicResult.name(), partitionResult.partitionIndex()), Errors.forCode(partitionResult.partitionErrorCode())); - } - } - return topicResults; - }); + public AddPartitionsToTxnTopicResultCollection getTransactionTopicResults(String transactionalId) { + return data.resultsByTransaction().find(transactionalId).topicResults(); } - - public Map> allErrors() { - if (cachedAllErrorsMap != null && cachedAllErrorsMap.size() == data.resultsByTransaction().size()) { - return cachedAllErrorsMap; - } - for (AddPartitionsToTxnResult result : this.data.resultsByTransaction()) { - if (cachedAllErrorsMap == null || !cachedAllErrorsMap.containsKey(result.transactionalId())) { - errorsPerTransaction(result.transactionalId()); + public Map errorsForTransaction(AddPartitionsToTxnTopicResultCollection topicCollection) { + Map topicResults = new HashMap<>(); + for (AddPartitionsToTxnTopicResult topicResult : topicCollection) { + for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.resultsByPartition()) { + topicResults.put( + new TopicPartition(topicResult.name(), partitionResult.partitionIndex()), Errors.forCode(partitionResult.partitionErrorCode())); } } - return cachedAllErrorsMap; + return topicResults; } @Override public Map errorCounts() { - if (data.resultsByTransaction().size() > 0) { - List allErrors = new ArrayList<>(); - allErrors().forEach((txnId, errors) -> - allErrors.addAll(errors.values()) - ); - return errorCounts(allErrors); - } - return errorCounts(errors().values()); + List allErrors = new ArrayList<>(); + errors().forEach((txnId, errors) -> + allErrors.addAll(errors.values()) + ); + return errorCounts(allErrors); } @Override diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json index f53d8dcc5ec61..54a4a92614b44 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json @@ -45,7 +45,7 @@ { "name": "AddPartitionsToTxnTopicResult", "versions": "0+", "fields": [ { "name": "Name", "type": "string", "versions": "0+", "mapKey": true, "entityType": "topicName", "about": "The topic name." }, - { "name": "Results", "type": "[]AddPartitionsToTxnPartitionResult", "versions": "0+", + { "name": "ResultsByPartition", "type": "[]AddPartitionsToTxnPartitionResult", "versions": "0+", "about": "The results for each partition" } ]}, { "name": "AddPartitionsToTxnPartitionResult", "versions": "0+", "fields": [ diff --git a/clients/src/test/java/org/apache/kafka/clients/producer/internals/SenderTest.java b/clients/src/test/java/org/apache/kafka/clients/producer/internals/SenderTest.java index bdbc1bd92e936..adee050da4230 100644 --- a/clients/src/test/java/org/apache/kafka/clients/producer/internals/SenderTest.java +++ b/clients/src/test/java/org/apache/kafka/clients/producer/internals/SenderTest.java @@ -40,6 +40,7 @@ import org.apache.kafka.common.errors.UnsupportedForMessageFormatException; import org.apache.kafka.common.errors.UnsupportedVersionException; import org.apache.kafka.common.internals.ClusterResourceListeners; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; import org.apache.kafka.common.message.ApiMessageType; import org.apache.kafka.common.message.EndTxnResponseData; import org.apache.kafka.common.message.InitProducerIdResponseData; @@ -1542,7 +1543,7 @@ public void testUnresolvedSequencesAreNotFatal() throws Exception { txnManager.beginTransaction(); txnManager.maybeAddPartition(tp0); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp0, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp0, Errors.NONE))); sender.runOnce(); // Send first ProduceRequest @@ -1828,7 +1829,7 @@ public void testTransactionalUnknownProducerHandlingWhenRetentionLimitReached() transactionManager.beginTransaction(); transactionManager.maybeAddPartition(tp0); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp0, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp0, Errors.NONE))); sender.runOnce(); // Receive AddPartitions response assertEquals(0, transactionManager.sequenceNumber(tp0).longValue()); @@ -2384,7 +2385,7 @@ public void testTransactionalSplitBatchAndSend() throws Exception { txnManager.beginTransaction(); txnManager.maybeAddPartition(tp); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp, Errors.NONE))); sender.runOnce(); testSplitBatchAndSend(txnManager, producerIdAndEpoch, tp); @@ -2731,7 +2732,7 @@ public void testTransactionalRequestsSentOnShutdown() { txnManager.beginTransaction(); txnManager.maybeAddPartition(tp); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp, Errors.NONE))); sender.runOnce(); sender.initiateClose(); txnManager.beginCommit(); @@ -2851,7 +2852,7 @@ public void testAwaitPendingRecordsBeforeCommittingTransaction() throws Exceptio private void addPartitionToTxn(Sender sender, TransactionManager txnManager, TopicPartition tp) { txnManager.maybeAddPartition(tp); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp, Errors.NONE))); runUntil(sender, () -> txnManager.isPartitionAdded(tp)); assertFalse(txnManager.hasInFlightRequest()); } @@ -2892,7 +2893,7 @@ public void testIncompleteTransactionAbortOnShutdown() { txnManager.beginTransaction(); txnManager.maybeAddPartition(tp); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp, Errors.NONE))); sender.runOnce(); sender.initiateClose(); AssertEndTxnRequestMatcher endTxnMatcher = new AssertEndTxnRequestMatcher(TransactionResult.ABORT); @@ -2926,7 +2927,7 @@ public void testForceShutdownWithIncompleteTransaction() { txnManager.beginTransaction(); txnManager.maybeAddPartition(tp); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp, Errors.NONE))); sender.runOnce(); // Try to commit the transaction but it won't happen as we'll forcefully close the sender @@ -2951,7 +2952,7 @@ public void testTransactionAbortedExceptionOnAbortWithoutError() throws Interrup // Begin the transaction txnManager.beginTransaction(); txnManager.maybeAddPartition(tp0); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp0, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp0, Errors.NONE))); // Run it once so that the partition is added to the transaction. sender.runOnce(); // Append a record to the accumulator. @@ -2989,7 +2990,7 @@ public void testTooLargeBatchesAreSafelyRemoved() throws InterruptedException { txnManager.beginTransaction(); txnManager.maybeAddPartition(tp0); - client.prepareResponse(new AddPartitionsToTxnResponse(0, Collections.singletonMap(tp0, Errors.NONE))); + client.prepareResponse(buildAddPartitionsToTxnResponseData(0, Collections.singletonMap(tp0, Errors.NONE))); sender.runOnce(); // create a producer batch with more than one record so it is eligible for splitting @@ -3338,4 +3339,11 @@ private void waitForProducerId(TransactionManager transactionManager, ProducerId assertTrue(transactionManager.hasProducerId()); assertEquals(producerIdAndEpoch, transactionManager.producerIdAndEpoch()); } + + private AddPartitionsToTxnResponse buildAddPartitionsToTxnResponseData(int throttleMs, Map errors) { + AddPartitionsToTxnResponseData.AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction( + AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, errors); + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(throttleMs); + return new AddPartitionsToTxnResponse(data); + } } diff --git a/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java b/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java index 8d0349a8f5aea..8c0f5c51d462f 100644 --- a/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java +++ b/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java @@ -40,6 +40,8 @@ import org.apache.kafka.common.header.Header; import org.apache.kafka.common.internals.ClusterResourceListeners; import org.apache.kafka.common.message.AddOffsetsToTxnResponseData; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult; import org.apache.kafka.common.message.ApiVersionsResponseData.ApiVersion; import org.apache.kafka.common.message.EndTxnResponseData; import org.apache.kafka.common.message.InitProducerIdResponseData; @@ -1303,11 +1305,13 @@ public void testCommitWithTopicAuthorizationFailureInAddPartitionsInFlight() thr Map errors = new HashMap<>(); errors.put(tp0, Errors.TOPIC_AUTHORIZATION_FAILED); errors.put(tp1, Errors.OPERATION_NOT_ATTEMPTED); + AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, errors); + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(0); client.respond(body -> { AddPartitionsToTxnRequest request = (AddPartitionsToTxnRequest) body; - assertEquals(new HashSet<>(request.partitions()), new HashSet<>(errors.keySet())); + assertEquals(new HashSet<>(AddPartitionsToTxnRequest.getPartitions(request.data().v3AndBelowTopics())), new HashSet<>(errors.keySet())); return true; - }, new AddPartitionsToTxnResponse(0, errors)); + }, new AddPartitionsToTxnResponse(data)); sender.runOnce(); assertTrue(transactionManager.hasError()); @@ -3439,11 +3443,13 @@ private void verifyCommitOrAbortTransactionRetriable(TransactionResult firstTran } private void prepareAddPartitionsToTxn(final Map errors) { + AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, errors); + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(0); client.prepareResponse(body -> { AddPartitionsToTxnRequest request = (AddPartitionsToTxnRequest) body; - assertEquals(new HashSet<>(request.partitions()), new HashSet<>(errors.keySet())); + assertEquals(new HashSet<>(AddPartitionsToTxnRequest.getPartitions(request.data().v3AndBelowTopics())), new HashSet<>(errors.keySet())); return true; - }, new AddPartitionsToTxnResponse(0, errors)); + }, new AddPartitionsToTxnResponse(data)); } private void prepareAddPartitionsToTxn(final TopicPartition tp, final Errors error) { @@ -3522,14 +3528,22 @@ private MockClient.RequestMatcher produceRequestMatcher(final long producerId, f private void prepareAddPartitionsToTxnResponse(Errors error, final TopicPartition topicPartition, final short epoch, final long producerId) { + AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction( + AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, singletonMap(topicPartition, error)); client.prepareResponse(addPartitionsRequestMatcher(topicPartition, epoch, producerId), - new AddPartitionsToTxnResponse(0, singletonMap(topicPartition, error))); + new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData() + .setThrottleTimeMs(0) + .setResultsByTopicV3AndBelow(result.topicResults()))); } private void sendAddPartitionsToTxnResponse(Errors error, final TopicPartition topicPartition, final short epoch, final long producerId) { + AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction( + AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, singletonMap(topicPartition, error)); client.respond(addPartitionsRequestMatcher(topicPartition, epoch, producerId), - new AddPartitionsToTxnResponse(0, singletonMap(topicPartition, error))); + new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData() + .setThrottleTimeMs(0) + .setResultsByTopicV3AndBelow(result.topicResults()))); } private MockClient.RequestMatcher addPartitionsRequestMatcher(final TopicPartition topicPartition, @@ -3538,7 +3552,7 @@ private MockClient.RequestMatcher addPartitionsRequestMatcher(final TopicPartiti AddPartitionsToTxnRequest addPartitionsToTxnRequest = (AddPartitionsToTxnRequest) body; assertEquals(producerId, addPartitionsToTxnRequest.data().v3AndBelowProducerId()); assertEquals(epoch, addPartitionsToTxnRequest.data().v3AndBelowProducerEpoch()); - assertEquals(singletonList(topicPartition), addPartitionsToTxnRequest.partitions()); + assertEquals(singletonList(topicPartition), AddPartitionsToTxnRequest.getPartitions(addPartitionsToTxnRequest.data().v3AndBelowTopics())); assertEquals(transactionalId, addPartitionsToTxnRequest.data().v3AndBelowTransactionalId()); return true; }; diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java index 59ec7ff9bb795..0d7c299401786 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java @@ -63,7 +63,7 @@ public void testConstructor(short version) { assertEquals(transactionalId1, request.data().v3AndBelowTransactionalId()); assertEquals(producerId, request.data().v3AndBelowProducerId()); assertEquals(producerEpoch, request.data().v3AndBelowProducerEpoch()); - assertEquals(partitions, request.partitions()); + assertEquals(partitions, AddPartitionsToTxnRequest.getPartitions(request.data().v3AndBelowTopics())); } else { AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); @@ -80,12 +80,15 @@ public void testConstructor(short version) { assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); assertEquals(throttleTimeMs, response.throttleTimeMs()); + + if (version >= 4) { + assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION.code(), response.data().errorCode()); + } } @Test public void testBatchedRequests() { AddPartitionsToTxnTransactionCollection transactions = createTwoTransactionCollection(); - boolean verifyOnly = true; AddPartitionsToTxnRequest.Builder builder = AddPartitionsToTxnRequest.Builder.forBroker(transactions); AddPartitionsToTxnRequest request = builder.build(ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion()); @@ -105,8 +108,8 @@ public void testBatchedRequests() { .setResultsByTransaction(results) .setThrottleTimeMs(throttleTimeMs)); - assertEquals(Collections.singletonMap(tp0, Errors.UNKNOWN_TOPIC_OR_PARTITION), response.errorsPerTransaction(transactionalId1)); - assertEquals(Collections.singletonMap(tp1, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED), response.errorsPerTransaction(transactionalId2)); + assertEquals(Collections.singletonMap(tp0, Errors.UNKNOWN_TOPIC_OR_PARTITION), response.errorsForTransaction(response.getTransactionTopicResults(transactionalId1))); + assertEquals(Collections.singletonMap(tp1, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED), response.errorsForTransaction(response.getTransactionTopicResults(transactionalId2))); } @Test diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index 3188a80c6b74e..044bd4b4884f0 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -25,8 +25,10 @@ import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResultCollection; import org.apache.kafka.common.protocol.ApiKeys; import org.apache.kafka.common.protocol.Errors; +import org.apache.kafka.common.utils.annotation.ApiKeyVersionsSource; import org.junit.jupiter.api.BeforeEach; import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; import java.util.Collections; import java.util.HashMap; @@ -62,67 +64,56 @@ public void setUp() { errorsMap.put(tp2, errorTwo); } - @Test - public void testConstructorWithErrorResponse() { - // This test only applies to versions 0-3. - AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(throttleTimeMs, errorsMap); - - assertEquals(expectedErrorCounts, response.errorCounts()); - assertEquals(throttleTimeMs, response.throttleTimeMs()); - } - - @Test - public void testParse() { + @ParameterizedTest + @ApiKeyVersionsSource(apiKey = ApiKeys.ADD_PARTITIONS_TO_TXN) + public void testParse(short version) { AddPartitionsToTxnTopicResultCollection topicCollection = new AddPartitionsToTxnTopicResultCollection(); AddPartitionsToTxnTopicResult topicResult = new AddPartitionsToTxnTopicResult(); topicResult.setName(topicOne); - topicResult.results().add(new AddPartitionsToTxnPartitionResult() + topicResult.resultsByPartition().add(new AddPartitionsToTxnPartitionResult() .setPartitionErrorCode(errorOne.code()) .setPartitionIndex(partitionOne)); - topicResult.results().add(new AddPartitionsToTxnPartitionResult() + topicResult.resultsByPartition().add(new AddPartitionsToTxnPartitionResult() .setPartitionErrorCode(errorTwo.code()) .setPartitionIndex(partitionTwo)); topicCollection.add(topicResult); - - for (short version : ApiKeys.ADD_PARTITIONS_TO_TXN.allVersions()) { - if (version < 4) { - AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() - .setResultsByTopicV3AndBelow(topicCollection) - .setThrottleTimeMs(throttleTimeMs); - AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); - - AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); - assertEquals(expectedErrorCounts, parsedResponse.errorCounts()); - assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); - assertEquals(version >= 1, parsedResponse.shouldClientThrottle(version)); - } else { - AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); - results.add(new AddPartitionsToTxnResult().setTransactionalId("txn1").setTopicResults(topicCollection)); - - // Create another transaction with new name and errorOne for a single partition. - Map txnTwoExpectedErrors = Collections.singletonMap(tp2, errorOne); - results.add(AddPartitionsToTxnResponse.resultForTransaction("txn2", txnTwoExpectedErrors)); - - AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() - .setResultsByTransaction(results) - .setThrottleTimeMs(throttleTimeMs); - AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); - - Map newExpectedErrorCounts = new HashMap<>(); - newExpectedErrorCounts.put(errorOne, 2); - newExpectedErrorCounts.put(errorTwo, 1); - - AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); - assertEquals(txnTwoExpectedErrors, parsedResponse.errorsPerTransaction("txn2")); - assertEquals(newExpectedErrorCounts, parsedResponse.errorCounts()); - assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); - assertTrue(parsedResponse.shouldClientThrottle(version)); - } + if (version < 4) { + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() + .setResultsByTopicV3AndBelow(topicCollection) + .setThrottleTimeMs(throttleTimeMs); + AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); + + AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); + assertEquals(expectedErrorCounts, parsedResponse.errorCounts()); + assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); + assertEquals(version >= 1, parsedResponse.shouldClientThrottle(version)); + } else { + AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); + results.add(new AddPartitionsToTxnResult().setTransactionalId("txn1").setTopicResults(topicCollection)); + + // Create another transaction with new name and errorOne for a single partition. + Map txnTwoExpectedErrors = Collections.singletonMap(tp2, errorOne); + results.add(AddPartitionsToTxnResponse.resultForTransaction("txn2", txnTwoExpectedErrors)); + + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData() + .setResultsByTransaction(results) + .setThrottleTimeMs(throttleTimeMs); + AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); + + Map newExpectedErrorCounts = new HashMap<>(); + newExpectedErrorCounts.put(errorOne, 2); + newExpectedErrorCounts.put(errorTwo, 1); + + AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); + assertEquals(txnTwoExpectedErrors, parsedResponse.errorsForTransaction(response.getTransactionTopicResults("txn2"))); + assertEquals(newExpectedErrorCounts, parsedResponse.errorCounts()); + assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); + assertTrue(parsedResponse.shouldClientThrottle(version)); } } @@ -140,7 +131,7 @@ public void testBatchedErrors() { AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setResultsByTransaction(results)); - assertEquals(txn1Errors, response.errorsPerTransaction("txn1")); - assertEquals(txn2Errors, response.errorsPerTransaction("txn2")); + assertEquals(txn1Errors, response.errorsForTransaction(response.getTransactionTopicResults("txn1"))); + assertEquals(txn2Errors, response.errorsForTransaction(response.getTransactionTopicResults("txn2"))); } } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java index a87a5d72a1e3b..e7c465c6bc3a2 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java @@ -35,6 +35,7 @@ import org.apache.kafka.common.errors.SecurityDisabledException; import org.apache.kafka.common.errors.UnknownServerException; import org.apache.kafka.common.errors.UnsupportedVersionException; +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData; import org.apache.kafka.common.message.AddOffsetsToTxnRequestData; import org.apache.kafka.common.message.AddOffsetsToTxnResponseData; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopic; @@ -2576,7 +2577,10 @@ private AddPartitionsToTxnRequest createAddPartitionsToTxnRequest(short version) } private AddPartitionsToTxnResponse createAddPartitionsToTxnResponse() { - return new AddPartitionsToTxnResponse(0, Collections.singletonMap(new TopicPartition("t", 0), Errors.NONE)); + AddPartitionsToTxnResponseData.AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction( + AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, Collections.singletonMap(new TopicPartition("t", 0), Errors.NONE)); + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(0); + return new AddPartitionsToTxnResponse(data); } private AddOffsetsToTxnRequest createAddOffsetsToTxnRequest(short version) { diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 25878f8c3bedf..3cb4388c1d55b 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2386,8 +2386,11 @@ class KafkaApis(val requestChannel: RequestChannel, } def handleAddPartitionsToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { ensureInterBrokerVersion(IBP_0_11_0_IV0) - val lock = new Object - val addPartitionsToTxnRequest = if (request.context.apiVersion() < 4) request.body[AddPartitionsToTxnRequest].normalizeRequest() else request.body[AddPartitionsToTxnRequest] + val addPartitionsToTxnRequest = + if (request.context.apiVersion() < 4) + request.body[AddPartitionsToTxnRequest].normalizeRequest() + else + request.body[AddPartitionsToTxnRequest] val version = addPartitionsToTxnRequest.version val responses = new AddPartitionsToTxnResultCollection() val partitionsByTransaction = addPartitionsToTxnRequest.partitionsByTransaction() @@ -2413,15 +2416,19 @@ class KafkaApis(val requestChannel: RequestChannel, val txns = addPartitionsToTxnRequest.data.transactions def maybeSendResponse(): Unit = { - lock synchronized { + var canSend = false + responses.synchronized { if (responses.size() == txns.size()) { - requestHelper.sendResponseMaybeThrottle(request, createResponse) + canSend = true } } + if (canSend) { + requestHelper.sendResponseMaybeThrottle(request, createResponse) + } } txns.forEach( transaction => { - val transactionalId = transaction.transactionalId() + val transactionalId = transaction.transactionalId val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala // Versions < 4 come from clients and must be authorized to write for the given transaction and for the given topics. @@ -2463,7 +2470,7 @@ class KafkaApis(val requestChannel: RequestChannel, error } } - lock synchronized { + responses.synchronized { responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) } maybeSendResponse() diff --git a/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala b/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala index fca36644a3154..e8f2ea88c491d 100644 --- a/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala +++ b/core/src/test/scala/integration/kafka/api/AuthorizerIntegrationTest.scala @@ -231,7 +231,7 @@ class AuthorizerIntegrationTest extends BaseRequestTest { resp.errors.get(new ConfigResource(ConfigResource.Type.TOPIC, tp.topic)).error), ApiKeys.INIT_PRODUCER_ID -> ((resp: InitProducerIdResponse) => resp.error), ApiKeys.WRITE_TXN_MARKERS -> ((resp: WriteTxnMarkersResponse) => resp.errorsByProducerId.get(producerId).get(tp)), - ApiKeys.ADD_PARTITIONS_TO_TXN -> ((resp: AddPartitionsToTxnResponse) => resp.errors.get(tp)), + ApiKeys.ADD_PARTITIONS_TO_TXN -> ((resp: AddPartitionsToTxnResponse) => resp.errors.get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID).get(tp)), ApiKeys.ADD_OFFSETS_TO_TXN -> ((resp: AddOffsetsToTxnResponse) => Errors.forCode(resp.data.errorCode)), ApiKeys.END_TXN -> ((resp: EndTxnResponse) => resp.error), ApiKeys.TXN_OFFSET_COMMIT -> ((resp: TxnOffsetCommitResponse) => resp.errors.get(tp)), diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index 0050c84989b22..32c489342c56c 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -302,7 +302,7 @@ class TransactionCoordinatorTest { any() ) } - + @Test def shouldRespondWithErrorsNoneOnAddPartitionWhenNoErrorsAndPartitionsTheSame(): Unit = { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index c04bc2b4a656f..9eeb04b3a6639 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -93,7 +93,11 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val leaderId = brokers.head.config.brokerId val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(leaderId)) - val errors = if (version < 4) response.errors else response.errorsPerTransaction(transactionalId) + val errors = + if (version < 4) + response.errors.get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID) + else + response.errorsForTransaction(response.getTransactionTopicResults(transactionalId)) assertEquals(2, errors.size) @@ -155,7 +159,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(coordinatorId)) - val errors = response.allErrors() + val errors = response.errors() assertTrue(errors.containsKey(transactionalId1)) assertTrue(errors.get(transactionalId1).containsKey(tp0)) From 2c7d25ca272345dc141b9b1f07185678823c243c Mon Sep 17 00:00:00 2001 From: Justine Date: Wed, 22 Feb 2023 16:36:55 -0800 Subject: [PATCH 10/17] Per Partition error codes for verify only --- .../requests/AddPartitionsToTxnResponse.java | 2 +- .../message/AddPartitionsToTxnResponse.json | 2 +- .../transaction/TransactionCoordinator.scala | 27 ++++-- .../main/scala/kafka/server/KafkaApis.scala | 9 ++ ...ransactionCoordinatorConcurrencyTest.scala | 3 + .../TransactionCoordinatorTest.scala | 30 ++++--- .../AddPartitionsToTxnRequestServerTest.scala | 87 ++++++++++++------- .../unit/kafka/server/KafkaApisTest.scala | 2 + 8 files changed, 113 insertions(+), 49 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index f305cb805bbee..804fdb0d7a52a 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -50,7 +50,7 @@ public class AddPartitionsToTxnResponse extends AbstractResponse { private final AddPartitionsToTxnResponseData data; - + public static final String V3_AND_BELOW_TXN_ID = ""; public AddPartitionsToTxnResponse(AddPartitionsToTxnResponseData data) { diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json index 54a4a92614b44..f96b42219a18e 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json @@ -23,7 +23,7 @@ // // Version 3 enables flexible versions. // - // Version 4 adds support to batch multiple transactions. + // Version 4 adds support to batch multiple transactions and a top level error code. "validVersions": "0-4", "flexibleVersions": "3+", "fields": [ diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index 31b64d7b2fdaf..3a2891a9ad029 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -30,6 +30,8 @@ import org.apache.kafka.common.requests.TransactionResult import org.apache.kafka.common.utils.{LogContext, ProducerIdAndEpoch, Time} import org.apache.kafka.server.util.Scheduler +import scala.collection.mutable + object TransactionCoordinator { def apply(config: KafkaConfig, @@ -92,6 +94,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, type InitProducerIdCallback = InitProducerIdResult => Unit type AddPartitionsCallback = Errors => Unit + type VerifyPartitionsCallback = Map[TopicPartition, Errors] => Unit type EndTxnCallback = Errors => Unit type ApiResult[T] = Either[Errors, T] @@ -318,12 +321,24 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } } + def verifyAddPartitionsToTransaction(partitions: collection.Set[TopicPartition], + txnMetadataPartitions: Set[TopicPartition], + responseCallback: VerifyPartitionsCallback): Unit = { + val addedPartitions = partitions.intersect(txnMetadataPartitions) + val nonAddedPartitions = partitions.diff(txnMetadataPartitions) + val errors = mutable.Map[TopicPartition, Errors]() + addedPartitions.foreach(errors.put(_, Errors.NONE)) + nonAddedPartitions.foreach(errors.put(_, Errors.INVALID_TXN_STATE)) + responseCallback(errors.toMap) + } + def handleAddPartitionsToTransaction(transactionalId: String, producerId: Long, producerEpoch: Short, partitions: collection.Set[TopicPartition], verifyOnly: Boolean, responseCallback: AddPartitionsCallback, + verifyCallback: VerifyPartitionsCallback, requestLocal: RequestLocal = RequestLocal.NoCaching): Unit = { if (transactionalId == null || transactionalId.isEmpty) { debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") @@ -355,10 +370,9 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } else { // If verifyOnly, we should have returned in the step above. If we didn't the partitions are not present in the transaction. if (verifyOnly) { - Left(Errors.INVALID_TXN_STATE) - } else { - Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds())) - } + verifyAddPartitionsToTransaction(partitions, txnMetadata.topicPartitions.toSet, verifyCallback) + } + Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds())) } } } @@ -368,9 +382,12 @@ class TransactionCoordinator(txnConfig: TransactionConfig, debug(s"Returning $err error code to client for $transactionalId's AddPartitions request") responseCallback(err) - case Right((coordinatorEpoch, newMetadata)) => + case Right((coordinatorEpoch, newMetadata)) if !verifyOnly => txnManager.appendTransactionToLog(transactionalId, coordinatorEpoch, newMetadata, responseCallback, requestLocal = requestLocal) + + case _ => + // We only hit this case if verifyOnly and some partitions were not present. If so, we already handled the response. } } } diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 3cb4388c1d55b..c1427190d6c58 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2475,6 +2475,13 @@ class KafkaApis(val requestChannel: RequestChannel, } maybeSendResponse() } + + def sendVerifyResponseCallback(errors: Map[TopicPartition, Errors]): Unit = { + responses.synchronized { + responses.add(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, errors.asJava)) + } + maybeSendResponse() + } txnCoordinator.handleAddPartitionsToTransaction(transactionalId, transaction.producerId, @@ -2482,6 +2489,7 @@ class KafkaApis(val requestChannel: RequestChannel, authorizedPartitions, transaction.verifyOnly, sendResponseCallback, + sendVerifyResponseCallback, requestLocal) } } @@ -2535,6 +2543,7 @@ class KafkaApis(val requestChannel: RequestChannel, Set(offsetTopicPartition), false, sendResponseCallback, + null, requestLocal) } } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala index e1c910d0b5cbe..07ad6baf96b50 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala @@ -498,7 +498,9 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren abstract class TxnOperation[R] extends Operation { @volatile var result: Option[R] = None + @volatile var results: Map[TopicPartition, R] = _ def resultCallback(r: R): Unit = this.result = Some(r) + def resultPerPartitionCallback(r: Map[TopicPartition, R]): Unit = this.results = r } class InitProducerIdOperation(val producerIdAndEpoch: Option[ProducerIdAndEpoch] = None) extends TxnOperation[InitProducerIdResult] { @@ -523,6 +525,7 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren partitions, false, resultCallback, + resultPerPartitionCallback, RequestLocal.withThreadConfinedCaching) replicaManager.tryCompleteActions() } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index 32c489342c56c..389fb0a28795d 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -63,6 +63,7 @@ class TransactionCoordinatorTest { val transactionStatePartitionCount = 1 var result: InitProducerIdResult = _ var error: Errors = Errors.NONE + var errors: Map[TopicPartition, Errors] = _ private def mockPidGenerator(): Unit = { when(pidGenerator.generateProducerId()).thenAnswer(_ => { @@ -200,19 +201,19 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(None)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsEmpty(): Unit = { - coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.INVALID_REQUEST, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsNull(): Unit = { - coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.INVALID_REQUEST, error) } @@ -221,7 +222,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.NOT_COORDINATOR)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.NOT_COORDINATOR, error) } @@ -230,7 +231,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.COORDINATOR_LOAD_IN_PROGRESS)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.COORDINATOR_LOAD_IN_PROGRESS, error) } @@ -249,7 +250,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, state, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.CONCURRENT_TRANSACTIONS, error) } @@ -259,7 +260,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 10, 9, 0, PrepareCommit, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.PRODUCER_FENCED, error) } @@ -290,7 +291,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, txnMetadata)))) - coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, false, errorsCallback, errorsPerPartitionCallback) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) verify(transactionManager).appendTransactionToLog( @@ -309,7 +310,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -320,7 +321,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Ongoing, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, true, errorsCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, true, errorsCallback, errorsPerPartitionCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -334,8 +335,9 @@ class TransactionCoordinatorTest { val extraPartitions = partitions ++ Set(new TopicPartition("topic2", 0)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, extraPartitions, true, errorsCallback) - assertEquals(Errors.INVALID_TXN_STATE, error) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, extraPartitions, true, errorsCallback, errorsPerPartitionCallback) + assertEquals(Errors.INVALID_TXN_STATE, errors(new TopicPartition("topic2", 0))) + assertEquals(Errors.NONE, errors(new TopicPartition("topic1", 0))) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -1210,4 +1212,8 @@ class TransactionCoordinatorTest { def errorsCallback(ret: Errors): Unit = { error = ret } + + def errorsPerPartitionCallback(ret: Map[TopicPartition, Errors]): Unit = { + errors = ret + } } diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index 9eeb04b3a6639..47f37317b8733 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -110,55 +110,32 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { @Test def testOneSuccessOneErrorInBatchedRequest(): Unit = { + val tp0 = new TopicPartition(topic1, 0) val transactionalId1 = "foobar" - - val findCoordinatorRequest = new FindCoordinatorRequest.Builder(new FindCoordinatorRequestData().setKey(transactionalId1).setKeyType(CoordinatorType.TRANSACTION.id)).build() - // First find coordinator request creates the state topic, then wait for transactional topics to be created. - connectAndReceive[FindCoordinatorResponse](findCoordinatorRequest, brokerSocketServer(brokers.head.config.brokerId)) - TestUtils.waitForAllPartitionsMetadata(brokers, "__transaction_state", 50) - val findCoordinatorResponse = connectAndReceive[FindCoordinatorResponse](findCoordinatorRequest, brokerSocketServer(brokers.head.config.brokerId)) - val coordinatorId = findCoordinatorResponse.data().coordinators().get(0).nodeId() - - val initPidRequest = new InitProducerIdRequest.Builder(new InitProducerIdRequestData().setTransactionalId(transactionalId1).setTransactionTimeoutMs(10000)).build() - val initPidResponse = connectAndReceive[InitProducerIdResponse](initPidRequest, brokerSocketServer(coordinatorId)) - - val producerId1 = initPidResponse.data().producerId() - val producerEpoch1 = initPidResponse.data().producerEpoch() - val transactionalId2 = "barfoo" // "barfoo" maps to the same transaction coordinator val producerId2 = 1000L val producerEpoch2: Short = 0 - val tp0 = new TopicPartition(topic1, 0) - - val txn1Topics = new AddPartitionsToTxnTopicCollection() - txn1Topics.add(new AddPartitionsToTxnTopic() - .setName(tp0.topic()) - .setPartitions(Collections.singletonList(tp0.partition()))) - - val txn2Topics = new AddPartitionsToTxnTopicCollection() + val txn2Topics = new AddPartitionsToTxnTopicCollection() txn2Topics.add(new AddPartitionsToTxnTopic() .setName(tp0.topic()) .setPartitions(Collections.singletonList(tp0.partition()))) + val (coordinatorId, txn1) = setUpTransactions(transactionalId1, false, Set(tp0)) + val transactions = new AddPartitionsToTxnTransactionCollection() - transactions.add(new AddPartitionsToTxnTransaction() - .setTransactionalId(transactionalId1) - .setProducerId(producerId1) - .setProducerEpoch(producerEpoch1) - .setVerifyOnly(false) - .setTopics(txn1Topics)) + transactions.add(txn1) transactions.add(new AddPartitionsToTxnTransaction() .setTransactionalId(transactionalId2) .setProducerId(producerId2) .setProducerEpoch(producerEpoch2) .setVerifyOnly(false) .setTopics(txn2Topics)) - + val request = AddPartitionsToTxnRequest.Builder.forBroker(transactions).build() val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(coordinatorId)) - + val errors = response.errors() assertTrue(errors.containsKey(transactionalId1)) @@ -169,6 +146,56 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { assertTrue(errors.get(transactionalId1).containsKey(tp0)) assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, errors.get(transactionalId2).get(tp0)) } + + @Test + def testVerifyOnly(): Unit = { + val tp0 = new TopicPartition(topic1, 0) + + val transactionalId = "foobar" + val (coordinatorId, txn) = setUpTransactions(transactionalId, true, Set(tp0)) + + val transactions = new AddPartitionsToTxnTransactionCollection() + transactions.add(txn) + + val verifyRequest = AddPartitionsToTxnRequest.Builder.forBroker(transactions).build() + + val verifyResponse = connectAndReceive[AddPartitionsToTxnResponse](verifyRequest, brokerSocketServer(coordinatorId)) + + val verifyErrors = verifyResponse.errors() + + assertTrue(verifyErrors.containsKey(transactionalId)) + assertTrue(verifyErrors.get(transactionalId).containsKey(tp0)) + assertEquals(Errors.INVALID_TXN_STATE, verifyErrors.get(transactionalId).get(tp0)) + } + + private def setUpTransactions(transactionalId: String, verifyOnly: Boolean, partitions: Set[TopicPartition]): (Int, AddPartitionsToTxnTransaction) = { + val findCoordinatorRequest = new FindCoordinatorRequest.Builder(new FindCoordinatorRequestData().setKey(transactionalId).setKeyType(CoordinatorType.TRANSACTION.id)).build() + // First find coordinator request creates the state topic, then wait for transactional topics to be created. + connectAndReceive[FindCoordinatorResponse](findCoordinatorRequest, brokerSocketServer(brokers.head.config.brokerId)) + TestUtils.waitForAllPartitionsMetadata(brokers, "__transaction_state", 50) + val findCoordinatorResponse = connectAndReceive[FindCoordinatorResponse](findCoordinatorRequest, brokerSocketServer(brokers.head.config.brokerId)) + val coordinatorId = findCoordinatorResponse.data().coordinators().get(0).nodeId() + + val initPidRequest = new InitProducerIdRequest.Builder(new InitProducerIdRequestData().setTransactionalId(transactionalId).setTransactionTimeoutMs(10000)).build() + val initPidResponse = connectAndReceive[InitProducerIdResponse](initPidRequest, brokerSocketServer(coordinatorId)) + + val producerId1 = initPidResponse.data().producerId() + val producerEpoch1 = initPidResponse.data().producerEpoch() + + val txn1Topics = new AddPartitionsToTxnTopicCollection() + partitions.foreach(tp => + txn1Topics.add(new AddPartitionsToTxnTopic() + .setName(tp.topic()) + .setPartitions(Collections.singletonList(tp.partition()))) + ) + + (coordinatorId, new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId) + .setProducerId(producerId1) + .setProducerEpoch(producerEpoch1) + .setVerifyOnly(verifyOnly) + .setTopics(txn1Topics)) + } } object AddPartitionsToTxnRequestServerTest { diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index 43559705c488d..2726d9900c7a0 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -1964,6 +1964,7 @@ class KafkaApisTest { ArgumentMatchers.eq(Set(new TopicPartition(Topic.GROUP_METADATA_TOPIC_NAME, partition))), ArgumentMatchers.eq(false), responseCallback.capture(), + ArgumentMatchers.any(), ArgumentMatchers.eq(requestLocal) )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) @@ -2019,6 +2020,7 @@ class KafkaApisTest { ArgumentMatchers.eq(Set(topicPartition)), ArgumentMatchers.eq(false), responseCallback.capture(), + ArgumentMatchers.any(), ArgumentMatchers.eq(requestLocal) )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) From 19dbfe60c458fb6b2f29d1d765af46ec2aebae6e Mon Sep 17 00:00:00 2001 From: Justine Date: Thu, 23 Feb 2023 14:15:32 -0800 Subject: [PATCH 11/17] unstable apis support --- .../resources/common/message/AddPartitionsToTxnRequest.json | 4 ++++ core/src/test/scala/unit/kafka/utils/TestUtils.scala | 2 ++ 2 files changed, 6 insertions(+) diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json index fbd34f73ed80c..a414dc53621ca 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json @@ -25,6 +25,10 @@ // Version 3 enables flexible versions. // // Version 4 adds VerifyOnly field to check if partitions are already in transaction and adds support to batch multiple transactions. + // The AddPartitionsToTxnRequest version 4 API is added as part of KIP-890 and is still + // under developement. Hence, the API is not exposed by default by brokers + // unless explicitely enabled. + "latestVersionUnstable": true, "validVersions": "0-4", "flexibleVersions": "3+", "fields": [ diff --git a/core/src/test/scala/unit/kafka/utils/TestUtils.scala b/core/src/test/scala/unit/kafka/utils/TestUtils.scala index 2c9c4ae6690b9..0af0e35b5a8c4 100755 --- a/core/src/test/scala/unit/kafka/utils/TestUtils.scala +++ b/core/src/test/scala/unit/kafka/utils/TestUtils.scala @@ -374,6 +374,8 @@ object TestUtils extends Logging { props.put(KafkaConfig.RackProp, nodeId.toString) props.put(KafkaConfig.ReplicaSelectorClassProp, "org.apache.kafka.common.replica.RackAwareReplicaSelector") } + + props.put(KafkaConfig.UnstableApiVersionsEnableProp, "true") props } From 42486de58f2a0083ccaae1f96556318dceefe77e Mon Sep 17 00:00:00 2001 From: Justine Date: Thu, 23 Feb 2023 15:30:34 -0800 Subject: [PATCH 12/17] Fix tests --- .../kafka/common/message/MessageTest.java | 23 +++++++++++-- .../common/requests/RequestResponseTest.java | 34 ++++++++++++------- .../unit/kafka/server/KafkaApisTest.scala | 8 ++--- 3 files changed, 47 insertions(+), 18 deletions(-) diff --git a/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java b/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java index 1f6ad3ed1133a..5762e84f60bb9 100644 --- a/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java +++ b/clients/src/test/java/org/apache/kafka/common/message/MessageTest.java @@ -24,6 +24,7 @@ import org.apache.kafka.common.errors.UnsupportedVersionException; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopic; import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTopicCollection; +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.AddPartitionsToTxnTransactionCollection; import org.apache.kafka.common.message.DescribeClusterResponseData.DescribeClusterBroker; import org.apache.kafka.common.message.DescribeClusterResponseData.DescribeClusterBrokerCollection; import org.apache.kafka.common.message.DescribeGroupsResponseData.DescribedGroup; @@ -96,14 +97,26 @@ public void testAddOffsetsToTxnVersions() throws Exception { @Test public void testAddPartitionsToTxnVersions() throws Exception { - testAllMessageRoundTrips(new AddPartitionsToTxnRequestData(). + AddPartitionsToTxnRequestData v3AndBelowData = new AddPartitionsToTxnRequestData(). setV3AndBelowTransactionalId("blah"). setV3AndBelowProducerId(0xbadcafebadcafeL). setV3AndBelowProducerEpoch((short) 30000). setV3AndBelowTopics(new AddPartitionsToTxnTopicCollection(singletonList( new AddPartitionsToTxnTopic(). setName("Topic"). - setPartitions(singletonList(1))).iterator()))); + setPartitions(singletonList(1))).iterator())); + testDuplication(v3AndBelowData); + testAllMessageRoundTripsUntilVersion((short) 3, v3AndBelowData); + + AddPartitionsToTxnRequestData data = new AddPartitionsToTxnRequestData(). + setTransactions(new AddPartitionsToTxnTransactionCollection(singletonList( + new AddPartitionsToTxnRequestData.AddPartitionsToTxnTransaction(). + setTransactionalId("blah"). + setProducerId(0xbadcafebadcafeL). + setProducerEpoch((short) 30000). + setTopics(v3AndBelowData.v3AndBelowTopics())).iterator())); + testDuplication(data); + testAllMessageRoundTripsFromVersion((short) 4, data); } @Test @@ -1032,6 +1045,12 @@ private void testAllMessageRoundTripsFromVersion(short fromVersion, Message mess } } + private void testAllMessageRoundTripsUntilVersion(short untilVersion, Message message) throws Exception { + for (short version = message.lowestSupportedVersion(); version <= untilVersion; version++) { + testEquivalentMessageRoundTrip(version, message); + } + } + private void testMessageRoundTrip(short version, Message message, Message expected) throws Exception { testByteBufferRoundTrip(version, message, expected); } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java index d1a07c63cdda9..587e111a46dd7 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java @@ -931,7 +931,8 @@ public void testDeletableTopicResultErrorMessageIsNullByDefault() { @Test public void testErrorCountsIncludesNone() { assertEquals(1, createAddOffsetsToTxnResponse().errorCounts().get(Errors.NONE)); - assertEquals(1, createAddPartitionsToTxnResponse().errorCounts().get(Errors.NONE)); + assertEquals(1, createAddPartitionsToTxnResponse((short) 3).errorCounts().get(Errors.NONE)); + assertEquals(1, createAddPartitionsToTxnResponse((short) 4).errorCounts().get(Errors.NONE)); assertEquals(1, createAlterClientQuotasResponse().errorCounts().get(Errors.NONE)); assertEquals(1, createAlterConfigsResponse().errorCounts().get(Errors.NONE)); assertEquals(2, createAlterPartitionReassignmentsResponse().errorCounts().get(Errors.NONE)); @@ -1085,7 +1086,7 @@ private AbstractResponse getResponse(ApiKeys apikey, short version) { case DELETE_RECORDS: return createDeleteRecordsResponse(); case INIT_PRODUCER_ID: return createInitPidResponse(); case OFFSET_FOR_LEADER_EPOCH: return createLeaderEpochResponse(); - case ADD_PARTITIONS_TO_TXN: return createAddPartitionsToTxnResponse(); + case ADD_PARTITIONS_TO_TXN: return createAddPartitionsToTxnResponse(version); case ADD_OFFSETS_TO_TXN: return createAddOffsetsToTxnResponse(); case END_TXN: return createEndTxnResponse(); case WRITE_TXN_MARKERS: return createWriteTxnMarkersResponse(); @@ -1616,7 +1617,7 @@ private void checkResponse(AbstractResponse response, short version) { serializedBytes.rewind(); assertEquals(serializedBytes, serializedBytes2, "Response " + response + "failed equality test"); } catch (Exception e) { - throw new RuntimeException("Failed to deserialize response " + response + " with type " + response.getClass(), e); + throw new RuntimeException("Failed to deserialize version " + version + " response " + response + " with type " + response.getClass(), e); } } @@ -2603,27 +2604,36 @@ private OffsetsForLeaderEpochResponse createLeaderEpochResponse() { } private AddPartitionsToTxnRequest createAddPartitionsToTxnRequest(short version) { - if (version < 3) { + if (version < 4) { return AddPartitionsToTxnRequest.Builder.forClient("tid", 21L, (short) 42, singletonList(new TopicPartition("topic", 73))).build(version); } else { - AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection(); - AddPartitionsToTxnTopicCollection topics = new AddPartitionsToTxnTopicCollection(); - topics.add(new AddPartitionsToTxnTopic().setName("topic").setPartitions(Collections.singletonList(73))); - transactions.add(new AddPartitionsToTxnTransaction() + AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection( + singletonList(new AddPartitionsToTxnTransaction() .setTransactionalId("tid") .setProducerId(21L) .setProducerEpoch((short) 42) .setVerifyOnly(false) - .setTopics(topics)); + .setTopics(new AddPartitionsToTxnTopicCollection( + singletonList(new AddPartitionsToTxnTopic() + .setName("topic") + .setPartitions(Collections.singletonList(73))).iterator()))) + .iterator()); return AddPartitionsToTxnRequest.Builder.forBroker(transactions).build(version); } } - private AddPartitionsToTxnResponse createAddPartitionsToTxnResponse() { + private AddPartitionsToTxnResponse createAddPartitionsToTxnResponse(short version) { + String txnId = version < 4 ? AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID : "tid"; AddPartitionsToTxnResponseData.AddPartitionsToTxnResult result = AddPartitionsToTxnResponse.resultForTransaction( - AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID, Collections.singletonMap(new TopicPartition("t", 0), Errors.NONE)); - AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(0); + txnId, Collections.singletonMap(new TopicPartition("t", 0), Errors.NONE)); + AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setThrottleTimeMs(0); + + if (version < 4) { + data.setResultsByTopicV3AndBelow(result.topicResults()); + } else { + data.setResultsByTransaction(new AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection(singletonList(result).iterator())); + } return new AddPartitionsToTxnResponse(data); } diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index bbea0fa385075..b962dbb6b0cdf 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -1990,7 +1990,7 @@ class KafkaApisTest { val topic = "topic" addTopicToMetadataCache(topic, numPartitions = 2) - for (version <- ApiKeys.ADD_PARTITIONS_TO_TXN.oldestVersion to ApiKeys.ADD_PARTITIONS_TO_TXN.latestVersion) { + for (version <- ApiKeys.ADD_PARTITIONS_TO_TXN.oldestVersion to 3) { reset(replicaManager, clientRequestQuotaManager, requestChannel, txnCoordinator) @@ -2034,9 +2034,9 @@ class KafkaApisTest { val response = capturedResponse.getValue if (version < 2) { - assertEquals(Collections.singletonMap(topicPartition, Errors.INVALID_PRODUCER_EPOCH), response.errors()) + assertEquals(Collections.singletonMap(topicPartition, Errors.INVALID_PRODUCER_EPOCH), response.errors().get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID)) } else { - assertEquals(Collections.singletonMap(topicPartition, Errors.PRODUCER_FENCED), response.errors()) + assertEquals(Collections.singletonMap(topicPartition, Errors.PRODUCER_FENCED), response.errors().get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID)) } } } @@ -2165,7 +2165,7 @@ class KafkaApisTest { createKafkaApis().handleAddPartitionsToTxnRequest(request, RequestLocal.withThreadConfinedCaching) val response = verifyNoThrottling[AddPartitionsToTxnResponse](request) - assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION, response.errors().get(invalidTopicPartition)) + assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION, response.errors().get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID).get(invalidTopicPartition)) } checkInvalidPartition(-1) From 32ca9bbaa78fc0e943d98376eee1f0558c7dffea Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 24 Feb 2023 13:19:45 -0800 Subject: [PATCH 13/17] just change the config for the test --- .../kafka/server/AddPartitionsToTxnRequestServerTest.scala | 4 +++- core/src/test/scala/unit/kafka/utils/TestUtils.scala | 3 --- 2 files changed, 3 insertions(+), 4 deletions(-) diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index 47f37317b8733..70587ae21db80 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -42,8 +42,10 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { private val topic1 = "topic1" val numPartitions = 1 - override def brokerPropertyOverrides(properties: Properties): Unit = + override def brokerPropertyOverrides(properties: Properties): Unit = { + properties.put(KafkaConfig.UnstableApiVersionsEnableProp, "true") properties.put(KafkaConfig.AutoCreateTopicsEnableProp, false.toString) + } @BeforeEach override def setUp(testInfo: TestInfo): Unit = { diff --git a/core/src/test/scala/unit/kafka/utils/TestUtils.scala b/core/src/test/scala/unit/kafka/utils/TestUtils.scala index 0af0e35b5a8c4..88cf3c11165a5 100755 --- a/core/src/test/scala/unit/kafka/utils/TestUtils.scala +++ b/core/src/test/scala/unit/kafka/utils/TestUtils.scala @@ -374,9 +374,6 @@ object TestUtils extends Logging { props.put(KafkaConfig.RackProp, nodeId.toString) props.put(KafkaConfig.ReplicaSelectorClassProp, "org.apache.kafka.common.replica.RackAwareReplicaSelector") } - - props.put(KafkaConfig.UnstableApiVersionsEnableProp, "true") - props } From a7abdae8465fca95f5243365219f7b16218dfbae Mon Sep 17 00:00:00 2001 From: Justine Date: Wed, 1 Mar 2023 15:13:52 -0800 Subject: [PATCH 14/17] Nits, changed to one callback, added kafka apis test for batched mode --- .../requests/AddPartitionsToTxnRequest.java | 9 +- .../requests/AddPartitionsToTxnResponse.java | 7 +- .../message/AddPartitionsToTxnRequest.json | 1 + .../message/AddPartitionsToTxnResponse.json | 2 +- .../transaction/TransactionCoordinator.scala | 16 ++-- .../main/scala/kafka/server/KafkaApis.scala | 78 ++++++++-------- ...ransactionCoordinatorConcurrencyTest.scala | 11 ++- .../TransactionCoordinatorTest.scala | 29 +++--- .../AddPartitionsToTxnRequestServerTest.scala | 51 +++++----- .../unit/kafka/server/KafkaApisTest.scala | 93 +++++++++++++++++-- 10 files changed, 190 insertions(+), 107 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 2d86eceec5290..730e2da4d02fa 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -43,8 +43,6 @@ public class AddPartitionsToTxnRequest extends AbstractRequest { private final AddPartitionsToTxnRequestData data; - private final short version; - public static class Builder extends AbstractRequest.Builder { public final AddPartitionsToTxnRequestData data; @@ -56,7 +54,7 @@ public static Builder forClient(String transactionalId, AddPartitionsToTxnTopicCollection topics = buildTxnTopicCollection(partitions); return new Builder(ApiKeys.ADD_PARTITIONS_TO_TXN.oldestVersion(), - (short) 3, + (short) 3, new AddPartitionsToTxnRequestData() .setV3AndBelowTransactionalId(transactionalId) .setV3AndBelowProducerId(producerId) @@ -113,7 +111,6 @@ public String toString() { public AddPartitionsToTxnRequest(final AddPartitionsToTxnRequestData data, short version) { super(ApiKeys.ADD_PARTITIONS_TO_TXN, version); this.data = data; - this.version = version; } @Override @@ -125,7 +122,7 @@ public AddPartitionsToTxnRequestData data() { public AddPartitionsToTxnResponse getErrorResponse(int throttleTimeMs, Throwable e) { Errors error = Errors.forException(e); AddPartitionsToTxnResponseData response = new AddPartitionsToTxnResponseData(); - if (version < 4) { + if (version() < 4) { response.setResultsByTopicV3AndBelow(errorResponseForTopics(data.v3AndBelowTopics(), error)); } else { AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); @@ -161,7 +158,7 @@ public Map> partitionsByTransaction() { // Takes a version 3 or below request and returns a v4+ singleton (one transaction ID) request. public AddPartitionsToTxnRequest normalizeRequest() { - return new AddPartitionsToTxnRequest(new AddPartitionsToTxnRequestData().setTransactions(singletonTransaction()), version); + return new AddPartitionsToTxnRequest(new AddPartitionsToTxnRequestData().setTransactions(singletonTransaction()), version()); } private AddPartitionsToTxnTransactionCollection singletonTransaction() { diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index 804fdb0d7a52a..e21e0661abbd2 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -71,11 +71,12 @@ public void maybeSetThrottleTimeMs(int throttleTimeMs) { public Map> errors() { Map> errorsMap = new HashMap<>(); - errorsMap.put(V3_AND_BELOW_TXN_ID, errorsForTransaction(this.data.resultsByTopicV3AndBelow())); + if (this.data.resultsByTopicV3AndBelow().size() != 0) { + errorsMap.put(V3_AND_BELOW_TXN_ID, errorsForTransaction(this.data.resultsByTopicV3AndBelow())); + } for (AddPartitionsToTxnResult result : this.data.resultsByTransaction()) { - String transactionalId = result.transactionalId(); - errorsMap.put(transactionalId, errorsForTransaction(data().resultsByTransaction().find(transactionalId).topicResults())); + errorsMap.put(result.transactionalId(), errorsForTransaction(result.topicResults())); } return errorsMap; diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json index a414dc53621ca..32bb9b8d1f763 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnRequest.json @@ -25,6 +25,7 @@ // Version 3 enables flexible versions. // // Version 4 adds VerifyOnly field to check if partitions are already in transaction and adds support to batch multiple transactions. + // Versions 3 and below will be exclusively used by clients and versions 4 and above will be used by brokers. // The AddPartitionsToTxnRequest version 4 API is added as part of KIP-890 and is still // under developement. Hence, the API is not exposed by default by brokers // unless explicitely enabled. diff --git a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json index f96b42219a18e..326b4acdb446a 100644 --- a/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json +++ b/clients/src/main/resources/common/message/AddPartitionsToTxnResponse.json @@ -29,7 +29,7 @@ "fields": [ { "name": "ThrottleTimeMs", "type": "int32", "versions": "0+", "about": "Duration in milliseconds for which the request was throttled due to a quota violation, or zero if the request did not violate any quota." }, - { "name": "ErrorCode", "type": "int16", "versions": "4+", + { "name": "ErrorCode", "type": "int16", "versions": "4+", "ignorable": true, "about": "The response top level error code." }, { "name": "ResultsByTransaction", "type": "[]AddPartitionsToTxnResult", "versions": "4+", "about": "Results categorized by transactional ID.", "fields": [ diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index 3a2891a9ad029..c62b635ed6544 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -93,8 +93,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, import TransactionCoordinator._ type InitProducerIdCallback = InitProducerIdResult => Unit - type AddPartitionsCallback = Errors => Unit - type VerifyPartitionsCallback = Map[TopicPartition, Errors] => Unit + type AddPartitionsCallback = (Option[Errors], Map[TopicPartition, Errors]) => Unit type EndTxnCallback = Errors => Unit type ApiResult[T] = Either[Errors, T] @@ -323,13 +322,13 @@ class TransactionCoordinator(txnConfig: TransactionConfig, def verifyAddPartitionsToTransaction(partitions: collection.Set[TopicPartition], txnMetadataPartitions: Set[TopicPartition], - responseCallback: VerifyPartitionsCallback): Unit = { + responseCallback: AddPartitionsCallback): Unit = { val addedPartitions = partitions.intersect(txnMetadataPartitions) val nonAddedPartitions = partitions.diff(txnMetadataPartitions) val errors = mutable.Map[TopicPartition, Errors]() addedPartitions.foreach(errors.put(_, Errors.NONE)) nonAddedPartitions.foreach(errors.put(_, Errors.INVALID_TXN_STATE)) - responseCallback(errors.toMap) + responseCallback(None, errors.toMap) } def handleAddPartitionsToTransaction(transactionalId: String, @@ -337,9 +336,12 @@ class TransactionCoordinator(txnConfig: TransactionConfig, producerEpoch: Short, partitions: collection.Set[TopicPartition], verifyOnly: Boolean, - responseCallback: AddPartitionsCallback, - verifyCallback: VerifyPartitionsCallback, + addPartitionsCallback: AddPartitionsCallback, requestLocal: RequestLocal = RequestLocal.NoCaching): Unit = { + def responseCallback(error: Errors): Unit = { + addPartitionsCallback(Some(error), Map.empty) + } + if (transactionalId == null || transactionalId.isEmpty) { debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") responseCallback(Errors.INVALID_REQUEST) @@ -370,7 +372,7 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } else { // If verifyOnly, we should have returned in the step above. If we didn't the partitions are not present in the transaction. if (verifyOnly) { - verifyAddPartitionsToTransaction(partitions, txnMetadata.topicPartitions.toSet, verifyCallback) + verifyAddPartitionsToTransaction(partitions, txnMetadata.topicPartitions.toSet, addPartitionsCallback) } Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds())) } diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 5706a3e3dc2fc..6a1df8e7d9f82 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -34,6 +34,7 @@ import org.apache.kafka.common.config.ConfigResource import org.apache.kafka.common.errors._ import org.apache.kafka.common.internals.Topic.{GROUP_METADATA_TOPIC_NAME, TRANSACTION_STATE_TOPIC_NAME, isInternal} import org.apache.kafka.common.internals.{FatalExitError, Topic} +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection import org.apache.kafka.common.message.AlterConfigsResponseData.AlterConfigsResourceResponse import org.apache.kafka.common.message.AlterPartitionReassignmentsResponseData.{ReassignablePartitionResponse, ReassignableTopicResponse} @@ -2405,10 +2406,10 @@ class KafkaApis(val requestChannel: RequestChannel, if (version < 4) { // There will only be one response in data. Add it to the response data object. val data = new AddPartitionsToTxnResponseData() - responses.forEach(result => { + responses.forEach { result => data.setResultsByTopicV3AndBelow(result.topicResults()) data.setThrottleTimeMs(requestThrottleMs) - }) + } new AddPartitionsToTxnResponse(data) } else { new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setThrottleTimeMs(requestThrottleMs).setResultsByTransaction(responses)) @@ -2416,33 +2417,34 @@ class KafkaApis(val requestChannel: RequestChannel, } val txns = addPartitionsToTxnRequest.data.transactions - def maybeSendResponse(): Unit = { - var canSend = false - responses.synchronized { - if (responses.size() == txns.size()) { - canSend = true - } + def addResultAndMaybeSendResponse(result: AddPartitionsToTxnResult): Unit = { + val canSend = responses.synchronized { + responses.add(result) + responses.size() == txns.size() } if (canSend) { requestHelper.sendResponseMaybeThrottle(request, createResponse) } } - txns.forEach( transaction => { + txns.forEach { transaction => val transactionalId = transaction.transactionalId val partitionsToAdd = partitionsByTransaction.get(transactionalId).asScala - + // Versions < 4 come from clients and must be authorized to write for the given transaction and for the given topics. if (version < 4 && !authHelper.authorize(request.context, WRITE, TRANSACTIONAL_ID, transactionalId)) { - responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) - maybeSendResponse() + addResultAndMaybeSendResponse(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED)) } else { val unauthorizedTopicErrors = mutable.Map[TopicPartition, Errors]() val nonExistingTopicErrors = mutable.Map[TopicPartition, Errors]() val authorizedPartitions = mutable.Set[TopicPartition]() - val authorizedTopics = if (version < 4) authHelper.filterByAuthorized(request.context, WRITE, TOPIC, - partitionsToAdd.filterNot(tp => Topic.isInternal(tp.topic)))(_.topic) else partitionsToAdd.map(_.topic).toSet + // Only request versions less than 4 need write authorization since they come from clients. + val authorizedTopics = + if (version < 4) + authHelper.filterByAuthorized(request.context, WRITE, TOPIC, partitionsToAdd.filterNot(tp => Topic.isInternal(tp.topic)))(_.topic) + else + partitionsToAdd.map(_.topic).toSet for (topicPartition <- partitionsToAdd) { if (!authorizedTopics.contains(topicPartition.topic)) unauthorizedTopicErrors += topicPartition -> Errors.TOPIC_AUTHORIZATION_FAILED @@ -2458,30 +2460,25 @@ class KafkaApis(val requestChannel: RequestChannel, // the authorization check to indicate that they were not added to the transaction. val partitionErrors = unauthorizedTopicErrors ++ nonExistingTopicErrors ++ authorizedPartitions.map(_ -> Errors.OPERATION_NOT_ATTEMPTED) - responses.add(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitionErrors.asJava)) - maybeSendResponse() + addResultAndMaybeSendResponse(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitionErrors.asJava)) } else { - def sendResponseCallback(error: Errors): Unit = { - val finalError = { - if (version < 2 && error == Errors.PRODUCER_FENCED) { - // For older clients, they could not understand the new PRODUCER_FENCED error code, - // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. - Errors.INVALID_PRODUCER_EPOCH - } else { - error - } - } - responses.synchronized { - responses.add(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) - } - maybeSendResponse() - } - def sendVerifyResponseCallback(errors: Map[TopicPartition, Errors]): Unit = { - responses.synchronized { - responses.add(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, errors.asJava)) + def sendResponseCallback(error: Option[Errors], errors: Map[TopicPartition, Errors]): Unit = { + error match { + case Some(error) => + val finalError = { + if (version < 2 && error == Errors.PRODUCER_FENCED) { + // For older clients, they could not understand the new PRODUCER_FENCED error code, + // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. + Errors.INVALID_PRODUCER_EPOCH + } else { + error + } + } + addResultAndMaybeSendResponse(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) + case None => + addResultAndMaybeSendResponse(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, errors.asJava)) } - maybeSendResponse() } txnCoordinator.handleAddPartitionsToTransaction(transactionalId, @@ -2490,11 +2487,10 @@ class KafkaApis(val requestChannel: RequestChannel, authorizedPartitions, transaction.verifyOnly, sendResponseCallback, - sendVerifyResponseCallback, requestLocal) } } - }) + } } def handleAddOffsetsToTxnRequest(request: RequestChannel.Request, requestLocal: RequestLocal): Unit = { @@ -2516,15 +2512,16 @@ class KafkaApis(val requestChannel: RequestChannel, .setThrottleTimeMs(requestThrottleMs)) ) else { - def sendResponseCallback(error: Errors): Unit = { + def sendResponseCallback(error: Option[Errors], errors: Map[TopicPartition, Errors]): Unit = { + // This will always have a single error, so error will always be defined. def createResponse(requestThrottleMs: Int): AbstractResponse = { val finalError = - if (addOffsetsToTxnRequest.version < 2 && error == Errors.PRODUCER_FENCED) { + if (addOffsetsToTxnRequest.version < 2 && error.get == Errors.PRODUCER_FENCED) { // For older clients, they could not understand the new PRODUCER_FENCED error code, // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. Errors.INVALID_PRODUCER_EPOCH } else { - error + error.get } val responseBody: AddOffsetsToTxnResponse = new AddOffsetsToTxnResponse( @@ -2544,7 +2541,6 @@ class KafkaApis(val requestChannel: RequestChannel, Set(offsetTopicPartition), false, sendResponseCallback, - null, requestLocal) } } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala index 625d94b959615..0b8e2064b043b 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala @@ -499,8 +499,14 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren abstract class TxnOperation[R] extends Operation { @volatile var result: Option[R] = None @volatile var results: Map[TopicPartition, R] = _ + def resultCallback(r: R): Unit = this.result = Some(r) - def resultPerPartitionCallback(r: Map[TopicPartition, R]): Unit = this.results = r + def addPartitionsResultCallback(rOpt: Option[R], rs: Map[TopicPartition, R]): Unit = { + rOpt match { + case Some(r) => this.result = rOpt + case None => this.results = rs + } + } } class InitProducerIdOperation(val producerIdAndEpoch: Option[ProducerIdAndEpoch] = None) extends TxnOperation[InitProducerIdResult] { @@ -524,8 +530,7 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren txnMetadata.producerEpoch, partitions, false, - resultCallback, - resultPerPartitionCallback, + addPartitionsResultCallback, RequestLocal.withThreadConfinedCaching) replicaManager.tryCompleteActions() } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index 389fb0a28795d..c22b1d4976930 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -201,19 +201,19 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(None)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsEmpty(): Unit = { - coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.INVALID_REQUEST, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsNull(): Unit = { - coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.INVALID_REQUEST, error) } @@ -222,7 +222,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.NOT_COORDINATOR)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.NOT_COORDINATOR, error) } @@ -231,7 +231,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.COORDINATOR_LOAD_IN_PROGRESS)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.COORDINATOR_LOAD_IN_PROGRESS, error) } @@ -250,7 +250,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, state, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.CONCURRENT_TRANSACTIONS, error) } @@ -260,7 +260,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 10, 9, 0, PrepareCommit, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.PRODUCER_FENCED, error) } @@ -291,7 +291,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, txnMetadata)))) - coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, false, addPartitionsToTxnCallback) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) verify(transactionManager).appendTransactionToLog( @@ -310,7 +310,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, addPartitionsToTxnCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -321,7 +321,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Ongoing, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, true, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, true, addPartitionsToTxnCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -335,7 +335,7 @@ class TransactionCoordinatorTest { val extraPartitions = partitions ++ Set(new TopicPartition("topic2", 0)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, extraPartitions, true, errorsCallback, errorsPerPartitionCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, extraPartitions, true, addPartitionsToTxnCallback) assertEquals(Errors.INVALID_TXN_STATE, errors(new TopicPartition("topic2", 0))) assertEquals(Errors.NONE, errors(new TopicPartition("topic1", 0))) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) @@ -1213,7 +1213,10 @@ class TransactionCoordinatorTest { error = ret } - def errorsPerPartitionCallback(ret: Map[TopicPartition, Errors]): Unit = { - errors = ret + def addPartitionsToTxnCallback(retOpt: Option[Errors], rets: Map[TopicPartition, Errors]): Unit = { + retOpt match { + case Some(ret) => error = ret + case None => errors = rets + } } } diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index 70587ae21db80..f12b266bc9c05 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -71,16 +71,16 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { transactionalId, producerId, producerEpoch, - List(createdTopicPartition, nonExistentTopic).asJava) - .build() + List(createdTopicPartition, nonExistentTopic).asJava + ).build() } else { val topics = new AddPartitionsToTxnTopicCollection() topics.add(new AddPartitionsToTxnTopic() - .setName(createdTopicPartition.topic()) - .setPartitions(Collections.singletonList(createdTopicPartition.partition()))) + .setName(createdTopicPartition.topic) + .setPartitions(Collections.singletonList(createdTopicPartition.partition))) topics.add(new AddPartitionsToTxnTopic() - .setName(nonExistentTopic.topic()) - .setPartitions(Collections.singletonList(nonExistentTopic.partition()))) + .setName(nonExistentTopic.topic) + .setPartitions(Collections.singletonList(nonExistentTopic.partition))) val transactions = new AddPartitionsToTxnTransactionCollection() transactions.add(new AddPartitionsToTxnTransaction() @@ -99,7 +99,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { if (version < 4) response.errors.get(AddPartitionsToTxnResponse.V3_AND_BELOW_TXN_ID) else - response.errorsForTransaction(response.getTransactionTopicResults(transactionalId)) + response.errors.get(transactionalId) assertEquals(2, errors.size) @@ -120,8 +120,8 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val txn2Topics = new AddPartitionsToTxnTopicCollection() txn2Topics.add(new AddPartitionsToTxnTopic() - .setName(tp0.topic()) - .setPartitions(Collections.singletonList(tp0.partition()))) + .setName(tp0.topic) + .setPartitions(Collections.singletonList(tp0.partition))) val (coordinatorId, txn1) = setUpTransactions(transactionalId1, false, Set(tp0)) @@ -139,14 +139,13 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val response = connectAndReceive[AddPartitionsToTxnResponse](request, brokerSocketServer(coordinatorId)) val errors = response.errors() + + val expectedErrors = Map( + transactionalId1 -> Collections.singletonMap(tp0, Errors.NONE), + transactionalId2 -> Collections.singletonMap(tp0, Errors.INVALID_PRODUCER_ID_MAPPING) + ).asJava - assertTrue(errors.containsKey(transactionalId1)) - assertTrue(errors.get(transactionalId1).containsKey(tp0)) - assertEquals(Errors.NONE, errors.get(transactionalId1).get(tp0)) - - assertTrue(errors.containsKey(transactionalId2)) - assertTrue(errors.get(transactionalId1).containsKey(tp0)) - assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, errors.get(transactionalId2).get(tp0)) + assertEquals(expectedErrors, errors) } @Test @@ -165,9 +164,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val verifyErrors = verifyResponse.errors() - assertTrue(verifyErrors.containsKey(transactionalId)) - assertTrue(verifyErrors.get(transactionalId).containsKey(tp0)) - assertEquals(Errors.INVALID_TXN_STATE, verifyErrors.get(transactionalId).get(tp0)) + assertEquals(Collections.singletonMap(transactionalId, Collections.singletonMap(tp0, Errors.INVALID_TXN_STATE)), verifyErrors) } private def setUpTransactions(transactionalId: String, verifyOnly: Boolean, partitions: Set[TopicPartition]): (Int, AddPartitionsToTxnTransaction) = { @@ -185,11 +182,11 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { val producerEpoch1 = initPidResponse.data().producerEpoch() val txn1Topics = new AddPartitionsToTxnTopicCollection() - partitions.foreach(tp => + partitions.foreach { tp => txn1Topics.add(new AddPartitionsToTxnTopic() - .setName(tp.topic()) - .setPartitions(Collections.singletonList(tp.partition()))) - ) + .setName(tp.topic) + .setPartitions(Collections.singletonList(tp.partition))) + } (coordinatorId, new AddPartitionsToTxnTransaction() .setTransactionalId(transactionalId) @@ -203,11 +200,11 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { object AddPartitionsToTxnRequestServerTest { def parameters: JStream[Arguments] = { val arguments = mutable.ListBuffer[Arguments]() - ApiKeys.ADD_PARTITIONS_TO_TXN.allVersions().forEach( version => - Array("kraft", "zk").foreach( quorum => + ApiKeys.ADD_PARTITIONS_TO_TXN.allVersions().forEach { version => + Array("kraft", "zk").foreach { quorum => arguments += Arguments.of(quorum, version) - ) - ) + } + } arguments.asJava.stream() } } diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index b962dbb6b0cdf..79fad74e18571 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -40,6 +40,7 @@ import org.apache.kafka.common.errors.UnsupportedVersionException import org.apache.kafka.common.internals.{KafkaFutureImpl, Topic} import org.apache.kafka.common.memory.MemoryPool import org.apache.kafka.common.config.ConfigResource.Type.{BROKER, BROKER_LOGGER} +import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.{AddPartitionsToTxnTopic, AddPartitionsToTxnTopicCollection, AddPartitionsToTxnTransaction, AddPartitionsToTxnTransactionCollection} import org.apache.kafka.common.message.AlterConfigsRequestData.{AlterConfigsResourceCollection => LAlterConfigsResourceCollection} import org.apache.kafka.common.message.AlterConfigsRequestData.{AlterConfigsResource => LAlterConfigsResource} import org.apache.kafka.common.message.AlterConfigsRequestData.{AlterableConfigCollection => LAlterableConfigCollection} @@ -1935,7 +1936,7 @@ class KafkaApisTest { reset(replicaManager, clientRequestQuotaManager, requestChannel, groupCoordinator, txnCoordinator) val capturedResponse: ArgumentCaptor[AddOffsetsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddOffsetsToTxnResponse]) - val responseCallback: ArgumentCaptor[Errors => Unit] = ArgumentCaptor.forClass(classOf[Errors => Unit]) + val responseCallback: ArgumentCaptor[(Option[Errors], Map[TopicPartition, Errors]) => Unit] = ArgumentCaptor.forClass(classOf[(Option[Errors], Map[TopicPartition, Errors]) => Unit]) val groupId = "groupId" val transactionalId = "txnId" @@ -1964,9 +1965,8 @@ class KafkaApisTest { ArgumentMatchers.eq(Set(new TopicPartition(Topic.GROUP_METADATA_TOPIC_NAME, partition))), ArgumentMatchers.eq(false), responseCallback.capture(), - ArgumentMatchers.any(), ArgumentMatchers.eq(requestLocal) - )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) + )).thenAnswer(_ => responseCallback.getValue.apply(Some(Errors.PRODUCER_FENCED), Map.empty)) createKafkaApis().handleAddOffsetsToTxnRequest(request, requestLocal) @@ -1995,7 +1995,7 @@ class KafkaApisTest { reset(replicaManager, clientRequestQuotaManager, requestChannel, txnCoordinator) val capturedResponse: ArgumentCaptor[AddPartitionsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResponse]) - val responseCallback: ArgumentCaptor[Errors => Unit] = ArgumentCaptor.forClass(classOf[Errors => Unit]) + val responseCallback: ArgumentCaptor[(Option[Errors], Map[TopicPartition, Errors]) => Unit] = ArgumentCaptor.forClass(classOf[(Option[Errors], Map[TopicPartition, Errors]) => Unit]) val transactionalId = "txnId" val producerId = 15L @@ -2020,9 +2020,8 @@ class KafkaApisTest { ArgumentMatchers.eq(Set(topicPartition)), ArgumentMatchers.eq(false), responseCallback.capture(), - ArgumentMatchers.any(), ArgumentMatchers.eq(requestLocal) - )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) + )).thenAnswer(_ => responseCallback.getValue.apply(Some(Errors.PRODUCER_FENCED), Map.empty)) createKafkaApis().handleAddPartitionsToTxnRequest(request, requestLocal) @@ -2041,6 +2040,88 @@ class KafkaApisTest { } } + @Test + def testBatchedRequest(): Unit = { + val topic = "topic" + addTopicToMetadataCache(topic, numPartitions = 2) + + val capturedResponse: ArgumentCaptor[AddPartitionsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResponse]) + val responseCallback: ArgumentCaptor[(Option[Errors], Map[TopicPartition, Errors]) => Unit] = ArgumentCaptor.forClass(classOf[(Option[Errors], Map[TopicPartition, Errors]) => Unit]) + + val transactionalId1 = "txnId1" + val transactionalId2 = "txnId2" + val producerId = 15L + val epoch = 0.toShort + + val tp0 = new TopicPartition(topic, 0) + val tp1 = new TopicPartition(topic, 1) + + val addPartitionsToTxnRequest = AddPartitionsToTxnRequest.Builder.forBroker( + new AddPartitionsToTxnTransactionCollection( + List(new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId1) + .setProducerId(producerId) + .setProducerEpoch(epoch) + .setVerifyOnly(true) + .setTopics(new AddPartitionsToTxnTopicCollection( + Collections.singletonList(new AddPartitionsToTxnTopic() + .setName(tp0.topic) + .setPartitions(Collections.singletonList(tp0.partition)) + ).iterator()) + ), new AddPartitionsToTxnTransaction() + .setTransactionalId(transactionalId2) + .setProducerId(producerId) + .setProducerEpoch(epoch) + .setVerifyOnly(false) + .setTopics(new AddPartitionsToTxnTopicCollection( + Collections.singletonList(new AddPartitionsToTxnTopic() + .setName(tp1.topic) + .setPartitions(Collections.singletonList(tp1.partition)) + ).iterator()) + ) + ).asJava.iterator() + ) + ).build(4.toShort) + val request = buildRequest(addPartitionsToTxnRequest) + + val requestLocal = RequestLocal.withThreadConfinedCaching + when(txnCoordinator.handleAddPartitionsToTransaction( + ArgumentMatchers.eq(transactionalId1), + ArgumentMatchers.eq(producerId), + ArgumentMatchers.eq(epoch), + ArgumentMatchers.eq(Set(tp0)), + ArgumentMatchers.eq(true), + responseCallback.capture(), + ArgumentMatchers.eq(requestLocal) + )).thenAnswer(_ => responseCallback.getValue.apply(None, Map(tp0 -> Errors.NONE))) + + when(txnCoordinator.handleAddPartitionsToTransaction( + ArgumentMatchers.eq(transactionalId2), + ArgumentMatchers.eq(producerId), + ArgumentMatchers.eq(epoch), + ArgumentMatchers.eq(Set(tp1)), + ArgumentMatchers.eq(false), + responseCallback.capture(), + ArgumentMatchers.eq(requestLocal) + )).thenAnswer(_ => responseCallback.getValue.apply(None, Map(tp1 -> Errors.PRODUCER_FENCED))) + + createKafkaApis().handleAddPartitionsToTxnRequest(request, requestLocal) + + verify(requestChannel).sendResponse( + ArgumentMatchers.eq(request), + capturedResponse.capture(), + ArgumentMatchers.eq(None) + ) + val response = capturedResponse.getValue + + val expectedErrors = Map( + transactionalId1 -> Collections.singletonMap(tp0, Errors.NONE), + transactionalId2 -> Collections.singletonMap(tp1, Errors.PRODUCER_FENCED) + ).asJava + + assertEquals(expectedErrors, response.errors()) + } + @Test def shouldReplaceProducerFencedWithInvalidProducerEpochInEndTxnWithOlderClient(): Unit = { val topic = "topic" From fd1f92b897ceaf27ebde7f654a72d0bb320d414f Mon Sep 17 00:00:00 2001 From: Justine Date: Thu, 2 Mar 2023 14:36:26 -0800 Subject: [PATCH 15/17] Redo callbacks by splitting into two methods, clean up error handling for top level errors, minor changes --- .../requests/AddPartitionsToTxnRequest.java | 3 - .../requests/AddPartitionsToTxnResponse.java | 8 +- .../internals/TransactionManagerTest.java | 10 +- .../AddPartitionsToTxnRequestTest.java | 10 +- .../AddPartitionsToTxnResponseTest.java | 8 +- .../transaction/TransactionCoordinator.scala | 123 ++++++++++-------- .../main/scala/kafka/server/KafkaApis.scala | 56 ++++---- ...ransactionCoordinatorConcurrencyTest.scala | 12 +- .../TransactionCoordinatorTest.scala | 32 +++-- .../unit/kafka/server/KafkaApisTest.scala | 29 ++--- 10 files changed, 155 insertions(+), 136 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 730e2da4d02fa..7e13e616ac1be 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -126,9 +126,6 @@ public AddPartitionsToTxnResponse getErrorResponse(int throttleTimeMs, Throwable response.setResultsByTopicV3AndBelow(errorResponseForTopics(data.v3AndBelowTopics(), error)); } else { AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); - for (AddPartitionsToTxnTransaction transaction : data().transactions()) { - results.add(errorResponseForTransaction(transaction.transactionalId(), error)); - } response.setResultsByTransaction(results); response.setErrorCode(error.code()); } diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index e21e0661abbd2..29518c251f2b2 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -119,7 +119,7 @@ public AddPartitionsToTxnTopicResultCollection getTransactionTopicResults(String return data.resultsByTransaction().find(transactionalId).topicResults(); } - public Map errorsForTransaction(AddPartitionsToTxnTopicResultCollection topicCollection) { + public static Map errorsForTransaction(AddPartitionsToTxnTopicResultCollection topicCollection) { Map topicResults = new HashMap<>(); for (AddPartitionsToTxnTopicResult topicResult : topicCollection) { for (AddPartitionsToTxnPartitionResult partitionResult : topicResult.resultsByPartition()) { @@ -133,6 +133,12 @@ public Map errorsForTransaction(AddPartitionsToTxnTopicR @Override public Map errorCounts() { List allErrors = new ArrayList<>(); + + // If we are not using this field, we have request 4 or later + if (this.data.resultsByTopicV3AndBelow().size() == 0) { + allErrors.add(Errors.forCode(data.errorCode())); + } + errors().forEach((txnId, errors) -> allErrors.addAll(errors.values()) ); diff --git a/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java b/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java index 8c0f5c51d462f..06bad27220525 100644 --- a/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java +++ b/clients/src/test/java/org/apache/kafka/clients/producer/internals/TransactionManagerTest.java @@ -1309,7 +1309,7 @@ public void testCommitWithTopicAuthorizationFailureInAddPartitionsInFlight() thr AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(0); client.respond(body -> { AddPartitionsToTxnRequest request = (AddPartitionsToTxnRequest) body; - assertEquals(new HashSet<>(AddPartitionsToTxnRequest.getPartitions(request.data().v3AndBelowTopics())), new HashSet<>(errors.keySet())); + assertEquals(new HashSet<>(getPartitionsFromV3Request(request)), new HashSet<>(errors.keySet())); return true; }, new AddPartitionsToTxnResponse(data)); @@ -3447,7 +3447,7 @@ private void prepareAddPartitionsToTxn(final Map errors) AddPartitionsToTxnResponseData data = new AddPartitionsToTxnResponseData().setResultsByTopicV3AndBelow(result.topicResults()).setThrottleTimeMs(0); client.prepareResponse(body -> { AddPartitionsToTxnRequest request = (AddPartitionsToTxnRequest) body; - assertEquals(new HashSet<>(AddPartitionsToTxnRequest.getPartitions(request.data().v3AndBelowTopics())), new HashSet<>(errors.keySet())); + assertEquals(new HashSet<>(getPartitionsFromV3Request(request)), new HashSet<>(errors.keySet())); return true; }, new AddPartitionsToTxnResponse(data)); } @@ -3552,11 +3552,15 @@ private MockClient.RequestMatcher addPartitionsRequestMatcher(final TopicPartiti AddPartitionsToTxnRequest addPartitionsToTxnRequest = (AddPartitionsToTxnRequest) body; assertEquals(producerId, addPartitionsToTxnRequest.data().v3AndBelowProducerId()); assertEquals(epoch, addPartitionsToTxnRequest.data().v3AndBelowProducerEpoch()); - assertEquals(singletonList(topicPartition), AddPartitionsToTxnRequest.getPartitions(addPartitionsToTxnRequest.data().v3AndBelowTopics())); + assertEquals(singletonList(topicPartition), getPartitionsFromV3Request(addPartitionsToTxnRequest)); assertEquals(transactionalId, addPartitionsToTxnRequest.data().v3AndBelowTransactionalId()); return true; }; } + + private List getPartitionsFromV3Request(AddPartitionsToTxnRequest request) { + return AddPartitionsToTxnRequest.getPartitions(request.data().v3AndBelowTopics()); + } private void prepareEndTxnResponse(Errors error, final TransactionResult result, final long producerId, final short epoch) { this.prepareEndTxnResponse(error, result, producerId, epoch, false); diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java index 0d7c299401786..92bb8741be09d 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequestTest.java @@ -35,6 +35,7 @@ import org.junit.jupiter.api.Test; import org.junit.jupiter.params.ParameterizedTest; +import static org.apache.kafka.common.requests.AddPartitionsToTxnResponse.errorsForTransaction; import static org.junit.jupiter.api.Assertions.assertEquals; public class AddPartitionsToTxnRequestTest { @@ -78,11 +79,14 @@ public void testConstructor(short version) { } AddPartitionsToTxnResponse response = request.getErrorResponse(throttleTimeMs, Errors.UNKNOWN_TOPIC_OR_PARTITION.exception()); - assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); assertEquals(throttleTimeMs, response.throttleTimeMs()); if (version >= 4) { assertEquals(Errors.UNKNOWN_TOPIC_OR_PARTITION.code(), response.data().errorCode()); + // Since the error is top level, we count it as one error in the counts. + assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 1), response.errorCounts()); + } else { + assertEquals(Collections.singletonMap(Errors.UNKNOWN_TOPIC_OR_PARTITION, 2), response.errorCounts()); } } @@ -108,8 +112,8 @@ public void testBatchedRequests() { .setResultsByTransaction(results) .setThrottleTimeMs(throttleTimeMs)); - assertEquals(Collections.singletonMap(tp0, Errors.UNKNOWN_TOPIC_OR_PARTITION), response.errorsForTransaction(response.getTransactionTopicResults(transactionalId1))); - assertEquals(Collections.singletonMap(tp1, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED), response.errorsForTransaction(response.getTransactionTopicResults(transactionalId2))); + assertEquals(Collections.singletonMap(tp0, Errors.UNKNOWN_TOPIC_OR_PARTITION), errorsForTransaction(response.getTransactionTopicResults(transactionalId1))); + assertEquals(Collections.singletonMap(tp1, Errors.TRANSACTIONAL_ID_AUTHORIZATION_FAILED), errorsForTransaction(response.getTransactionTopicResults(transactionalId2))); } @Test diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index 044bd4b4884f0..6f81dad0279cc 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -34,6 +34,7 @@ import java.util.HashMap; import java.util.Map; +import static org.apache.kafka.common.requests.AddPartitionsToTxnResponse.errorsForTransaction; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertTrue; @@ -106,11 +107,12 @@ public void testParse(short version) { AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(data); Map newExpectedErrorCounts = new HashMap<>(); + newExpectedErrorCounts.put(Errors.NONE, 1); // top level error newExpectedErrorCounts.put(errorOne, 2); newExpectedErrorCounts.put(errorTwo, 1); AddPartitionsToTxnResponse parsedResponse = AddPartitionsToTxnResponse.parse(response.serialize(version), version); - assertEquals(txnTwoExpectedErrors, parsedResponse.errorsForTransaction(response.getTransactionTopicResults("txn2"))); + assertEquals(txnTwoExpectedErrors, errorsForTransaction(response.getTransactionTopicResults("txn2"))); assertEquals(newExpectedErrorCounts, parsedResponse.errorCounts()); assertEquals(throttleTimeMs, parsedResponse.throttleTimeMs()); assertTrue(parsedResponse.shouldClientThrottle(version)); @@ -131,7 +133,7 @@ public void testBatchedErrors() { AddPartitionsToTxnResponse response = new AddPartitionsToTxnResponse(new AddPartitionsToTxnResponseData().setResultsByTransaction(results)); - assertEquals(txn1Errors, response.errorsForTransaction(response.getTransactionTopicResults("txn1"))); - assertEquals(txn2Errors, response.errorsForTransaction(response.getTransactionTopicResults("txn2"))); + assertEquals(txn1Errors, errorsForTransaction(response.getTransactionTopicResults("txn1"))); + assertEquals(txn2Errors, errorsForTransaction(response.getTransactionTopicResults("txn2"))); } } diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index c62b635ed6544..73ebb7603391d 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -22,15 +22,17 @@ import kafka.server.{KafkaConfig, MetadataCache, ReplicaManager, RequestLocal} import kafka.utils.Logging import org.apache.kafka.common.TopicPartition import org.apache.kafka.common.internals.Topic +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult import org.apache.kafka.common.message.{DescribeTransactionsResponseData, ListTransactionsResponseData} import org.apache.kafka.common.metrics.Metrics import org.apache.kafka.common.protocol.Errors import org.apache.kafka.common.record.RecordBatch -import org.apache.kafka.common.requests.TransactionResult +import org.apache.kafka.common.requests.{AddPartitionsToTxnResponse, TransactionResult} import org.apache.kafka.common.utils.{LogContext, ProducerIdAndEpoch, Time} import org.apache.kafka.server.util.Scheduler import scala.collection.mutable +import scala.jdk.CollectionConverters._ object TransactionCoordinator { @@ -93,7 +95,8 @@ class TransactionCoordinator(txnConfig: TransactionConfig, import TransactionCoordinator._ type InitProducerIdCallback = InitProducerIdResult => Unit - type AddPartitionsCallback = (Option[Errors], Map[TopicPartition, Errors]) => Unit + type AddPartitionsCallback = Errors => Unit + type VerifyPartitionsCallback = AddPartitionsToTxnResult => Unit type EndTxnCallback = Errors => Unit type ApiResult[T] = Either[Errors, T] @@ -319,80 +322,94 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } } } - - def verifyAddPartitionsToTransaction(partitions: collection.Set[TopicPartition], - txnMetadataPartitions: Set[TopicPartition], - responseCallback: AddPartitionsCallback): Unit = { - val addedPartitions = partitions.intersect(txnMetadataPartitions) - val nonAddedPartitions = partitions.diff(txnMetadataPartitions) - val errors = mutable.Map[TopicPartition, Errors]() - addedPartitions.foreach(errors.put(_, Errors.NONE)) - nonAddedPartitions.foreach(errors.put(_, Errors.INVALID_TXN_STATE)) - responseCallback(None, errors.toMap) + + def handleVerifyPartitionsInTransaction(transactionalId: String, + producerId: Long, + producerEpoch: Short, + partitions: collection.Set[TopicPartition], + responseCallback: VerifyPartitionsCallback): Unit = { + if (transactionalId == null || transactionalId.isEmpty) { + debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") + responseCallback(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitions.map(_ -> Errors.INVALID_REQUEST).toMap.asJava)) + } else { + val result: ApiResult[(Int, TransactionMetadata)] = getTransactionMetadata(transactionalId, producerId, producerEpoch, partitions) + + result match { + case Left(err) => + debug(s"Returning $err error code to client for $transactionalId's AddPartitions request") + responseCallback(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitions.map(_ -> err).toMap.asJava)) + + case Right((_, txnMetadata)) => + val txnMetadataPartitions = txnMetadata.topicPartitions + val addedPartitions = partitions.intersect(txnMetadataPartitions) + val nonAddedPartitions = partitions.diff(txnMetadataPartitions) + val errors = mutable.Map[TopicPartition, Errors]() + addedPartitions.foreach(errors.put(_, Errors.NONE)) + nonAddedPartitions.foreach(errors.put(_, Errors.INVALID_TXN_STATE)) + responseCallback(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, errors.asJava)) + } + } } def handleAddPartitionsToTransaction(transactionalId: String, producerId: Long, producerEpoch: Short, partitions: collection.Set[TopicPartition], - verifyOnly: Boolean, - addPartitionsCallback: AddPartitionsCallback, + responseCallback: AddPartitionsCallback, requestLocal: RequestLocal = RequestLocal.NoCaching): Unit = { - def responseCallback(error: Errors): Unit = { - addPartitionsCallback(Some(error), Map.empty) - } - if (transactionalId == null || transactionalId.isEmpty) { debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") responseCallback(Errors.INVALID_REQUEST) } else { // try to update the transaction metadata and append the updated metadata to txn log; // if there is no such metadata treat it as invalid producerId mapping error. - val result: ApiResult[(Int, TxnTransitMetadata)] = txnManager.getTransactionState(transactionalId).flatMap { - case None => Left(Errors.INVALID_PRODUCER_ID_MAPPING) - - case Some(epochAndMetadata) => - val coordinatorEpoch = epochAndMetadata.coordinatorEpoch - val txnMetadata = epochAndMetadata.transactionMetadata - - // generate the new transaction metadata with added partitions - txnMetadata.inLock { - if (txnMetadata.producerId != producerId) { - Left(Errors.INVALID_PRODUCER_ID_MAPPING) - } else if (txnMetadata.producerEpoch != producerEpoch) { - Left(Errors.PRODUCER_FENCED) - } else if (txnMetadata.pendingTransitionInProgress) { - // return a retriable exception to let the client backoff and retry - Left(Errors.CONCURRENT_TRANSACTIONS) - } else if (txnMetadata.state == PrepareCommit || txnMetadata.state == PrepareAbort) { - Left(Errors.CONCURRENT_TRANSACTIONS) - } else if (txnMetadata.state == Ongoing && partitions.subsetOf(txnMetadata.topicPartitions)) { - // this is an optimization: if the partitions are already in the metadata reply OK immediately - Left(Errors.NONE) - } else { - // If verifyOnly, we should have returned in the step above. If we didn't the partitions are not present in the transaction. - if (verifyOnly) { - verifyAddPartitionsToTransaction(partitions, txnMetadata.topicPartitions.toSet, addPartitionsCallback) - } - Right(coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds())) - } - } - } + val result: ApiResult[(Int, TransactionMetadata)] = getTransactionMetadata(transactionalId, producerId, producerEpoch, partitions) result match { case Left(err) => debug(s"Returning $err error code to client for $transactionalId's AddPartitions request") responseCallback(err) - case Right((coordinatorEpoch, newMetadata)) if !verifyOnly => - txnManager.appendTransactionToLog(transactionalId, coordinatorEpoch, newMetadata, + case Right((coordinatorEpoch, txnMetadata)) => + txnManager.appendTransactionToLog(transactionalId, coordinatorEpoch, txnMetadata.prepareAddPartitions(partitions.toSet, time.milliseconds()), responseCallback, requestLocal = requestLocal) - - case _ => - // We only hit this case if verifyOnly and some partitions were not present. If so, we already handled the response. } } } + + private def getTransactionMetadata(transactionalId: String, + producerId: Long, + producerEpoch: Short, + partitions: collection.Set[TopicPartition]): ApiResult[(Int, TransactionMetadata)] = { + // try to update the transaction metadata and append the updated metadata to txn log; + // if there is no such metadata treat it as invalid producerId mapping error. + txnManager.getTransactionState(transactionalId).flatMap { + case None => Left(Errors.INVALID_PRODUCER_ID_MAPPING) + + case Some(epochAndMetadata) => + val coordinatorEpoch = epochAndMetadata.coordinatorEpoch + val txnMetadata = epochAndMetadata.transactionMetadata + + // generate the new transaction metadata with added partitions + txnMetadata.inLock { + if (txnMetadata.producerId != producerId) { + Left(Errors.INVALID_PRODUCER_ID_MAPPING) + } else if (txnMetadata.producerEpoch != producerEpoch) { + Left(Errors.PRODUCER_FENCED) + } else if (txnMetadata.pendingTransitionInProgress) { + // return a retriable exception to let the client backoff and retry + Left(Errors.CONCURRENT_TRANSACTIONS) + } else if (txnMetadata.state == PrepareCommit || txnMetadata.state == PrepareAbort) { + Left(Errors.CONCURRENT_TRANSACTIONS) + } else if (txnMetadata.state == Ongoing && partitions.subsetOf(txnMetadata.topicPartitions)) { + // this is an optimization: if the partitions are already in the metadata reply OK immediately + Left(Errors.NONE) + } else { + Right(coordinatorEpoch, txnMetadata) + } + } + } + } /** * Load state from the given partition and begin handling requests for groups which map to this partition. diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index 6a1df8e7d9f82..c489658e33880 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2462,32 +2462,34 @@ class KafkaApis(val requestChannel: RequestChannel, authorizedPartitions.map(_ -> Errors.OPERATION_NOT_ATTEMPTED) addResultAndMaybeSendResponse(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitionErrors.asJava)) } else { - - def sendResponseCallback(error: Option[Errors], errors: Map[TopicPartition, Errors]): Unit = { - error match { - case Some(error) => - val finalError = { - if (version < 2 && error == Errors.PRODUCER_FENCED) { - // For older clients, they could not understand the new PRODUCER_FENCED error code, - // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. - Errors.INVALID_PRODUCER_EPOCH - } else { - error - } - } - addResultAndMaybeSendResponse(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) - case None => - addResultAndMaybeSendResponse(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, errors.asJava)) + def sendResponseCallback(error: Errors): Unit = { + val finalError = { + if (version < 2 && error == Errors.PRODUCER_FENCED) { + // For older clients, they could not understand the new PRODUCER_FENCED error code, + // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. + Errors.INVALID_PRODUCER_EPOCH + } else { + error + } } + addResultAndMaybeSendResponse(addPartitionsToTxnRequest.errorResponseForTransaction(transactionalId, finalError)) } - txnCoordinator.handleAddPartitionsToTransaction(transactionalId, - transaction.producerId, - transaction.producerEpoch, - authorizedPartitions, - transaction.verifyOnly, - sendResponseCallback, - requestLocal) + + if (!transaction.verifyOnly) { + txnCoordinator.handleAddPartitionsToTransaction(transactionalId, + transaction.producerId, + transaction.producerEpoch, + authorizedPartitions, + sendResponseCallback, + requestLocal) + } else { + txnCoordinator.handleVerifyPartitionsInTransaction(transactionalId, + transaction.producerId, + transaction.producerEpoch, + authorizedPartitions, + addResultAndMaybeSendResponse) + } } } } @@ -2512,16 +2514,15 @@ class KafkaApis(val requestChannel: RequestChannel, .setThrottleTimeMs(requestThrottleMs)) ) else { - def sendResponseCallback(error: Option[Errors], errors: Map[TopicPartition, Errors]): Unit = { - // This will always have a single error, so error will always be defined. + def sendResponseCallback(error: Errors): Unit = { def createResponse(requestThrottleMs: Int): AbstractResponse = { val finalError = - if (addOffsetsToTxnRequest.version < 2 && error.get == Errors.PRODUCER_FENCED) { + if (addOffsetsToTxnRequest.version < 2 && error == Errors.PRODUCER_FENCED) { // For older clients, they could not understand the new PRODUCER_FENCED error code, // so we need to return the old INVALID_PRODUCER_EPOCH to have the same client handling logic. Errors.INVALID_PRODUCER_EPOCH } else { - error.get + error } val responseBody: AddOffsetsToTxnResponse = new AddOffsetsToTxnResponse( @@ -2539,7 +2540,6 @@ class KafkaApis(val requestChannel: RequestChannel, addOffsetsToTxnRequest.data.producerId, addOffsetsToTxnRequest.data.producerEpoch, Set(offsetTopicPartition), - false, sendResponseCallback, requestLocal) } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala index 0b8e2064b043b..c458ac191c0f8 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorConcurrencyTest.scala @@ -501,12 +501,7 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren @volatile var results: Map[TopicPartition, R] = _ def resultCallback(r: R): Unit = this.result = Some(r) - def addPartitionsResultCallback(rOpt: Option[R], rs: Map[TopicPartition, R]): Unit = { - rOpt match { - case Some(r) => this.result = rOpt - case None => this.results = rs - } - } + } class InitProducerIdOperation(val producerIdAndEpoch: Option[ProducerIdAndEpoch] = None) extends TxnOperation[InitProducerIdResult] { @@ -528,9 +523,8 @@ class TransactionCoordinatorConcurrencyTest extends AbstractCoordinatorConcurren transactionCoordinator.handleAddPartitionsToTransaction(txn.transactionalId, txnMetadata.producerId, txnMetadata.producerEpoch, - partitions, - false, - addPartitionsResultCallback, + partitions, + resultCallback, RequestLocal.withThreadConfinedCaching) replicaManager.tryCompleteActions() } diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index c22b1d4976930..0b3ad663333f1 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -17,9 +17,10 @@ package kafka.coordinator.transaction import org.apache.kafka.common.TopicPartition +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult import org.apache.kafka.common.protocol.Errors import org.apache.kafka.common.record.RecordBatch -import org.apache.kafka.common.requests.TransactionResult +import org.apache.kafka.common.requests.{AddPartitionsToTxnResponse, TransactionResult} import org.apache.kafka.common.utils.{LogContext, MockTime, ProducerIdAndEpoch} import org.apache.kafka.server.util.MockScheduler import org.junit.jupiter.api.Assertions._ @@ -201,19 +202,19 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(None)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, errorsCallback) assertEquals(Errors.INVALID_PRODUCER_ID_MAPPING, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsEmpty(): Unit = { - coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction("", 0L, 1, partitions, errorsCallback) assertEquals(Errors.INVALID_REQUEST, error) } @Test def shouldRespondWithInvalidRequestAddPartitionsToTransactionWhenTransactionalIdIsNull(): Unit = { - coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(null, 0L, 1, partitions, errorsCallback) assertEquals(Errors.INVALID_REQUEST, error) } @@ -222,7 +223,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.NOT_COORDINATOR)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, errorsCallback) assertEquals(Errors.NOT_COORDINATOR, error) } @@ -231,7 +232,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Left(Errors.COORDINATOR_LOAD_IN_PROGRESS)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 1, partitions, errorsCallback) assertEquals(Errors.COORDINATOR_LOAD_IN_PROGRESS, error) } @@ -250,7 +251,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, state, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, errorsCallback) assertEquals(Errors.CONCURRENT_TRANSACTIONS, error) } @@ -260,7 +261,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 10, 9, 0, PrepareCommit, mutable.Set.empty, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, errorsCallback) assertEquals(Errors.PRODUCER_FENCED, error) } @@ -291,7 +292,7 @@ class TransactionCoordinatorTest { when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, txnMetadata)))) - coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, producerId, producerEpoch, partitions, errorsCallback) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) verify(transactionManager).appendTransactionToLog( @@ -310,7 +311,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, false, addPartitionsToTxnCallback) + coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, errorsCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -321,7 +322,7 @@ class TransactionCoordinatorTest { .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Ongoing, partitions, 0, 0))))) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, partitions, true, addPartitionsToTxnCallback) + coordinator.handleVerifyPartitionsInTransaction(transactionalId, 0L, 0, partitions, verifyPartitionsInTxnCallback) assertEquals(Errors.NONE, error) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) } @@ -335,7 +336,7 @@ class TransactionCoordinatorTest { val extraPartitions = partitions ++ Set(new TopicPartition("topic2", 0)) - coordinator.handleAddPartitionsToTransaction(transactionalId, 0L, 0, extraPartitions, true, addPartitionsToTxnCallback) + coordinator.handleVerifyPartitionsInTransaction(transactionalId, 0L, 0, extraPartitions, verifyPartitionsInTxnCallback) assertEquals(Errors.INVALID_TXN_STATE, errors(new TopicPartition("topic2", 0))) assertEquals(Errors.NONE, errors(new TopicPartition("topic1", 0))) verify(transactionManager).getTransactionState(ArgumentMatchers.eq(transactionalId)) @@ -1213,10 +1214,7 @@ class TransactionCoordinatorTest { error = ret } - def addPartitionsToTxnCallback(retOpt: Option[Errors], rets: Map[TopicPartition, Errors]): Unit = { - retOpt match { - case Some(ret) => error = ret - case None => errors = rets - } + def verifyPartitionsInTxnCallback(result: AddPartitionsToTxnResult): Unit = { + errors = AddPartitionsToTxnResponse.errorsForTransaction(result.topicResults()).asScala.toMap } } diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index 79fad74e18571..d3cb4eb77788e 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -41,6 +41,7 @@ import org.apache.kafka.common.internals.{KafkaFutureImpl, Topic} import org.apache.kafka.common.memory.MemoryPool import org.apache.kafka.common.config.ConfigResource.Type.{BROKER, BROKER_LOGGER} import org.apache.kafka.common.message.AddPartitionsToTxnRequestData.{AddPartitionsToTxnTopic, AddPartitionsToTxnTopicCollection, AddPartitionsToTxnTransaction, AddPartitionsToTxnTransactionCollection} +import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult import org.apache.kafka.common.message.AlterConfigsRequestData.{AlterConfigsResourceCollection => LAlterConfigsResourceCollection} import org.apache.kafka.common.message.AlterConfigsRequestData.{AlterConfigsResource => LAlterConfigsResource} import org.apache.kafka.common.message.AlterConfigsRequestData.{AlterableConfigCollection => LAlterableConfigCollection} @@ -1936,7 +1937,7 @@ class KafkaApisTest { reset(replicaManager, clientRequestQuotaManager, requestChannel, groupCoordinator, txnCoordinator) val capturedResponse: ArgumentCaptor[AddOffsetsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddOffsetsToTxnResponse]) - val responseCallback: ArgumentCaptor[(Option[Errors], Map[TopicPartition, Errors]) => Unit] = ArgumentCaptor.forClass(classOf[(Option[Errors], Map[TopicPartition, Errors]) => Unit]) + val responseCallback: ArgumentCaptor[Errors => Unit] = ArgumentCaptor.forClass(classOf[Errors => Unit]) val groupId = "groupId" val transactionalId = "txnId" @@ -1963,10 +1964,9 @@ class KafkaApisTest { ArgumentMatchers.eq(producerId), ArgumentMatchers.eq(epoch), ArgumentMatchers.eq(Set(new TopicPartition(Topic.GROUP_METADATA_TOPIC_NAME, partition))), - ArgumentMatchers.eq(false), responseCallback.capture(), ArgumentMatchers.eq(requestLocal) - )).thenAnswer(_ => responseCallback.getValue.apply(Some(Errors.PRODUCER_FENCED), Map.empty)) + )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) createKafkaApis().handleAddOffsetsToTxnRequest(request, requestLocal) @@ -1995,7 +1995,7 @@ class KafkaApisTest { reset(replicaManager, clientRequestQuotaManager, requestChannel, txnCoordinator) val capturedResponse: ArgumentCaptor[AddPartitionsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResponse]) - val responseCallback: ArgumentCaptor[(Option[Errors], Map[TopicPartition, Errors]) => Unit] = ArgumentCaptor.forClass(classOf[(Option[Errors], Map[TopicPartition, Errors]) => Unit]) + val responseCallback: ArgumentCaptor[Errors => Unit] = ArgumentCaptor.forClass(classOf[Errors => Unit]) val transactionalId = "txnId" val producerId = 15L @@ -2018,10 +2018,9 @@ class KafkaApisTest { ArgumentMatchers.eq(producerId), ArgumentMatchers.eq(epoch), ArgumentMatchers.eq(Set(topicPartition)), - ArgumentMatchers.eq(false), responseCallback.capture(), ArgumentMatchers.eq(requestLocal) - )).thenAnswer(_ => responseCallback.getValue.apply(Some(Errors.PRODUCER_FENCED), Map.empty)) + )).thenAnswer(_ => responseCallback.getValue.apply(Errors.PRODUCER_FENCED)) createKafkaApis().handleAddPartitionsToTxnRequest(request, requestLocal) @@ -2046,7 +2045,8 @@ class KafkaApisTest { addTopicToMetadataCache(topic, numPartitions = 2) val capturedResponse: ArgumentCaptor[AddPartitionsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResponse]) - val responseCallback: ArgumentCaptor[(Option[Errors], Map[TopicPartition, Errors]) => Unit] = ArgumentCaptor.forClass(classOf[(Option[Errors], Map[TopicPartition, Errors]) => Unit]) + val responseCallback: ArgumentCaptor[Errors => Unit] = ArgumentCaptor.forClass(classOf[Errors => Unit]) + val verifyPartitionsCallback: ArgumentCaptor[AddPartitionsToTxnResult => Unit] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResult => Unit]) val transactionalId1 = "txnId1" val transactionalId2 = "txnId2" @@ -2062,7 +2062,7 @@ class KafkaApisTest { .setTransactionalId(transactionalId1) .setProducerId(producerId) .setProducerEpoch(epoch) - .setVerifyOnly(true) + .setVerifyOnly(false) .setTopics(new AddPartitionsToTxnTopicCollection( Collections.singletonList(new AddPartitionsToTxnTopic() .setName(tp0.topic) @@ -2072,7 +2072,7 @@ class KafkaApisTest { .setTransactionalId(transactionalId2) .setProducerId(producerId) .setProducerEpoch(epoch) - .setVerifyOnly(false) + .setVerifyOnly(true) .setTopics(new AddPartitionsToTxnTopicCollection( Collections.singletonList(new AddPartitionsToTxnTopic() .setName(tp1.topic) @@ -2090,20 +2090,17 @@ class KafkaApisTest { ArgumentMatchers.eq(producerId), ArgumentMatchers.eq(epoch), ArgumentMatchers.eq(Set(tp0)), - ArgumentMatchers.eq(true), responseCallback.capture(), ArgumentMatchers.eq(requestLocal) - )).thenAnswer(_ => responseCallback.getValue.apply(None, Map(tp0 -> Errors.NONE))) + )).thenAnswer(_ => responseCallback.getValue.apply(Errors.NONE)) - when(txnCoordinator.handleAddPartitionsToTransaction( + when(txnCoordinator.handleVerifyPartitionsInTransaction( ArgumentMatchers.eq(transactionalId2), ArgumentMatchers.eq(producerId), ArgumentMatchers.eq(epoch), ArgumentMatchers.eq(Set(tp1)), - ArgumentMatchers.eq(false), - responseCallback.capture(), - ArgumentMatchers.eq(requestLocal) - )).thenAnswer(_ => responseCallback.getValue.apply(None, Map(tp1 -> Errors.PRODUCER_FENCED))) + verifyPartitionsCallback.capture(), + )).thenAnswer(_ => verifyPartitionsCallback.getValue.apply(AddPartitionsToTxnResponse.resultForTransaction(transactionalId2, Map(tp1 -> Errors.PRODUCER_FENCED).asJava))) createKafkaApis().handleAddPartitionsToTxnRequest(request, requestLocal) From 4589d365866ed208d176f93ac49e9add4510e5f4 Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 3 Mar 2023 09:16:56 -0800 Subject: [PATCH 16/17] test fix --- .../org/apache/kafka/common/requests/RequestResponseTest.java | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java index 587e111a46dd7..7c059b41c7677 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java @@ -932,7 +932,7 @@ public void testDeletableTopicResultErrorMessageIsNullByDefault() { public void testErrorCountsIncludesNone() { assertEquals(1, createAddOffsetsToTxnResponse().errorCounts().get(Errors.NONE)); assertEquals(1, createAddPartitionsToTxnResponse((short) 3).errorCounts().get(Errors.NONE)); - assertEquals(1, createAddPartitionsToTxnResponse((short) 4).errorCounts().get(Errors.NONE)); + assertEquals(2, createAddPartitionsToTxnResponse((short) 4).errorCounts().get(Errors.NONE)); assertEquals(1, createAlterClientQuotasResponse().errorCounts().get(Errors.NONE)); assertEquals(1, createAlterConfigsResponse().errorCounts().get(Errors.NONE)); assertEquals(2, createAlterPartitionReassignmentsResponse().errorCounts().get(Errors.NONE)); From ac16b133e8666aee1e9b36f9cc0b5c44d71a58ac Mon Sep 17 00:00:00 2001 From: Justine Date: Fri, 3 Mar 2023 14:26:46 -0800 Subject: [PATCH 17/17] Small changes --- .../requests/AddPartitionsToTxnRequest.java | 8 ++----- .../requests/AddPartitionsToTxnResponse.java | 4 ++-- .../AddPartitionsToTxnResponseTest.java | 5 ++++ .../common/requests/RequestResponseTest.java | 4 ++-- .../transaction/TransactionCoordinator.scala | 23 +++++++++---------- .../main/scala/kafka/server/KafkaApis.scala | 2 +- .../TransactionCoordinatorTest.scala | 16 ++++++++----- .../AddPartitionsToTxnRequestServerTest.scala | 4 ++-- .../unit/kafka/server/KafkaApisTest.scala | 10 ++------ 9 files changed, 37 insertions(+), 39 deletions(-) diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java index 7e13e616ac1be..c91374fc50795 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnRequest.java @@ -26,7 +26,6 @@ import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResult; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnPartitionResultCollection; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResult; -import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnResultCollection; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResult; import org.apache.kafka.common.message.AddPartitionsToTxnResponseData.AddPartitionsToTxnTopicResultCollection; import org.apache.kafka.common.protocol.ApiKeys; @@ -53,8 +52,7 @@ public static Builder forClient(String transactionalId, AddPartitionsToTxnTopicCollection topics = buildTxnTopicCollection(partitions); - return new Builder(ApiKeys.ADD_PARTITIONS_TO_TXN.oldestVersion(), - (short) 3, + return new Builder(ApiKeys.ADD_PARTITIONS_TO_TXN.oldestVersion(), (short) 3, new AddPartitionsToTxnRequestData() .setV3AndBelowTransactionalId(transactionalId) .setV3AndBelowProducerId(producerId) @@ -68,7 +66,7 @@ public static Builder forBroker(AddPartitionsToTxnTransactionCollection transact .setTransactions(transactions)); } - public Builder(short minVersion, short maxVersion, AddPartitionsToTxnRequestData data) { + private Builder(short minVersion, short maxVersion, AddPartitionsToTxnRequestData data) { super(ApiKeys.ADD_PARTITIONS_TO_TXN, minVersion, maxVersion); this.data = data; @@ -125,8 +123,6 @@ public AddPartitionsToTxnResponse getErrorResponse(int throttleTimeMs, Throwable if (version() < 4) { response.setResultsByTopicV3AndBelow(errorResponseForTopics(data.v3AndBelowTopics(), error)); } else { - AddPartitionsToTxnResultCollection results = new AddPartitionsToTxnResultCollection(); - response.setResultsByTransaction(results); response.setErrorCode(error.code()); } response.setThrottleTimeMs(throttleTimeMs); diff --git a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java index 29518c251f2b2..645a03038a8c5 100644 --- a/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java +++ b/clients/src/main/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponse.java @@ -71,7 +71,7 @@ public void maybeSetThrottleTimeMs(int throttleTimeMs) { public Map> errors() { Map> errorsMap = new HashMap<>(); - if (this.data.resultsByTopicV3AndBelow().size() != 0) { + if (!this.data.resultsByTopicV3AndBelow().isEmpty()) { errorsMap.put(V3_AND_BELOW_TXN_ID, errorsForTransaction(this.data.resultsByTopicV3AndBelow())); } @@ -135,7 +135,7 @@ public Map errorCounts() { List allErrors = new ArrayList<>(); // If we are not using this field, we have request 4 or later - if (this.data.resultsByTopicV3AndBelow().size() == 0) { + if (this.data.resultsByTopicV3AndBelow().isEmpty()) { allErrors.add(Errors.forCode(data.errorCode())); } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java index 6f81dad0279cc..3b2dbee332c88 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/AddPartitionsToTxnResponseTest.java @@ -135,5 +135,10 @@ public void testBatchedErrors() { assertEquals(txn1Errors, errorsForTransaction(response.getTransactionTopicResults("txn1"))); assertEquals(txn2Errors, errorsForTransaction(response.getTransactionTopicResults("txn2"))); + + Map> expectedErrors = new HashMap<>(); + expectedErrors.put("txn1", txn1Errors); + expectedErrors.put("txn2", txn2Errors); + assertEquals(expectedErrors, response.errors()); } } diff --git a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java index 7c059b41c7677..7b0ca0d233e4c 100644 --- a/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java +++ b/clients/src/test/java/org/apache/kafka/common/requests/RequestResponseTest.java @@ -2609,13 +2609,13 @@ private AddPartitionsToTxnRequest createAddPartitionsToTxnRequest(short version) singletonList(new TopicPartition("topic", 73))).build(version); } else { AddPartitionsToTxnTransactionCollection transactions = new AddPartitionsToTxnTransactionCollection( - singletonList(new AddPartitionsToTxnTransaction() + singletonList(new AddPartitionsToTxnTransaction() .setTransactionalId("tid") .setProducerId(21L) .setProducerEpoch((short) 42) .setVerifyOnly(false) .setTopics(new AddPartitionsToTxnTopicCollection( - singletonList(new AddPartitionsToTxnTopic() + singletonList(new AddPartitionsToTxnTopic() .setName("topic") .setPartitions(Collections.singletonList(73))).iterator()))) .iterator()); diff --git a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala index 73ebb7603391d..02142f938a878 100644 --- a/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala +++ b/core/src/main/scala/kafka/coordinator/transaction/TransactionCoordinator.scala @@ -329,23 +329,24 @@ class TransactionCoordinator(txnConfig: TransactionConfig, partitions: collection.Set[TopicPartition], responseCallback: VerifyPartitionsCallback): Unit = { if (transactionalId == null || transactionalId.isEmpty) { - debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request") + debug(s"Returning ${Errors.INVALID_REQUEST} error code to client for $transactionalId's AddPartitions request for verification") responseCallback(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitions.map(_ -> Errors.INVALID_REQUEST).toMap.asJava)) } else { val result: ApiResult[(Int, TransactionMetadata)] = getTransactionMetadata(transactionalId, producerId, producerEpoch, partitions) result match { case Left(err) => - debug(s"Returning $err error code to client for $transactionalId's AddPartitions request") + debug(s"Returning $err error code to client for $transactionalId's AddPartitions request for verification") responseCallback(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, partitions.map(_ -> err).toMap.asJava)) case Right((_, txnMetadata)) => - val txnMetadataPartitions = txnMetadata.topicPartitions - val addedPartitions = partitions.intersect(txnMetadataPartitions) - val nonAddedPartitions = partitions.diff(txnMetadataPartitions) val errors = mutable.Map[TopicPartition, Errors]() - addedPartitions.foreach(errors.put(_, Errors.NONE)) - nonAddedPartitions.foreach(errors.put(_, Errors.INVALID_TXN_STATE)) + partitions.foreach { tp => + if (txnMetadata.topicPartitions.contains(tp)) + errors.put(tp, Errors.NONE) + else + errors.put(tp, Errors.INVALID_TXN_STATE) + } responseCallback(AddPartitionsToTxnResponse.resultForTransaction(transactionalId, errors.asJava)) } } @@ -378,11 +379,9 @@ class TransactionCoordinator(txnConfig: TransactionConfig, } private def getTransactionMetadata(transactionalId: String, - producerId: Long, - producerEpoch: Short, - partitions: collection.Set[TopicPartition]): ApiResult[(Int, TransactionMetadata)] = { - // try to update the transaction metadata and append the updated metadata to txn log; - // if there is no such metadata treat it as invalid producerId mapping error. + producerId: Long, + producerEpoch: Short, + partitions: collection.Set[TopicPartition]): ApiResult[(Int, TransactionMetadata)] = { txnManager.getTransactionState(transactionalId).flatMap { case None => Left(Errors.INVALID_PRODUCER_ID_MAPPING) diff --git a/core/src/main/scala/kafka/server/KafkaApis.scala b/core/src/main/scala/kafka/server/KafkaApis.scala index c489658e33880..7f15e7c4ef3d9 100644 --- a/core/src/main/scala/kafka/server/KafkaApis.scala +++ b/core/src/main/scala/kafka/server/KafkaApis.scala @@ -2420,7 +2420,7 @@ class KafkaApis(val requestChannel: RequestChannel, def addResultAndMaybeSendResponse(result: AddPartitionsToTxnResult): Unit = { val canSend = responses.synchronized { responses.add(result) - responses.size() == txns.size() + responses.size == txns.size } if (canSend) { requestHelper.sendResponseMaybeThrottle(request, createResponse) diff --git a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala index 0b3ad663333f1..fc84244cf2197 100644 --- a/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala +++ b/core/src/test/scala/unit/kafka/coordinator/transaction/TransactionCoordinatorTest.scala @@ -64,7 +64,6 @@ class TransactionCoordinatorTest { val transactionStatePartitionCount = 1 var result: InitProducerIdResult = _ var error: Errors = Errors.NONE - var errors: Map[TopicPartition, Errors] = _ private def mockPidGenerator(): Unit = { when(pidGenerator.generateProducerId()).thenAnswer(_ => { @@ -318,6 +317,11 @@ class TransactionCoordinatorTest { @Test def shouldRespondWithErrorsNoneOnAddPartitionWhenOngoingVerifyOnlyAndPartitionsTheSame(): Unit = { + var errors: Map[TopicPartition, Errors] = Map.empty + def verifyPartitionsInTxnCallback(result: AddPartitionsToTxnResult): Unit = { + errors = AddPartitionsToTxnResponse.errorsForTransaction(result.topicResults()).asScala.toMap + } + when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Ongoing, partitions, 0, 0))))) @@ -329,10 +333,14 @@ class TransactionCoordinatorTest { @Test def shouldRespondWithInvalidTxnStateWhenVerifyOnlyAndPartitionNotPresent(): Unit = { + var errors: Map[TopicPartition, Errors] = Map.empty + def verifyPartitionsInTxnCallback(result: AddPartitionsToTxnResult): Unit = { + errors = AddPartitionsToTxnResponse.errorsForTransaction(result.topicResults()).asScala.toMap + } + when(transactionManager.getTransactionState(ArgumentMatchers.eq(transactionalId))) .thenReturn(Right(Some(CoordinatorEpochAndTxnMetadata(coordinatorEpoch, new TransactionMetadata(transactionalId, 0, 0, 0, RecordBatch.NO_PRODUCER_EPOCH, 0, Empty, partitions, 0, 0))))) - val extraPartitions = partitions ++ Set(new TopicPartition("topic2", 0)) @@ -1213,8 +1221,4 @@ class TransactionCoordinatorTest { def errorsCallback(ret: Errors): Unit = { error = ret } - - def verifyPartitionsInTxnCallback(result: AddPartitionsToTxnResult): Unit = { - errors = AddPartitionsToTxnResponse.errorsForTransaction(result.topicResults()).asScala.toMap - } } diff --git a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala index f12b266bc9c05..5673315cf31e6 100644 --- a/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala +++ b/core/src/test/scala/unit/kafka/server/AddPartitionsToTxnRequestServerTest.scala @@ -72,7 +72,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { producerId, producerEpoch, List(createdTopicPartition, nonExistentTopic).asJava - ).build() + ).build(version) } else { val topics = new AddPartitionsToTxnTopicCollection() topics.add(new AddPartitionsToTxnTopic() @@ -89,7 +89,7 @@ class AddPartitionsToTxnRequestServerTest extends BaseRequestTest { .setProducerEpoch(producerEpoch) .setVerifyOnly(false) .setTopics(topics)) - AddPartitionsToTxnRequest.Builder.forBroker(transactions).build() + AddPartitionsToTxnRequest.Builder.forBroker(transactions).build(version) } val leaderId = brokers.head.config.brokerId diff --git a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala index d3cb4eb77788e..a142adece0c86 100644 --- a/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala +++ b/core/src/test/scala/unit/kafka/server/KafkaApisTest.scala @@ -2040,11 +2040,10 @@ class KafkaApisTest { } @Test - def testBatchedRequest(): Unit = { + def testBatchedAddPartitionsToTxnRequest(): Unit = { val topic = "topic" addTopicToMetadataCache(topic, numPartitions = 2) - val capturedResponse: ArgumentCaptor[AddPartitionsToTxnResponse] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResponse]) val responseCallback: ArgumentCaptor[Errors => Unit] = ArgumentCaptor.forClass(classOf[Errors => Unit]) val verifyPartitionsCallback: ArgumentCaptor[AddPartitionsToTxnResult => Unit] = ArgumentCaptor.forClass(classOf[AddPartitionsToTxnResult => Unit]) @@ -2104,12 +2103,7 @@ class KafkaApisTest { createKafkaApis().handleAddPartitionsToTxnRequest(request, requestLocal) - verify(requestChannel).sendResponse( - ArgumentMatchers.eq(request), - capturedResponse.capture(), - ArgumentMatchers.eq(None) - ) - val response = capturedResponse.getValue + val response = verifyNoThrottling[AddPartitionsToTxnResponse](request) val expectedErrors = Map( transactionalId1 -> Collections.singletonMap(tp0, Errors.NONE),