1717 */
1818package org .apache .beam .sdk .io .kafka ;
1919
20+ import static java .nio .charset .StandardCharsets .UTF_8 ;
2021import static org .apache .beam .vendor .guava .v32_1_2_jre .com .google .common .base .Preconditions .checkArgument ;
2122import static org .apache .beam .vendor .guava .v32_1_2_jre .com .google .common .base .Preconditions .checkState ;
2223import static org .apache .kafka .clients .consumer .ConsumerConfig .AUTO_OFFSET_RESET_CONFIG ;
2526import com .google .auto .value .AutoValue ;
2627import edu .umd .cs .findbugs .annotations .SuppressFBWarnings ;
2728import io .confluent .kafka .serializers .KafkaAvroDeserializer ;
29+ import io .opentelemetry .api .trace .Span ;
30+ import io .opentelemetry .api .trace .Tracer ;
31+ import io .opentelemetry .api .trace .propagation .W3CTraceContextPropagator ;
32+ import io .opentelemetry .context .Context ;
33+ import io .opentelemetry .context .Scope ;
34+ import io .opentelemetry .context .propagation .TextMapGetter ;
35+ import io .opentelemetry .context .propagation .TextMapSetter ;
2836import java .io .InputStream ;
2937import java .io .OutputStream ;
3038import java .lang .reflect .Method ;
4048import java .util .Set ;
4149import java .util .regex .Pattern ;
4250import java .util .stream .Collectors ;
51+ import java .util .stream .StreamSupport ;
4352import org .apache .beam .sdk .annotations .Internal ;
4453import org .apache .beam .sdk .coders .AtomicCoder ;
4554import org .apache .beam .sdk .coders .ByteArrayCoder ;
6170import org .apache .beam .sdk .options .Default ;
6271import org .apache .beam .sdk .options .ExperimentalOptions ;
6372import org .apache .beam .sdk .options .PipelineOptions ;
73+ import org .apache .beam .sdk .options .SdkHarnessOptions ;
6474import org .apache .beam .sdk .options .StreamingOptions ;
6575import org .apache .beam .sdk .options .ValueProvider ;
6676import org .apache .beam .sdk .runners .AppliedPTransform ;
125135import org .apache .kafka .common .TopicPartition ;
126136import org .apache .kafka .common .config .SaslConfigs ;
127137import org .apache .kafka .common .header .Header ;
138+ import org .apache .kafka .common .header .Headers ;
128139import org .apache .kafka .common .header .internals .RecordHeader ;
129140import org .apache .kafka .common .serialization .ByteArrayDeserializer ;
130141import org .apache .kafka .common .serialization .Deserializer ;
@@ -614,6 +625,7 @@ public static <K, V> Read<K, V> read() {
614625 .setTimestampPolicyFactory (TimestampPolicyFactory .withProcessingTime ())
615626 .setConsumerPollingTimeout (2L )
616627 .setRedistributed (false )
628+ .setEnableOpenTelemetryTracing (false )
617629 .setAllowDuplicates (false )
618630 .setRedistributeNumKeys (0 )
619631 .build ();
@@ -742,6 +754,9 @@ public abstract static class Read<K, V>
742754 @ Pure
743755 public abstract @ Nullable Duration getWatchTopicPartitionDuration ();
744756
757+ @ Pure
758+ public abstract boolean isEnableOpenTelemetryTracing ();
759+
745760 @ Pure
746761 public abstract TimestampPolicyFactory <K , V > getTimestampPolicyFactory ();
747762
@@ -832,6 +847,8 @@ Builder<K, V> setCheckStopReadingFn(
832847 return setCheckStopReadingFn (CheckStopReadingFnWrapper .of (checkStopReadingFn ));
833848 }
834849
850+ abstract Builder <K , V > setEnableOpenTelemetryTracing (boolean enableOpenTelemetryTracing );
851+
835852 abstract Builder <K , V > setConsumerPollingTimeout (long consumerPollingTimeout );
836853
837854 abstract Builder <K , V > setLogTopicVerification (@ Nullable Boolean logTopicVerification );
@@ -865,6 +882,7 @@ static <K, V> void setupExternalBuilder(
865882
866883 // Set required defaults
867884 builder .setTopicPartitions (Collections .emptyList ());
885+ builder .setEnableOpenTelemetryTracing (false );
868886 builder .setConsumerFactoryFn (KafkaIOUtils .KAFKA_CONSUMER_FACTORY_FN );
869887 if (config .maxReadTime != null ) {
870888 builder .setMaxReadTime (Duration .standardSeconds (config .maxReadTime ));
@@ -1302,6 +1320,10 @@ public Read<K, V> withValueDeserializer(DeserializerProvider<V> deserializerProv
13021320 return toBuilder ().setValueDeserializerProvider (deserializerProvider ).build ();
13031321 }
13041322
1323+ public Read <K , V > withEnableOpenTelemetryTracing () {
1324+ return toBuilder ().setEnableOpenTelemetryTracing (true ).build ();
1325+ }
1326+
13051327 public Read <K , V > withValueDeserializerProviderAndCoder (
13061328 DeserializerProvider <V > deserializerProvider , Coder <V > valueCoder ) {
13071329 return toBuilder ()
@@ -1920,6 +1942,14 @@ public PCollection<KafkaRecord<K, V>> expand(PBegin input) {
19201942 .withMaxNumRecords (kafkaRead .getMaxNumRecords ());
19211943 }
19221944 PCollection <KafkaRecord <K , V >> output = input .getPipeline ().apply (transform );
1945+
1946+ if (kafkaRead .isEnableOpenTelemetryTracing ()) {
1947+ output =
1948+ output .apply (
1949+ "Extract OpenTelemetry context from Header" ,
1950+ ParDo .of (new OpenTelemetryHeaderConsumer <>()));
1951+ }
1952+
19231953 if (kafkaRead .getOffsetDeduplication () != null && kafkaRead .getOffsetDeduplication ()) {
19241954 output =
19251955 output .apply (
@@ -2041,9 +2071,15 @@ public PCollection<KafkaRecord<K, V>> expand(PBegin input) {
20412071 .apply (ParDo .of (new GenerateKafkaSourceDescriptor (kafkaRead )));
20422072 }
20432073 }
2074+ PCollection <KafkaRecord <K , V >> pcol =
2075+ output .apply (readTransform ).setCoder (KafkaRecordCoder .of (keyCoder , valueCoder ));
2076+ if (kafkaRead .isEnableOpenTelemetryTracing ()) {
2077+ pcol =
2078+ pcol .apply (
2079+ "Extract OpenTelemetry context from Header" ,
2080+ ParDo .of (new OpenTelemetryHeaderConsumer <>()));
2081+ }
20442082 if (kafkaRead .isRedistributed ()) {
2045- PCollection <KafkaRecord <K , V >> pcol =
2046- output .apply (readTransform ).setCoder (KafkaRecordCoder .of (keyCoder , valueCoder ));
20472083 if (kafkaRead .getRedistributeNumKeys () == 0 ) {
20482084 return pcol .apply (
20492085 "Insert Redistribute" ,
@@ -2057,7 +2093,7 @@ public PCollection<KafkaRecord<K, V>> expand(PBegin input) {
20572093 .withNumBuckets ((int ) kafkaRead .getRedistributeNumKeys ()));
20582094 }
20592095 }
2060- return output . apply ( readTransform ). setCoder ( KafkaRecordCoder . of ( keyCoder , valueCoder )) ;
2096+ return pcol ;
20612097 }
20622098 }
20632099
@@ -2218,6 +2254,101 @@ public void populateDisplayData(DisplayData.Builder builder) {
22182254 }
22192255 }
22202256
2257+ static class OpenTelemetryHeaderConsumer <K , V >
2258+ extends DoFn <KafkaRecord <K , V >, KafkaRecord <K , V >> {
2259+ @ Nullable Tracer tracer = null ;
2260+
2261+ @ Setup
2262+ public void setup (PipelineOptions options ) {
2263+ // inject tracer via options
2264+ io .opentelemetry .api .OpenTelemetry openTelemetry =
2265+ options .as (SdkHarnessOptions .class ).getOpenTelemetry ();
2266+ if (openTelemetry != null ) {
2267+ tracer = openTelemetry .getTracer ("KafkaIO" );
2268+ }
2269+ }
2270+
2271+ Context extractSpanContext (KafkaRecord <K , V > message ) {
2272+ TextMapGetter <KafkaRecord <K , V >> extractMessageAttributes =
2273+ new TextMapGetter <KafkaRecord <K , V >>() {
2274+
2275+ @ Override
2276+ public @ Nullable String get (@ Nullable KafkaRecord <K , V > carrier , String key ) {
2277+ if (carrier == null ) {
2278+ return null ;
2279+ }
2280+ Headers headers = carrier .getHeaders ();
2281+ if (headers == null ) {
2282+ return null ;
2283+ }
2284+ Header header = headers .lastHeader (key );
2285+ if (header == null ) {
2286+ return null ;
2287+ }
2288+ return new String (header .value (), UTF_8 );
2289+ }
2290+
2291+ @ Override
2292+ public Iterable <String > keys (@ Nullable KafkaRecord <K , V > carrier ) {
2293+ if (carrier == null || carrier .getHeaders () == null ) {
2294+ return ImmutableList .of ();
2295+ }
2296+ return StreamSupport .stream (carrier .getHeaders ().spliterator (), false )
2297+ .map (Header ::key )
2298+ .collect (Collectors .toList ());
2299+ }
2300+ };
2301+ return W3CTraceContextPropagator .getInstance ()
2302+ .extract (Context .current (), message , extractMessageAttributes );
2303+ }
2304+
2305+ @ ProcessElement
2306+ public void processElement (
2307+ @ Element KafkaRecord <K , V > element , OutputReceiver <KafkaRecord <K , V >> receiver ) {
2308+ Context context = extractSpanContext (element );
2309+ Span span =
2310+ Preconditions .checkArgumentNotNull (tracer )
2311+ .spanBuilder ("KafkaIO.Read" )
2312+ .setParent (context )
2313+ .startSpan ();
2314+ try (Scope ignored = span .makeCurrent ()) {
2315+ receiver .output (element );
2316+ } finally {
2317+ span .end ();
2318+ }
2319+ }
2320+ }
2321+
2322+ static class OpenTelemetryHeaderPropagator <K , V >
2323+ extends DoFn <ProducerRecord <K , V >, ProducerRecord <K , V >> {
2324+ ProducerRecord <K , V > injectTraceContext (ProducerRecord <K , V > message ) {
2325+ org .apache .kafka .common .header .internals .RecordHeaders headers =
2326+ new org .apache .kafka .common .header .internals .RecordHeaders (message .headers ());
2327+ TextMapSetter <org .apache .kafka .common .header .internals .RecordHeaders >
2328+ injectMessageAttributes =
2329+ (carrier , key , value ) -> {
2330+ if (carrier != null ) {
2331+ carrier .add (key , value .getBytes (UTF_8 ));
2332+ }
2333+ };
2334+ W3CTraceContextPropagator .getInstance ()
2335+ .inject (Context .current (), headers , injectMessageAttributes );
2336+ return new ProducerRecord <>(
2337+ message .topic (),
2338+ message .partition (),
2339+ message .timestamp (),
2340+ message .key (),
2341+ message .value (),
2342+ headers );
2343+ }
2344+
2345+ @ ProcessElement
2346+ public void processElement (
2347+ @ Element ProducerRecord <K , V > element , OutputReceiver <ProducerRecord <K , V >> receiver ) {
2348+ receiver .output (injectTraceContext (element ));
2349+ }
2350+ }
2351+
22212352 /**
22222353 * A {@link PTransform} to read from Kafka topics. Similar to {@link KafkaIO.Read}, but removes
22232354 * Kafka metatdata and returns a {@link PCollection} of {@link KV}. See {@link KafkaIO} for more
@@ -3162,6 +3293,8 @@ public abstract static class WriteRecords<K, V>
31623293 // we shouldn't have to duplicate the same API for similar transforms like {@link Write} and
31633294 // {@link WriteRecords}. See example at {@link PubsubIO.Write}.
31643295
3296+ public abstract boolean isEnableOpenTelemetryTracing ();
3297+
31653298 @ Pure
31663299 public abstract @ Nullable String getTopic ();
31673300
@@ -3212,6 +3345,8 @@ public abstract static class WriteRecords<K, V>
32123345 abstract static class Builder <K , V > {
32133346 abstract Builder <K , V > setTopic (String topic );
32143347
3348+ abstract Builder <K , V > setEnableOpenTelemetryTracing (boolean enableOpenTelemetryTracing );
3349+
32153350 abstract Builder <K , V > setProducerConfig (Map <String , Object > producerConfig );
32163351
32173352 abstract Builder <K , V > setProducerFactoryFn (
@@ -3277,6 +3412,10 @@ public WriteRecords<K, V> withValueSerializer(Class<? extends Serializer<V>> val
32773412 return toBuilder ().setValueSerializer (valueSerializer ).build ();
32783413 }
32793414
3415+ public WriteRecords <K , V > withEnableOpenTelemetryTracing () {
3416+ return toBuilder ().setEnableOpenTelemetryTracing (true ).build ();
3417+ }
3418+
32803419 /**
32813420 * Adds the given producer properties, overriding old values of properties with the same key.
32823421 *
@@ -3413,7 +3552,11 @@ public PDone expand(PCollection<ProducerRecord<K, V>> input) {
34133552
34143553 checkArgument (getKeySerializer () != null , "withKeySerializer() is required" );
34153554 checkArgument (getValueSerializer () != null , "withValueSerializer() is required" );
3416-
3555+ if (this .isEnableOpenTelemetryTracing ()) {
3556+ input =
3557+ input .apply (
3558+ "Propagate OpenTelemetry Tracing" , ParDo .of (new OpenTelemetryHeaderPropagator <>()));
3559+ }
34173560 if (isEOS ()) {
34183561 checkArgument (getTopic () != null , "withTopic() is required when isEOS() is true" );
34193562 checkArgument (
@@ -3653,6 +3796,10 @@ public Write<K, V> withInputTimestamp() {
36533796 return withWriteRecordsTransform (getWriteRecordsTransform ().withInputTimestamp ());
36543797 }
36553798
3799+ public Write <K , V > withEnableOpenTelemetryTracing () {
3800+ return withWriteRecordsTransform (getWriteRecordsTransform ().withEnableOpenTelemetryTracing ());
3801+ }
3802+
36563803 /**
36573804 * Wrapper method over {@link
36583805 * WriteRecords#withPublishTimestampFunction(KafkaPublishTimestampFunction)}, used to keep the
0 commit comments