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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -41,9 +41,15 @@ public class KinesisRecord implements Record<byte[]> {
private final Optional<String> key;
private final byte[] value;
private final HashMap<String, String> userProperties = new HashMap<>();

private final String sequenceNumber;
private final KinesisRecordProcessor recordProcessor;

public KinesisRecord(KinesisClientRecord record, String shardId, long millisBehindLatest,
Set<String> propertiesToInclude) {
Set<String> propertiesToInclude, KinesisRecordProcessor recordProcessor) {
this.key = Optional.of(record.partitionKey());
this.sequenceNumber = record.sequenceNumber();
this.recordProcessor = recordProcessor;
// encryption type can (annoyingly) be null, so we default to NONE
EncryptionType encType = EncryptionType.NONE;
if (record.encryptionType() != null) {
Expand Down Expand Up @@ -94,6 +100,16 @@ public byte[] getValue() {
return value;
}

@Override
public void ack() {
this.recordProcessor.updateSequenceNumberToCheckpoint(this.sequenceNumber);
}

@Override
public void fail() {
this.recordProcessor.failed();
}

public Map<String, String> getProperties() {
return userProperties;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,8 @@
import java.util.Set;
import java.util.concurrent.LinkedBlockingQueue;
import lombok.extern.slf4j.Slf4j;
import org.apache.pulsar.client.api.PulsarClientException;
import org.apache.pulsar.io.core.SourceContext;
import software.amazon.kinesis.exceptions.InvalidStateException;
import software.amazon.kinesis.exceptions.KinesisClientLibDependencyException;
import software.amazon.kinesis.exceptions.ShutdownException;
Expand All @@ -42,34 +44,43 @@ public class KinesisRecordProcessor implements ShardRecordProcessor {
private final long backoffTime;

private final LinkedBlockingQueue<KinesisRecord> queue;
private final SourceContext sourceContext;
private final Set<String> propertiesToInclude;

private long nextCheckpointTimeInNanos;
private String kinesisShardId;
private final Set<String> propertiesToInclude;
public KinesisRecordProcessor(LinkedBlockingQueue<KinesisRecord> queue, KinesisSourceConfig config) {
private volatile String sequenceNumberToCheckpoint = null;
private String lastCheckpointedSequenceNumber = null;

public KinesisRecordProcessor(LinkedBlockingQueue<KinesisRecord> queue, KinesisSourceConfig config,
SourceContext sourceContext) {
this.queue = queue;
this.checkpointInterval = config.getCheckpointInterval();
this.numRetries = config.getNumRetries();
this.backoffTime = config.getBackoffTime();
this.propertiesToInclude = config.getPropertiesToInclude();
this.sourceContext = sourceContext;
}

private void checkpoint(RecordProcessorCheckpointer checkpointer) {
log.info("Checkpointing shard " + kinesisShardId);
private void checkpoint(RecordProcessorCheckpointer checkpointer, String sequenceNumber) {
log.info("Checkpointing shard {} at sequence number {}", kinesisShardId, sequenceNumber);
for (int i = 0; i < numRetries; i++) {
try {
checkpointer.checkpoint();
checkpointer.checkpoint(sequenceNumber);
lastCheckpointedSequenceNumber = sequenceNumber;
break;
} catch (ShutdownException se) {
// Ignore checkpoint if the processor instance has been shutdown.
log.info("Caught shutdown exception, skipping checkpoint.", se);
sourceContext.fatal(se);
break;
} catch (InvalidStateException e) {
log.error("Cannot save checkpoint to the DynamoDB table.", e);
sourceContext.fatal(e);
break;
} catch (ThrottlingException | KinesisClientLibDependencyException e) {
// Back off and re-attempt checkpoint upon transient failures
if (i >= (numRetries - 1)) {
log.error("Checkpoint failed after " + (i + 1) + "attempts.", e);
log.error("Checkpoint failed after {} attempts.", (i + 1), e);
sourceContext.fatal(e);
break;
}
}
Expand All @@ -82,49 +93,66 @@ private void checkpoint(RecordProcessorCheckpointer checkpointer) {
}
}

public void updateSequenceNumberToCheckpoint(String sequenceNumber) {
this.sequenceNumberToCheckpoint = sequenceNumber;
}

public void failed() {
sourceContext.fatal(new PulsarClientException("Failed to process Kinesis records due send to pulsar topic"));
}

@Override
public void initialize(InitializationInput initializationInput) {
kinesisShardId = initializationInput.shardId();
log.info("Initializing KinesisRecordProcessor for shard {}. Config: checkpointInterval={}ms, numRetries={}, "
+ "backoffTime={}ms, propertiesToInclude={}",
kinesisShardId, checkpointInterval, numRetries, backoffTime, propertiesToInclude);
nextCheckpointTimeInNanos = System.nanoTime() + checkpointInterval;
}

@Override
public void processRecords(ProcessRecordsInput processRecordsInput) {

log.info("Processing " + processRecordsInput.records().size() + " records from " + kinesisShardId);
log.info("Processing {} records from {}", processRecordsInput.records().size(), kinesisShardId);
long millisBehindLatest = processRecordsInput.millisBehindLatest();

for (KinesisClientRecord record : processRecordsInput.records()) {
try {
queue.put(new KinesisRecord(record, this.kinesisShardId, millisBehindLatest, propertiesToInclude));
queue.put(new KinesisRecord(record, this.kinesisShardId, millisBehindLatest,
propertiesToInclude, this));
} catch (InterruptedException e) {
log.warn("unable to create KinesisRecord ", e);
}
}

// Checkpoint once every checkpoint interval.
if (System.nanoTime() > nextCheckpointTimeInNanos) {
checkpoint(processRecordsInput.checkpointer());
if (sequenceNumberToCheckpoint != null
&& !sequenceNumberToCheckpoint.equals(lastCheckpointedSequenceNumber)) {
checkpoint(processRecordsInput.checkpointer(), sequenceNumberToCheckpoint);
}
Comment thread
RobertIndie marked this conversation as resolved.
nextCheckpointTimeInNanos = System.nanoTime() + checkpointInterval;
}
}

@Override
public void leaseLost(LeaseLostInput leaseLostInput) {
log.info("lease lost, will terminate soon");
log.info("Lease lost for shard {} lastCheckPointedSequenceNumber {}, will terminate soon.",
kinesisShardId, lastCheckpointedSequenceNumber);
}

@Override
public void shardEnded(ShardEndedInput shardEndedInput) {
log.info("reached end of shard, will checkpoint");
checkpoint(shardEndedInput.checkpointer());
log.info("Reached end of shard {}, will checkpoint.", kinesisShardId);
if (sequenceNumberToCheckpoint != null) {
checkpoint(shardEndedInput.checkpointer(), sequenceNumberToCheckpoint);
}
}

@Override
public void shutdownRequested(ShutdownRequestedInput shutdownRequestedInput) {
log.info("Shutting down record processor for shard: " + kinesisShardId);
checkpoint(shutdownRequestedInput.checkpointer());
log.info("Shutdown requested for record processor on shard {}, will checkpoint.", kinesisShardId);
if (sequenceNumberToCheckpoint != null) {
checkpoint(shutdownRequestedInput.checkpointer(), sequenceNumberToCheckpoint);
}
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -19,21 +19,25 @@
package org.apache.pulsar.io.kinesis;

import java.util.concurrent.LinkedBlockingQueue;
import org.apache.pulsar.io.core.SourceContext;
import software.amazon.kinesis.processor.ShardRecordProcessor;
import software.amazon.kinesis.processor.ShardRecordProcessorFactory;

public class KinesisRecordProcessorFactory implements ShardRecordProcessorFactory {

private final LinkedBlockingQueue<KinesisRecord> queue;
private final KinesisSourceConfig config;
private final SourceContext sourceContext;
public KinesisRecordProcessorFactory(LinkedBlockingQueue<KinesisRecord> queue,
KinesisSourceConfig kinesisSourceConfig) {
KinesisSourceConfig kinesisSourceConfig,
SourceContext sourceContext) {
this.queue = queue;
this.config = kinesisSourceConfig;
this.sourceContext = sourceContext;
}

@Override
public ShardRecordProcessor shardRecordProcessor() {
return new KinesisRecordProcessor(queue, config);
return new KinesisRecordProcessor(queue, config, sourceContext);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -73,7 +73,7 @@ public void open(Map<String, Object> config, SourceContext sourceContext) throws
kinesisSourceConfig.getAwsCredentialPluginParam());

KinesisAsyncClient kClient = kinesisSourceConfig.buildKinesisAsyncClient(credentialsProvider);
recordProcessorFactory = new KinesisRecordProcessorFactory(queue, kinesisSourceConfig);
recordProcessorFactory = new KinesisRecordProcessorFactory(queue, kinesisSourceConfig, sourceContext);
configsBuilder = new ConfigsBuilder(kinesisSourceConfig.getAwsKinesisStreamName(),
kinesisSourceConfig.getApplicationName(),
kClient,
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,147 @@
/*
* Licensed to the Apache Software Foundation (ASF) under one
* or more contributor license agreements. See the NOTICE file
* distributed with this work for additional information
* regarding copyright ownership. The ASF licenses this file
* to you under the Apache License, Version 2.0 (the
* "License"); you may not use this file except in compliance
* with the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing,
* software distributed under the License is distributed on an
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
* KIND, either express or implied. See the License for the
* specific language governing permissions and limitations
* under the License.
*/
package org.apache.pulsar.io.kinesis;

import static org.mockito.Mockito.when;
import static org.testng.Assert.assertEquals;
import java.nio.ByteBuffer;
import java.time.Instant;
import java.util.Arrays;
import java.util.Collections;
import java.util.concurrent.LinkedBlockingQueue;
import org.apache.pulsar.io.core.SourceContext;
import org.mockito.Mockito;
import org.testng.annotations.BeforeMethod;
import org.testng.annotations.Test;
import software.amazon.awssdk.services.kinesis.model.EncryptionType;
import software.amazon.kinesis.lifecycle.events.ProcessRecordsInput;
import software.amazon.kinesis.processor.RecordProcessorCheckpointer;
import software.amazon.kinesis.retrieval.KinesisClientRecord;

public class KinesisRecordProcessorTest {

private KinesisSourceConfig config;
private SourceContext sourceContext;
private LinkedBlockingQueue<KinesisRecord> queue;
private KinesisRecordProcessor recordProcessor;
private RecordProcessorCheckpointer checkpointer;

@BeforeMethod
public void setup() {
config = Mockito.mock(KinesisSourceConfig.class);
sourceContext = Mockito.mock(SourceContext.class);
queue = new LinkedBlockingQueue<>();
checkpointer = Mockito.mock(RecordProcessorCheckpointer.class);

// Configure the mock config for the processor
when(config.getCheckpointInterval()).thenReturn(100L);
when(config.getNumRetries()).thenReturn(1);
when(config.getBackoffTime()).thenReturn(10L);
when(config.getPropertiesToInclude()).thenReturn(Collections.emptySet());

recordProcessor = new KinesisRecordProcessor(queue, config, sourceContext);
}

@Test
public void testCheckpointAfterAck() throws Exception {
// --- Setup: Prepare mock inputs ---
String seqNum1 = "seq-1";
String seqNum2 = "seq-2";
KinesisClientRecord kcr1 = createMockKinesisRecord(seqNum1);
KinesisClientRecord kcr2 = createMockKinesisRecord(seqNum2);
ProcessRecordsInput processRecordsInput = createMockProcessRecordsInput(kcr1, kcr2);

// --- Action 1: Process records ---
recordProcessor.processRecords(processRecordsInput);

// --- Assert 1: Records are in the queue ---
assertEquals(queue.size(), 2);

// --- Action 2: Simulate source reading and acking the first record ---
KinesisRecord recordFromQueue1 = queue.take();
recordFromQueue1.ack(); // This updates sequenceNumberToCheckpoint in the processor

// --- Action 3: Advance time and trigger checkpoint logic ---
Thread.sleep(config.getCheckpointInterval() + 50);
Comment thread
shibd marked this conversation as resolved.
recordProcessor.processRecords(createMockProcessRecordsInput()); // Empty input to trigger checkpoint

// --- Assert 3: Verify checkpoint was called with the correct sequence number ---
Mockito.verify(checkpointer, Mockito.times(1)).checkpoint(seqNum1);

// --- Action 4: Ack the second record ---
KinesisRecord recordFromQueue2 = queue.take();
recordFromQueue2.ack();

// --- Action 5: Trigger checkpoint again ---
Thread.sleep(config.getCheckpointInterval() + 50);
recordProcessor.processRecords(createMockProcessRecordsInput());

// --- Assert 5: Verify checkpoint was called with the new, correct sequence number ---
Mockito.verify(checkpointer, Mockito.times(1)).checkpoint(seqNum2);
}

@Test
public void testNoCheckpointWithoutAck() throws Exception {
// --- Setup ---
String seqNum1 = "seq-1";
KinesisClientRecord kcr1 = createMockKinesisRecord(seqNum1);
ProcessRecordsInput processRecordsInput = createMockProcessRecordsInput(kcr1);

// --- Action 1: Process a record ---
recordProcessor.processRecords(processRecordsInput);
assertEquals(queue.size(), 1);
queue.take(); // Simulate reading but not acking

// --- Action 2: Advance time and trigger checkpoint logic ---
Thread.sleep(config.getCheckpointInterval() + 50);
recordProcessor.processRecords(processRecordsInput);

// --- Assert 2: Verify checkpoint was NEVER called because no record was acked ---
Mockito.verify(checkpointer, Mockito.never()).checkpoint(Mockito.anyString());
}

@Test
public void testFailTriggersFatal() throws Exception {
KinesisClientRecord kcr1 = createMockKinesisRecord("seq-fail");
ProcessRecordsInput processRecordsInput = createMockProcessRecordsInput(kcr1);

recordProcessor.processRecords(processRecordsInput);
KinesisRecord recordToFail = queue.take();
recordToFail.fail();

Mockito.verify(sourceContext, Mockito.times(1)).fatal(Mockito.any(Exception.class));
}

private KinesisClientRecord createMockKinesisRecord(String sequenceNumber) {
KinesisClientRecord mockRecord = Mockito.mock(KinesisClientRecord.class);
when(mockRecord.partitionKey()).thenReturn("test-key");
when(mockRecord.sequenceNumber()).thenReturn(sequenceNumber);
when(mockRecord.approximateArrivalTimestamp()).thenReturn(Instant.now());
when(mockRecord.encryptionType()).thenReturn(EncryptionType.NONE);
when(mockRecord.data()).thenReturn(ByteBuffer.wrap("data".getBytes()));
return mockRecord;
}

private ProcessRecordsInput createMockProcessRecordsInput(KinesisClientRecord... records) {
ProcessRecordsInput input = Mockito.mock(ProcessRecordsInput.class);
when(input.records()).thenReturn(Arrays.asList(records));
when(input.checkpointer()).thenReturn(checkpointer);
return input;
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,8 @@ public void testAllPropertiesIncluded() {
KinesisRecord.MILLIS_BEHIND_LATEST
));

KinesisRecord kinesisRecord = new KinesisRecord(mockRecord, shardId, millisBehindLatest, propertiesToInclude);
KinesisRecord kinesisRecord = new KinesisRecord(mockRecord, shardId, millisBehindLatest,
propertiesToInclude, null);
Map<String, String> properties = kinesisRecord.getProperties();

assertEquals(properties.size(), 6);
Expand All @@ -85,7 +86,8 @@ public void testSomePropertiesIncluded() {
KinesisRecord.SEQUENCE_NUMBER
));

KinesisRecord kinesisRecord = new KinesisRecord(mockRecord, shardId, millisBehindLatest, propertiesToInclude);
KinesisRecord kinesisRecord = new KinesisRecord(mockRecord, shardId, millisBehindLatest,
propertiesToInclude, null);
Map<String, String> properties = kinesisRecord.getProperties();

assertEquals(properties.size(), 2);
Expand All @@ -102,7 +104,8 @@ public void testSomePropertiesIncluded() {
public void testNoPropertiesIncluded() {
Set<String> propertiesToInclude = Collections.emptySet();

KinesisRecord kinesisRecord = new KinesisRecord(mockRecord, shardId, millisBehindLatest, propertiesToInclude);
KinesisRecord kinesisRecord = new KinesisRecord(mockRecord, shardId, millisBehindLatest,
propertiesToInclude, null);
Map<String, String> properties = kinesisRecord.getProperties();

assertTrue(properties.isEmpty());
Expand Down
Loading