diff --git a/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/DataWriterMRJob.java b/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/DataWriterMRJob.java index 19c1b3af506..52f39800d39 100644 --- a/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/DataWriterMRJob.java +++ b/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/DataWriterMRJob.java @@ -34,6 +34,7 @@ import static com.linkedin.venice.vpj.VenicePushJobConstants.PARTITION_COUNT; import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_DUAL_WRITE_TARGET_REGIONS; import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_EXTERNAL_STORAGE_PROP_PREFIX; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_PROP_PREFIX; import static com.linkedin.venice.vpj.VenicePushJobConstants.REDUCER_SPECULATIVE_EXECUTION_ENABLE; import static com.linkedin.venice.vpj.VenicePushJobConstants.REPUSH_TTL_ENABLE; import static com.linkedin.venice.vpj.VenicePushJobConstants.REPUSH_TTL_POLICY; @@ -164,6 +165,10 @@ private void setupDefaultJobConf(JobConf conf, PushJobSetting pushJobSetting, Ve if (key.startsWith(PUSH_JOB_EXTERNAL_STORAGE_PROP_PREFIX)) { conf.set(key, props.getString(key)); } + // The factory receives these properties when it is initialized in the executor task. + if (key.startsWith(PUSH_JOB_WRITER_HOOK_PROP_PREFIX)) { + conf.set(key, props.getString(key)); + } } conf.set(PUSH_JOB_DUAL_WRITE_TARGET_REGIONS, String.join(",", pushJobSetting.dualWriteTargetRegions)); conf.setBoolean(ALLOW_DUPLICATE_KEY, pushJobSetting.isDuplicateKeyAllowed); diff --git a/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriter.java b/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriter.java index e3fefb15a59..62fd41dc631 100644 --- a/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriter.java +++ b/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriter.java @@ -22,6 +22,7 @@ import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_EXTERNAL_STORAGE_BATCHPUT_RETRY_BACKOFF_MS; import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_EXTERNAL_STORAGE_BATCH_SIZE; import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_EXTERNAL_STORAGE_WRITER_CLASS; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_FACTORY_CLASS; import static com.linkedin.venice.vpj.VenicePushJobConstants.RMD_SCHEMA_DIR; import static com.linkedin.venice.vpj.VenicePushJobConstants.RMD_SCHEMA_ID_PROP; import static com.linkedin.venice.vpj.VenicePushJobConstants.RMD_SCHEMA_PROP; @@ -66,6 +67,7 @@ import com.linkedin.venice.utils.ByteUtils; import com.linkedin.venice.utils.DictionaryUtils; import com.linkedin.venice.utils.PartitionUtils; +import com.linkedin.venice.utils.ReflectUtils; import com.linkedin.venice.utils.SystemTime; import com.linkedin.venice.utils.Time; import com.linkedin.venice.utils.Utils; @@ -82,6 +84,7 @@ import com.linkedin.venice.writer.PutMetadata; import com.linkedin.venice.writer.VeniceWriter; import com.linkedin.venice.writer.VeniceWriterFactory; +import com.linkedin.venice.writer.VeniceWriterHook; import com.linkedin.venice.writer.VeniceWriterOptions; import java.io.Closeable; import java.io.IOException; @@ -249,6 +252,7 @@ public int getValueSchemaId() { private AbstractVeniceWriter veniceWriter = null; private VeniceWriter mainWriter = null; private ComplexVeniceWriter[] childWriters = null; + private VeniceWriterHook writerHook = null; private int valueSchemaId = -1; private int rmdSchemaId = -1; @@ -562,7 +566,7 @@ protected AbstractVeniceWriter createBasicVeniceWriter() VenicePartitioner partitioner = PartitionUtils.getVenicePartitioner(props); String topicName = props.getString(TOPIC_PROP); - VeniceWriterOptions options = + VeniceWriterOptions.Builder optionsBuilder = new VeniceWriterOptions.Builder(topicName).setKeyPayloadSerializer(new DefaultSerializer()) .setValuePayloadSerializer(new DefaultSerializer()) .setWriteComputePayloadSerializer(new DefaultSerializer()) @@ -571,8 +575,11 @@ protected AbstractVeniceWriter createBasicVeniceWriter() .setTime(SystemTime.INSTANCE) .setPartitionCount(getPartitionCount()) .setPartitioner(partitioner) - .setMaxRecordSizeBytes(Integer.parseInt(maxRecordSizeBytesStr)) - .build(); + .setMaxRecordSizeBytes(Integer.parseInt(maxRecordSizeBytesStr)); + if (writerHook != null) { + optionsBuilder.setWriterHook(writerHook); + } + VeniceWriterOptions options = optionsBuilder.build(); String flatViewConfigMapString = props.getString(PUSH_JOB_VIEW_CONFIGS, ""); AbstractVeniceWriter baseWriter; if (!flatViewConfigMapString.isEmpty()) { @@ -971,6 +978,7 @@ protected void configureTask(VeniceProperties props) { } initStorageQuotaFields(props); initIncrementalPushThrottlers(props); + initWriterHookFactory(); /** * A dummy background task that reports progress every 5 minutes. */ @@ -1021,6 +1029,48 @@ protected void configureTask(VeniceProperties props) { }); } + private void initWriterHookFactory() { + String factoryClassName = props.getString(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, "").trim(); + if (factoryClassName.isEmpty()) { + return; + } + + VeniceWriterHookFactory factory = loadWriterHookFactory(factoryClassName); + String topicName = props.getString(TOPIC_PROP); + VeniceWriterHook hook = factory.createWriterHook(Version.parseStoreFromKafkaTopicName(topicName), props); + if (hook == null) { + throw new VeniceException( + VeniceWriterHookFactory.class.getSimpleName() + " '" + factoryClassName + "' returned a null hook"); + } + this.writerHook = hook; + } + + private VeniceWriterHookFactory loadWriterHookFactory(String className) { + Class loadedClass; + try { + loadedClass = ReflectUtils.loadClass(className); + } catch (Exception e) { + throw new VeniceException( + "Failed to load " + VeniceWriterHookFactory.class.getSimpleName() + " class '" + className + "'", + e); + } + Class factoryClass; + try { + factoryClass = loadedClass.asSubclass(VeniceWriterHookFactory.class); + } catch (ClassCastException e) { + throw new VeniceException( + "Configured class '" + className + "' does not implement " + VeniceWriterHookFactory.class.getName(), + e); + } + try { + return ReflectUtils.callConstructor(factoryClass, new Class[0], new Object[0]); + } catch (Exception e) { + throw new VeniceException( + "Failed to instantiate " + VeniceWriterHookFactory.class.getSimpleName() + " '" + className + "'", + e); + } + } + private void initStorageQuotaFields(VeniceProperties props) { Long storeStorageQuota = props.containsKey(STORAGE_QUOTA_PROP) ? props.getLong(STORAGE_QUOTA_PROP) : null; inputStorageQuotaTracker = new InputStorageQuotaTracker(storeStorageQuota); diff --git a/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/VeniceWriterHookFactory.java b/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/VeniceWriterHookFactory.java new file mode 100644 index 00000000000..c8e8f5c8402 --- /dev/null +++ b/clients/venice-push-job/src/main/java/com/linkedin/venice/hadoop/task/datawriter/VeniceWriterHookFactory.java @@ -0,0 +1,25 @@ +package com.linkedin.venice.hadoop.task.datawriter; + +import com.linkedin.venice.utils.VeniceProperties; +import com.linkedin.venice.writer.VeniceWriterHook; + + +/** + * Optional VPJ executor-side factory for the hook attached to the primary data {@code VeniceWriter}. + * + *

Implementations must have a public no-arg constructor. VPJ initializes the factory inside the executor task JVM + * and invokes {@link #createWriterHook(String, VeniceProperties)} exactly once per partition writer. + * + *

The hook is attached only to the primary data writer. It is not attached to control-message, + * heartbeat, or materialized-view child writers. + */ +public interface VeniceWriterHookFactory { + /** + * Creates the hook for a VPJ partition writer. + * + * @param storeName the destination Venice store name + * @param taskProperties the executor task properties, including any {@code push.job.writer.hook.*} settings + * @return the non-null hook to attach to the primary data writer + */ + VeniceWriterHook createWriterHook(String storeName, VeniceProperties taskProperties); +} diff --git a/clients/venice-push-job/src/main/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJob.java b/clients/venice-push-job/src/main/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJob.java index cbc44687807..5e6cb013752 100644 --- a/clients/venice-push-job/src/main/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJob.java +++ b/clients/venice-push-job/src/main/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJob.java @@ -50,6 +50,7 @@ import static com.linkedin.venice.vpj.VenicePushJobConstants.PARTITION_COUNT; import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_DUAL_WRITE_TARGET_REGIONS; import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_EXTERNAL_STORAGE_PROP_PREFIX; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_PROP_PREFIX; import static com.linkedin.venice.vpj.VenicePushJobConstants.REPUSH_TTL_ENABLE; import static com.linkedin.venice.vpj.VenicePushJobConstants.REPUSH_TTL_POLICY; import static com.linkedin.venice.vpj.VenicePushJobConstants.REPUSH_TTL_START_TIMESTAMP; @@ -154,6 +155,7 @@ import org.apache.spark.sql.types.StructType; import org.apache.spark.util.AccumulatorV2; import org.apache.spark.util.LongAccumulator; +import scala.collection.JavaConverters; /** @@ -236,6 +238,9 @@ private void setupDefaultSparkSessionForDataWriterJob(PushJobSetting pushJobSett sparkContext.setCallSite(jobGroupId); RuntimeConfig jobConf = sparkSession.conf(); + new ArrayList<>(JavaConverters.mapAsJavaMap(jobConf.getAll()).keySet()).stream() + .filter(key -> key.startsWith(PUSH_JOB_WRITER_HOOK_PROP_PREFIX)) + .forEach(jobConf::unset); setupCommonSparkConf(props, jobConf, pushJobSetting); jobConf.set(BATCH_NUM_BYTES_PROP, pushJobSetting.batchNumBytes); jobConf.set(TOPIC_PROP, pushJobSetting.topic); @@ -352,6 +357,10 @@ private void setupDefaultSparkSessionForDataWriterJob(PushJobSetting pushJobSett if (key.startsWith(PUSH_JOB_EXTERNAL_STORAGE_PROP_PREFIX)) { jobConf.set(key, props.getString(key)); } + // The factory receives these properties when it is initialized in the executor task. + if (key.startsWith(PUSH_JOB_WRITER_HOOK_PROP_PREFIX)) { + jobConf.set(key, props.getString(key)); + } } // Forward the DUAL_WRITE target-region list resolved by the VPJ driver (one entry per region whose // store-level storage mode is DUAL_WRITE) so the partition writer's gating predicate and per-region diff --git a/clients/venice-push-job/src/main/java/com/linkedin/venice/vpj/VenicePushJobConstants.java b/clients/venice-push-job/src/main/java/com/linkedin/venice/vpj/VenicePushJobConstants.java index 32b631004a1..4e36e6f2a06 100644 --- a/clients/venice-push-job/src/main/java/com/linkedin/venice/vpj/VenicePushJobConstants.java +++ b/clients/venice-push-job/src/main/java/com/linkedin/venice/vpj/VenicePushJobConstants.java @@ -513,6 +513,19 @@ private VenicePushJobConstants() { /** Enables Spark's pre-write quota check. Disabled by default. */ public static final String SPARK_PRE_WRITE_QUOTA_CHECK = "spark.pre.write.quota.check"; + /** + * Namespace for the optional VPJ primary-data-writer hook factory. Every property under this prefix is + * forwarded to executor task properties and passed to the factory when it is initialized in the task JVM. + */ + public static final String PUSH_JOB_WRITER_HOOK_PROP_PREFIX = "push.job.writer.hook."; + + /** + * Fully-qualified class name of the optional + * {@code com.linkedin.venice.hadoop.task.datawriter.VeniceWriterHookFactory}. The class must have a public + * no-arg constructor. When absent or empty, VPJ creates writers exactly as before, without a writer hook. + */ + public static final String PUSH_JOB_WRITER_HOOK_FACTORY_CLASS = PUSH_JOB_WRITER_HOOK_PROP_PREFIX + "factory.class"; + /** * Namespace for the external-storage dual-write subsystem. Every property whose key starts with this * prefix is forwarded verbatim from the VPJ driver into the Spark executor's {@code RuntimeConfig} so diff --git a/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/TestDataWriterMRJob.java b/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/TestDataWriterMRJob.java index cee42aca306..b72e1b3f160 100644 --- a/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/TestDataWriterMRJob.java +++ b/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/mapreduce/datawriter/jobs/TestDataWriterMRJob.java @@ -1,5 +1,7 @@ package com.linkedin.venice.hadoop.mapreduce.datawriter.jobs; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_FACTORY_CLASS; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_PROP_PREFIX; import static com.linkedin.venice.vpj.VenicePushJobConstants.WRITER_RMD_SCHEMA_STRING_PROP; import static com.linkedin.venice.vpj.VenicePushJobConstants.WRITER_VALUE_SCHEMA_STRING_PROP; import static org.mockito.Mockito.doReturn; @@ -11,9 +13,14 @@ import com.linkedin.venice.etl.ETLValueSchemaTransformation; import com.linkedin.venice.hadoop.PushJobSetting; +import com.linkedin.venice.hadoop.VenicePushJob; +import com.linkedin.venice.partitioner.DefaultVenicePartitioner; import com.linkedin.venice.schema.rmd.RmdSchemaGenerator; import com.linkedin.venice.utils.TestWriteUtils; +import com.linkedin.venice.utils.VeniceProperties; import java.io.IOException; +import java.util.Collections; +import java.util.Properties; import org.apache.avro.Schema; import org.apache.hadoop.conf.Configuration; import org.apache.hadoop.fs.FileSystem; @@ -174,6 +181,30 @@ public void testSetupInputFormatConfOmitsWriterSchemasWhenNotProjecting() { assertNull(jobConf.get(WRITER_RMD_SCHEMA_STRING_PROP)); } + @Test + public void testConfigureForwardsWriterHookPropertiesToTasks() { + String writerHookSetting = PUSH_JOB_WRITER_HOOK_PROP_PREFIX + "test.setting"; + Properties properties = new Properties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, "com.example.WriterHookFactory"); + properties.setProperty(writerHookSetting, "test-value"); + + PushJobSetting setting = avroProjectionPushJobSetting(); + setting.jobId = "test-job"; + setting.topic = "testStore_v1"; + setting.pushDestinationPubsubBroker = "test-broker"; + setting.partitionerClass = DefaultVenicePartitioner.class.getName(); + setting.dualWriteTargetRegions = Collections.emptyList(); + setting.partitionCount = 1; + setting.vpjEntryClass = VenicePushJob.class; + + CapturingDataWriterMRJob mrJob = new CapturingDataWriterMRJob(); + mrJob.configure(new VeniceProperties(properties), setting); + + Assert + .assertEquals(mrJob.configuredJobConf.get(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS), "com.example.WriterHookFactory"); + Assert.assertEquals(mrJob.configuredJobConf.get(writerHookSetting), "test-value"); + } + private PushJobSetting avroProjectionPushJobSetting() { PushJobSetting setting = new PushJobSetting(); setting.isSourceKafka = false; @@ -186,4 +217,14 @@ private PushJobSetting avroProjectionPushJobSetting() { setting.inputDataSchemaString = TestWriteUtils.STRING_TO_NAME_RECORD_V2_SCHEMA.toString(); return setting; } + + private static class CapturingDataWriterMRJob extends DataWriterMRJob { + private JobConf configuredJobConf; + + @Override + void setupMRConf(JobConf jobConf, PushJobSetting pushJobSetting, VeniceProperties props) { + configuredJobConf = jobConf; + super.setupMRConf(jobConf, pushJobSetting, props); + } + } } diff --git a/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriterHookFactoryTest.java b/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriterHookFactoryTest.java new file mode 100644 index 00000000000..5e3bd2d41e0 --- /dev/null +++ b/clients/venice-push-job/src/test/java/com/linkedin/venice/hadoop/task/datawriter/AbstractPartitionWriterHookFactoryTest.java @@ -0,0 +1,237 @@ +package com.linkedin.venice.hadoop.task.datawriter; + +import static com.linkedin.venice.vpj.VenicePushJobConstants.PARTITION_COUNT; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_FACTORY_CLASS; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_PROP_PREFIX; +import static com.linkedin.venice.vpj.VenicePushJobConstants.TELEMETRY_MESSAGE_INTERVAL; +import static com.linkedin.venice.vpj.VenicePushJobConstants.TOPIC_PROP; +import static com.linkedin.venice.vpj.VenicePushJobConstants.VALUE_SCHEMA_ID_PROP; +import static org.mockito.ArgumentMatchers.any; +import static org.mockito.Mockito.mock; +import static org.mockito.Mockito.verify; +import static org.mockito.Mockito.when; +import static org.testng.Assert.assertEquals; +import static org.testng.Assert.assertNull; +import static org.testng.Assert.assertSame; +import static org.testng.Assert.assertTrue; + +import com.linkedin.venice.ConfigKeys; +import com.linkedin.venice.exceptions.VeniceException; +import com.linkedin.venice.hadoop.engine.EngineTaskConfigProvider; +import com.linkedin.venice.meta.MaterializedViewParameters; +import com.linkedin.venice.meta.ViewConfig; +import com.linkedin.venice.meta.ViewConfigImpl; +import com.linkedin.venice.partitioner.DefaultVenicePartitioner; +import com.linkedin.venice.utils.VeniceProperties; +import com.linkedin.venice.views.MaterializedView; +import com.linkedin.venice.views.ViewUtils; +import com.linkedin.venice.writer.ComplexVeniceWriter; +import com.linkedin.venice.writer.VeniceWriter; +import com.linkedin.venice.writer.VeniceWriterFactory; +import com.linkedin.venice.writer.VeniceWriterHook; +import com.linkedin.venice.writer.VeniceWriterOptions; +import java.io.IOException; +import java.util.Collections; +import java.util.Properties; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.mockito.ArgumentCaptor; +import org.testng.Assert; +import org.testng.annotations.BeforeMethod; +import org.testng.annotations.Test; + + +public class AbstractPartitionWriterHookFactoryTest { + private static final String TOPIC_NAME = "testStore_v1"; + private static final int TASK_ID = 3; + private static final int PARTITIONS = 8; + + @BeforeMethod + public void resetFactories() { + RecordingFactory.reset(); + } + + @Test + public void testUnconfiguredWriterHasNoHook() throws IOException { + TestablePartitionWriter partitionWriter = configureWriter(createBaseProperties()); + try { + VeniceWriterOptions options = createAndCaptureMainWriterOptions(partitionWriter); + assertNull(options.getWriterHook()); + } finally { + partitionWriter.close(); + } + } + + @Test + public void testWhitespaceOnlyFactoryConfigurationHasNoHook() throws IOException { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, " \t "); + + TestablePartitionWriter partitionWriter = configureWriter(properties); + try { + VeniceWriterOptions options = createAndCaptureMainWriterOptions(partitionWriter); + assertNull(options.getWriterHook()); + assertEquals(RecordingFactory.CREATE_COUNT.get(), 0); + } finally { + partitionWriter.close(); + } + } + + @Test + public void testConfiguredFactoryInjectsHookAndReceivesStoreAndTaskProperties() throws IOException { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, " " + RecordingFactory.class.getName() + " "); + properties.setProperty(PUSH_JOB_WRITER_HOOK_PROP_PREFIX + "test.setting", "test-value"); + + TestablePartitionWriter partitionWriter = configureWriter(properties); + try { + VeniceWriterOptions options = createAndCaptureMainWriterOptions(partitionWriter); + assertSame(options.getWriterHook(), RecordingFactory.HOOK); + assertEquals(RecordingFactory.CREATE_COUNT.get(), 1); + assertEquals(RecordingFactory.STORE_NAME.get(), "testStore"); + assertEquals( + RecordingFactory.TASK_PROPERTIES.get().getString(PUSH_JOB_WRITER_HOOK_PROP_PREFIX + "test.setting"), + "test-value"); + } finally { + partitionWriter.close(); + } + } + + @Test + public void testHookIsNotAttachedToMaterializedViewChildWriter() throws IOException { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, RecordingFactory.class.getName()); + MaterializedViewParameters.Builder viewParametersBuilder = new MaterializedViewParameters.Builder("testView"); + viewParametersBuilder.setPartitionCount(4); + viewParametersBuilder.setPartitioner(DefaultVenicePartitioner.class.getName()); + ViewConfig viewConfig = new ViewConfigImpl(MaterializedView.class.getName(), viewParametersBuilder.build()); + properties.setProperty( + ConfigKeys.PUSH_JOB_VIEW_CONFIGS, + ViewUtils.flatViewConfigMapString(Collections.singletonMap("testView", viewConfig))); + + TestablePartitionWriter partitionWriter = configureWriter(properties); + VeniceWriterFactory writerFactory = mock(VeniceWriterFactory.class); + when(writerFactory.createVeniceWriter(any())).thenReturn(mock(VeniceWriter.class)); + when(writerFactory.createComplexVeniceWriter(any())).thenReturn(mock(ComplexVeniceWriter.class)); + partitionWriter.setVeniceWriterFactory(writerFactory); + try { + partitionWriter.createBasicVeniceWriter(); + + ArgumentCaptor mainOptions = ArgumentCaptor.forClass(VeniceWriterOptions.class); + ArgumentCaptor childOptions = ArgumentCaptor.forClass(VeniceWriterOptions.class); + verify(writerFactory).createVeniceWriter(mainOptions.capture()); + verify(writerFactory).createComplexVeniceWriter(childOptions.capture()); + assertSame(mainOptions.getValue().getWriterHook(), RecordingFactory.HOOK); + assertNull(childOptions.getValue().getWriterHook()); + } finally { + partitionWriter.close(); + } + } + + @Test + public void testFactoryReturningNullHookFails() { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, NullHookFactory.class.getName()); + + VeniceException exception = Assert.expectThrows(VeniceException.class, () -> configureWriter(properties)); + assertTrue(exception.getMessage().contains("returned a null hook")); + } + + @Test + public void testInvalidFactoryClassNameFailsClearly() { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, "com.example.DoesNotExist"); + + VeniceException exception = Assert.expectThrows(VeniceException.class, () -> configureWriter(properties)); + assertTrue(exception.getMessage().contains("Failed to load VeniceWriterHookFactory class")); + } + + @Test + public void testConfiguredClassMustImplementFactory() { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, String.class.getName()); + + VeniceException exception = Assert.expectThrows(VeniceException.class, () -> configureWriter(properties)); + assertTrue(exception.getMessage().contains("does not implement " + VeniceWriterHookFactory.class.getName())); + } + + @Test + public void testFactoryMustHavePublicNoArgConstructor() { + Properties properties = createBaseProperties(); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, FactoryWithoutNoArgConstructor.class.getName()); + + VeniceException exception = Assert.expectThrows(VeniceException.class, () -> configureWriter(properties)); + assertTrue(exception.getMessage().contains("Failed to instantiate VeniceWriterHookFactory")); + } + + private TestablePartitionWriter configureWriter(Properties properties) { + EngineTaskConfigProvider taskConfigProvider = mock(EngineTaskConfigProvider.class); + when(taskConfigProvider.getJobProps()).thenReturn(properties); + when(taskConfigProvider.getTaskId()).thenReturn(TASK_ID); + TestablePartitionWriter partitionWriter = new TestablePartitionWriter(); + partitionWriter.configure(taskConfigProvider); + return partitionWriter; + } + + private VeniceWriterOptions createAndCaptureMainWriterOptions(TestablePartitionWriter partitionWriter) { + VeniceWriterFactory writerFactory = mock(VeniceWriterFactory.class); + when(writerFactory.createVeniceWriter(any())).thenReturn(mock(VeniceWriter.class)); + partitionWriter.setVeniceWriterFactory(writerFactory); + partitionWriter.createBasicVeniceWriter(); + + ArgumentCaptor options = ArgumentCaptor.forClass(VeniceWriterOptions.class); + verify(writerFactory).createVeniceWriter(options.capture()); + return options.getValue(); + } + + private Properties createBaseProperties() { + Properties properties = new Properties(); + properties.setProperty(PARTITION_COUNT, Integer.toString(PARTITIONS)); + properties.setProperty(TOPIC_PROP, TOPIC_NAME); + properties.setProperty(VALUE_SCHEMA_ID_PROP, "1"); + properties.setProperty(TELEMETRY_MESSAGE_INTERVAL, "10000"); + properties.setProperty(ConfigKeys.PARTITIONER_CLASS, DefaultVenicePartitioner.class.getName()); + return properties; + } + + private static class TestablePartitionWriter extends AbstractPartitionWriter { + } + + public static class RecordingFactory implements VeniceWriterHookFactory { + private static final VeniceWriterHook HOOK = (operationType, keySizeBytes, valueSizeBytes) -> {}; + private static final AtomicReference STORE_NAME = new AtomicReference<>(); + private static final AtomicReference TASK_PROPERTIES = new AtomicReference<>(); + private static final AtomicInteger CREATE_COUNT = new AtomicInteger(); + + private static void reset() { + STORE_NAME.set(null); + TASK_PROPERTIES.set(null); + CREATE_COUNT.set(0); + } + + @Override + public VeniceWriterHook createWriterHook(String storeName, VeniceProperties taskProperties) { + STORE_NAME.set(storeName); + TASK_PROPERTIES.set(taskProperties); + CREATE_COUNT.incrementAndGet(); + return HOOK; + } + } + + public static class NullHookFactory implements VeniceWriterHookFactory { + @Override + public VeniceWriterHook createWriterHook(String storeName, VeniceProperties taskProperties) { + return null; + } + } + + public static class FactoryWithoutNoArgConstructor implements VeniceWriterHookFactory { + public FactoryWithoutNoArgConstructor(String ignored) { + } + + @Override + public VeniceWriterHook createWriterHook(String storeName, VeniceProperties taskProperties) { + return RecordingFactory.HOOK; + } + } +} diff --git a/clients/venice-push-job/src/test/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJobTest.java b/clients/venice-push-job/src/test/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJobTest.java index 1c9b4b03e43..eeed110fc50 100644 --- a/clients/venice-push-job/src/test/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJobTest.java +++ b/clients/venice-push-job/src/test/java/com/linkedin/venice/spark/datawriter/jobs/AbstractDataWriterSparkJobTest.java @@ -17,6 +17,8 @@ import static com.linkedin.venice.spark.SparkConstants.VALUE_COLUMN_NAME; import static com.linkedin.venice.vpj.VenicePushJobConstants.DEFAULT_KEY_FIELD_PROP; import static com.linkedin.venice.vpj.VenicePushJobConstants.DEFAULT_VALUE_FIELD_PROP; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_FACTORY_CLASS; +import static com.linkedin.venice.vpj.VenicePushJobConstants.PUSH_JOB_WRITER_HOOK_PROP_PREFIX; import static com.linkedin.venice.vpj.VenicePushJobConstants.SPARK_NATIVE_INPUT_FORMAT_ENABLED; import static org.apache.spark.sql.types.DataTypes.BinaryType; import static org.apache.spark.sql.types.DataTypes.IntegerType; @@ -89,11 +91,15 @@ public void testConfigure() throws IOException { String dummyKafkaConfigValue = "dummy.kafka.config.value"; String dummyConfig = "some.dummy.config"; String dummyConfigValue = "some.dummy.config.value"; + String writerHookFactoryClass = "com.example.WriterHookFactory"; + String writerHookSetting = PUSH_JOB_WRITER_HOOK_PROP_PREFIX + "test.setting"; Properties properties = new Properties(); properties.setProperty(SPARK_SESSION_CONF_PREFIX + SPARK_APP_NAME_CONFIG, sparkAppNameOverride); properties.setProperty(KAFKA_CONFIG_PREFIX + dummyKafkaConfig, dummyKafkaConfigValue); properties.setProperty(SPARK_DATA_WRITER_CONF_PREFIX + dummyConfig, dummyConfigValue); + properties.setProperty(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS, writerHookFactoryClass); + properties.setProperty(writerHookSetting, "test-value"); try (DataWriterSparkJob dataWriterSparkJob = new DataWriterSparkJob()) { dataWriterSparkJob.configure(new VeniceProperties(properties), setting); @@ -107,6 +113,18 @@ public void testConfigure() throws IOException { // Properties with SPARK_DATA_WRITER_CONF_PREFIX should get applied after stripping the prefix Assert.assertEquals(jobConf.get(dummyConfig), dummyConfigValue); + + // VPJ writer-hook factory config should reach executor task properties. + Assert.assertEquals(jobConf.get(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS), writerHookFactoryClass); + Assert.assertEquals(jobConf.get(writerHookSetting), "test-value"); + } + + try (DataWriterSparkJob dataWriterSparkJob = new DataWriterSparkJob()) { + dataWriterSparkJob.configure(VeniceProperties.empty(), setting); + + RuntimeConfig jobConf = dataWriterSparkJob.getSparkSession().conf(); + Assert.assertTrue(jobConf.getOption(PUSH_JOB_WRITER_HOOK_FACTORY_CLASS).isEmpty()); + Assert.assertTrue(jobConf.getOption(writerHookSetting).isEmpty()); } }