This is an automated email from the ASF dual-hosted git repository.

jrmccluskey pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/beam.git


The following commit(s) were added to refs/heads/master by this push:
     new dc508f69c1e [Python SDK] Add support for named data streams to provide 
bundle isolation (#40061)
dc508f69c1e is described below

commit dc508f69c1e41a7cf4ec0e5738be8fe13a9f2dac
Author: Sam Whittle <[email protected]>
AuthorDate: Tue Sep 15 16:27:53 2026 +0200

    [Python SDK] Add support for named data streams to provide bundle isolation 
(#40061)
    
    * Initial support for named data streams for python sdk.
    
    * fix format and lint::
    
    * fix format and lint
---
 .../apache_beam/runners/worker/bundle_processor.py | 72 ++++++++++++++----
 .../runners/worker/bundle_processor_test.py        | 74 +++++++++++++++++++
 .../apache_beam/runners/worker/data_plane.py       | 85 ++++++++++++++--------
 .../apache_beam/runners/worker/data_plane_test.py  | 42 +++++++++--
 .../apache_beam/runners/worker/sdk_worker.py       |  3 +-
 .../apache_beam/runners/worker/sdk_worker_test.py  | 20 +++++
 .../runners/worker/worker_id_interceptor.py        | 39 +++++++++-
 .../runners/worker/worker_id_interceptor_test.py   | 48 +++++++++++-
 sdks/python/apache_beam/transforms/environments.py |  1 +
 .../apache_beam/transforms/environments_test.py    |  2 +
 10 files changed, 331 insertions(+), 55 deletions(-)

diff --git a/sdks/python/apache_beam/runners/worker/bundle_processor.py 
b/sdks/python/apache_beam/runners/worker/bundle_processor.py
index 851efc81221..e62e522c00c 100644
--- a/sdks/python/apache_beam/runners/worker/bundle_processor.py
+++ b/sdks/python/apache_beam/runners/worker/bundle_processor.py
@@ -26,6 +26,7 @@ import bisect
 import collections
 import concurrent.futures
 import copy
+import functools
 import heapq
 import itertools
 import json
@@ -136,22 +137,48 @@ class RunnerIOOperation(operations.Operation):
       state_sampler: statesampler.StateSampler,
       windowed_coder: coders.Coder,
       transform_id: str,
-      data_channel: data_plane.DataChannel) -> None:
+      data_channel_factory: Callable[[Optional[str]], data_plane.DataChannel]
+  ) -> None:
     super().__init__(name_context, None, counter_factory, state_sampler)
     self.windowed_coder = windowed_coder
     self.windowed_coder_impl = windowed_coder.get_impl()
     # transform_id represents the consumer for the bytes in the data plane for 
a
     # DataInputOperation or a producer of these bytes for a 
DataOutputOperation.
     self.transform_id = transform_id
-    self.data_channel = data_channel
+    self.data_channel_factory = data_channel_factory
     for _, consumer_ops in consumers.items():
       for consumer in consumer_ops:
         self.add_receiver(consumer, 0)
 
+  def get_data_channel(
+      self, data_stream_id: Optional[str] = None) -> data_plane.DataChannel:
+    return self.data_channel_factory(data_stream_id)
+
 
 class DataOutputOperation(RunnerIOOperation):
   """A sink-like operation that gathers outputs to be sent back to the runner.
   """
+  def __init__(
+      self,
+      operation_name: common.NameContext,
+      step_name: Any,
+      consumers: Mapping[Any, list[operations.Operation]],
+      counter_factory: counters.CounterFactory,
+      state_sampler: statesampler.StateSampler,
+      windowed_coder: coders.Coder,
+      transform_id: str,
+      data_channel_factory: Callable[[Optional[str]], data_plane.DataChannel]
+  ) -> None:
+    super().__init__(
+        operation_name,
+        step_name,
+        consumers,
+        counter_factory,
+        state_sampler,
+        windowed_coder,
+        transform_id=transform_id,
+        data_channel_factory=data_channel_factory)
+
   def set_output_stream(
       self, output_stream: data_plane.ClosableOutputStream) -> None:
     self.output_stream = output_stream
@@ -171,13 +198,14 @@ class DataInputOperation(RunnerIOOperation):
   def __init__(
       self,
       operation_name: common.NameContext,
-      step_name,
+      step_name: Any,
       consumers: Mapping[Any, list[operations.Operation]],
       counter_factory: counters.CounterFactory,
       state_sampler: statesampler.StateSampler,
       windowed_coder: coders.Coder,
-      transform_id,
-      data_channel: data_plane.GrpcClientDataChannel) -> None:
+      transform_id: str,
+      data_channel_factory: Callable[[Optional[str]], data_plane.DataChannel]
+  ) -> None:
     super().__init__(
         operation_name,
         step_name,
@@ -186,7 +214,7 @@ class DataInputOperation(RunnerIOOperation):
         state_sampler,
         windowed_coder,
         transform_id=transform_id,
-        data_channel=data_channel)
+        data_channel_factory=data_channel_factory)
 
     self.consumer = next(iter(consumers.values()))
     self.splitting_lock = threading.Lock()
@@ -1238,7 +1266,9 @@ class BundleProcessor(object):
       op.reset()
 
   def process_bundle(
-      self, instruction_id: str
+      self,
+      instruction_id: str,
+      data_stream_id: Optional[str] = None
   ) -> tuple[list[beam_fn_api_pb2.DelayedBundleApplication], bool]:
 
     expected_input_ops: list[DataInputOperation] = []
@@ -1247,8 +1277,9 @@ class BundleProcessor(object):
       if isinstance(op, DataOutputOperation):
         # TODO(robertwb): Is there a better way to pass the instruction id to
         # the operation?
+        data_channel = op.get_data_channel(data_stream_id)
         op.set_output_stream(
-            op.data_channel.output_stream(instruction_id, op.transform_id))
+            data_channel.output_stream(instruction_id, op.transform_id))
       elif isinstance(op, DataInputOperation):
         # We must wait until we receive "end of stream" for each of these ops.
         expected_input_ops.append(op)
@@ -1274,18 +1305,27 @@ class BundleProcessor(object):
       # Add expected data inputs for each data channel.
       input_op_by_transform_id = {}
       for input_op in expected_input_ops:
-        data_channels[input_op.data_channel].append(input_op.transform_id)
+        data_channel = input_op.get_data_channel(data_stream_id)
+        data_channels[data_channel].append(input_op.transform_id)
         input_op_by_transform_id[input_op.transform_id] = input_op
 
       # Update timer_data channel with expected timer inputs.
-      if self.timer_data_channel:
-        data_channels[self.timer_data_channel].extend(
-            list(self.timers_info.keys()))
+      timer_data_channel = None
+      if self.process_bundle_descriptor.timer_api_service_descriptor.url:
+        timer_data_channel = (
+            self.data_channel_factory.create_data_channel_from_url(
+                
self.process_bundle_descriptor.timer_api_service_descriptor.url,
+                data_stream_id=data_stream_id))
+      elif self.timer_data_channel:
+        timer_data_channel = self.timer_data_channel
+
+      if timer_data_channel:
+        data_channels[timer_data_channel].extend(list(self.timers_info.keys()))
 
         # Set up timer output stream for DoOperation.
         for ((transform_id, timer_family_id),
              timer_info) in self.timers_info.items():
-          output_stream = self.timer_data_channel.output_timer_stream(
+          output_stream = timer_data_channel.output_timer_stream(
               instruction_id, transform_id, timer_family_id)
           timer_info.output_stream = output_stream
           self.ops[transform_id].add_timer_info(timer_family_id, timer_info)
@@ -1632,7 +1672,8 @@ def create_source_runner(
       factory.state_sampler,
       output_coder,
       transform_id=transform_id,
-      data_channel=factory.data_channel_factory.create_data_channel(grpc_port))
+      data_channel_factory=functools.partial(
+          factory.data_channel_factory.create_data_channel, grpc_port))
 
 
 @BeamTransformFactory.register_urn(
@@ -1652,7 +1693,8 @@ def create_sink_runner(
       factory.state_sampler,
       output_coder,
       transform_id=transform_id,
-      data_channel=factory.data_channel_factory.create_data_channel(grpc_port))
+      data_channel_factory=functools.partial(
+          factory.data_channel_factory.create_data_channel, grpc_port))
 
 
 @BeamTransformFactory.register_urn(OLD_DATAFLOW_RUNNER_HARNESS_READ_URN, None)
diff --git a/sdks/python/apache_beam/runners/worker/bundle_processor_test.py 
b/sdks/python/apache_beam/runners/worker/bundle_processor_test.py
index 0eb4dd9485f..ad2089c9938 100644
--- a/sdks/python/apache_beam/runners/worker/bundle_processor_test.py
+++ b/sdks/python/apache_beam/runners/worker/bundle_processor_test.py
@@ -20,6 +20,7 @@
 
 import random
 import unittest
+from unittest import mock
 
 import apache_beam as beam
 from apache_beam.coders import StrUtf8Coder
@@ -736,5 +737,78 @@ class OrderedListStateTest(unittest.TestCase):
     self.assertEqual([A1, A2, A7, B7, A8], list(self.state.read()))
 
 
+class NamedDataStreamsTest(unittest.TestCase):
+  def test_named_data_streams_routing(self):
+    descriptor = beam_fn_api_pb2.ProcessBundleDescriptor(id='descriptor_id')
+
+    # Coders
+    CODER_ID = 'coder'
+    descriptor.coders[
+        CODER_ID].spec.urn = common_urns.StandardCoders.Enum.BYTES.urn
+
+    # PCollections
+    PCOLLECTION_IN = 'pcoll_in'
+    descriptor.pcollections[PCOLLECTION_IN].unique_name = PCOLLECTION_IN
+    descriptor.pcollections[PCOLLECTION_IN].coder_id = CODER_ID
+
+    PCOLLECTION_OUT = 'pcoll_out'
+    descriptor.pcollections[PCOLLECTION_OUT].unique_name = PCOLLECTION_OUT
+    descriptor.pcollections[PCOLLECTION_OUT].coder_id = CODER_ID
+
+    # Source transform
+    SOURCE_ID = 'source'
+    source_transform = descriptor.transforms[SOURCE_ID]
+    source_transform.spec.urn = bundle_processor.DATA_INPUT_URN
+    source_port = beam_fn_api_pb2.RemoteGrpcPort(coder_id=CODER_ID)
+    source_port.api_service_descriptor.url = 'localhost:123'
+    source_transform.spec.payload = source_port.SerializeToString()
+    source_transform.outputs['None'] = PCOLLECTION_IN
+
+    # Sink transform
+    SINK_ID = 'sink'
+    sink_transform = descriptor.transforms[SINK_ID]
+    sink_transform.spec.urn = bundle_processor.DATA_OUTPUT_URN
+    sink_port = beam_fn_api_pb2.RemoteGrpcPort(coder_id=CODER_ID)
+    sink_port.api_service_descriptor.url = 'localhost:123'
+    sink_transform.spec.payload = sink_port.SerializeToString()
+    sink_transform.inputs['None'] = PCOLLECTION_IN
+    sink_transform.outputs['None'] = PCOLLECTION_OUT
+
+    data_channel_factory = mock.MagicMock()
+    mock_channel_default = mock.MagicMock()
+    mock_channel_named = mock.MagicMock()
+
+    def get_channel(port, data_stream_id):
+      if data_stream_id == 'named_stream':
+        return mock_channel_named
+      return mock_channel_default
+
+    data_channel_factory.create_data_channel.side_effect = get_channel
+
+    mock_channel_default.input_elements.return_value = []
+    mock_channel_named.input_elements.return_value = []
+
+    processor = BundleProcessor(
+        frozenset(), descriptor, None, data_channel_factory)
+
+    # Process on default stream
+    processor.process_bundle('inst_1')
+    data_channel_factory.create_data_channel.assert_any_call(source_port, None)
+    data_channel_factory.create_data_channel.assert_any_call(sink_port, None)
+    mock_channel_default.output_stream.assert_called_once_with(
+        'inst_1', SINK_ID)
+
+    processor.reset()
+
+    # Process on named stream
+    processor.process_bundle('inst_2', data_stream_id='named_stream')
+
+    data_channel_factory.create_data_channel.assert_any_call(
+        source_port, 'named_stream')
+    data_channel_factory.create_data_channel.assert_any_call(
+        sink_port, 'named_stream')
+    mock_channel_named.output_stream.assert_called_once_with('inst_2', SINK_ID)
+
+
 if __name__ == '__main__':
   unittest.main()
diff --git a/sdks/python/apache_beam/runners/worker/data_plane.py 
b/sdks/python/apache_beam/runners/worker/data_plane.py
index cfefa37d76b..2a626b163a6 100644
--- a/sdks/python/apache_beam/runners/worker/data_plane.py
+++ b/sdks/python/apache_beam/runners/worker/data_plane.py
@@ -48,6 +48,7 @@ from apache_beam.coders import coder_impl
 from apache_beam.portability.api import beam_fn_api_pb2
 from apache_beam.portability.api import beam_fn_api_pb2_grpc
 from apache_beam.runners.worker.channel_factory import GRPCChannelFactory
+from apache_beam.runners.worker.worker_id_interceptor import 
DataStreamIdInterceptor
 from apache_beam.runners.worker.worker_id_interceptor import 
WorkerIdInterceptor
 from apache_beam.utils.byte_limited_queue import ByteLimitedQueue
 
@@ -804,12 +805,14 @@ class 
BeamFnDataServicer(beam_fn_api_pb2_grpc.BeamFnDataServicer):
     self._lock = threading.Lock()
     self._connections_by_worker_id = collections.defaultdict(
         lambda: _GrpcDataChannel(data_buffer_time_limit_ms)
-    )  # type: DefaultDict[str, _GrpcDataChannel]
+    )  # type: DefaultDict[Tuple[str, str], _GrpcDataChannel]
 
-  def get_conn_by_worker_id(self, worker_id):
-    # type: (str) -> _GrpcDataChannel
+  def get_conn_by_worker_id(
+      self,
+      worker_id: str,
+      data_stream_id: Optional[str] = None) -> _GrpcDataChannel:
     with self._lock:
-      return self._connections_by_worker_id[worker_id]
+      return self._connections_by_worker_id[(worker_id, data_stream_id or '')]
 
   def Data(
       self,
@@ -817,8 +820,10 @@ class 
BeamFnDataServicer(beam_fn_api_pb2_grpc.BeamFnDataServicer):
       context  # type: Any
   ):
     # type: (...) -> Iterator[beam_fn_api_pb2.Elements]
-    worker_id = dict(context.invocation_metadata())['worker_id']
-    data_conn = self.get_conn_by_worker_id(worker_id)
+    metadata = dict(context.invocation_metadata())
+    worker_id = metadata['worker_id']
+    data_stream_id = metadata.get('data_stream_id', '')
+    data_conn = self.get_conn_by_worker_id(worker_id, data_stream_id)
     data_conn.set_inputs(elements_iterator)
     for elements in data_conn._write_outputs():
       yield elements
@@ -827,16 +832,18 @@ class 
BeamFnDataServicer(beam_fn_api_pb2_grpc.BeamFnDataServicer):
 class DataChannelFactory(metaclass=abc.ABCMeta):
   """An abstract factory for creating ``DataChannel``."""
   @abc.abstractmethod
-  def create_data_channel(self, remote_grpc_port):
-    # type: (beam_fn_api_pb2.RemoteGrpcPort) -> GrpcClientDataChannel
-
+  def create_data_channel(
+      self,
+      remote_grpc_port: beam_fn_api_pb2.RemoteGrpcPort,
+      data_stream_id: Optional[str] = None) -> DataChannel:
     """Returns a ``DataChannel`` from the given RemoteGrpcPort."""
     raise NotImplementedError(type(self))
 
   @abc.abstractmethod
-  def create_data_channel_from_url(self, url):
-    # type: (str) -> Optional[GrpcClientDataChannel]
-
+  def create_data_channel_from_url(
+      self,
+      url: str,
+      data_stream_id: Optional[str] = None) -> Optional[DataChannel]:
     """Returns a ``DataChannel`` from the given url."""
     raise NotImplementedError(type(self))
 
@@ -857,7 +864,7 @@ class DataChannelFactory(metaclass=abc.ABCMeta):
 class GrpcClientDataChannelFactory(DataChannelFactory):
   """A factory for ``GrpcClientDataChannel``.
 
-  Caches the created channels by ``data descriptor url``.
+  Caches the created channels by ``(data descriptor url, data_stream_id)``.
   """
   def __init__(
       self,
@@ -866,7 +873,8 @@ class GrpcClientDataChannelFactory(DataChannelFactory):
       data_buffer_time_limit_ms=0  # type: int
   ):
     # type: (...) -> None
-    self._data_channel_cache = {}  # type: Dict[str, GrpcClientDataChannel]
+    self._data_channel_cache = {
+    }  # type: Dict[Tuple[str, str], GrpcClientDataChannel]
     self._lock = threading.Lock()
     self._credentials = None
     self._worker_id = worker_id
@@ -875,14 +883,21 @@ class GrpcClientDataChannelFactory(DataChannelFactory):
       _LOGGER.info('Using secure channel creds.')
       self._credentials = credentials
 
-  def create_data_channel_from_url(self, url):
-    # type: (str) -> Optional[GrpcClientDataChannel]
+  def create_data_channel_from_url(
+      self,
+      url: str,
+      data_stream_id: Optional[str] = None) -> Optional[GrpcClientDataChannel]:
     if not url:
       return None
-    if url not in self._data_channel_cache:
+    data_stream_id = data_stream_id or ''
+    cache_key = (url, data_stream_id)
+    if cache_key not in self._data_channel_cache:
       with self._lock:
-        if url not in self._data_channel_cache:
-          _LOGGER.info('Creating client data channel for %s', url)
+        if cache_key not in self._data_channel_cache:
+          _LOGGER.info(
+              'Creating client data channel for %s (data_stream_id: %s)',
+              url,
+              data_stream_id)
           # Options to have no limits (-1) on the size of the messages
           # received or sent over the data plane. The actual buffer size
           # is controlled in a layer above.
@@ -897,22 +912,26 @@ class GrpcClientDataChannelFactory(DataChannelFactory):
             grpc_channel = GRPCChannelFactory.secure_channel(
                 url, self._credentials, options=channel_options)
           _LOGGER.info('Data channel established.')
-          # Add workerId to the grpc channel
-          grpc_channel = grpc.intercept_channel(
-              grpc_channel, WorkerIdInterceptor(self._worker_id))
-          self._data_channel_cache[url] = GrpcClientDataChannel(
+          # Add workerId and optional data_stream_id to the grpc channel
+          interceptors = [WorkerIdInterceptor(self._worker_id)]
+          if data_stream_id:
+            interceptors.append(DataStreamIdInterceptor(data_stream_id))
+          grpc_channel = grpc.intercept_channel(grpc_channel, *interceptors)
+          self._data_channel_cache[cache_key] = GrpcClientDataChannel(
               beam_fn_api_pb2_grpc.BeamFnDataStub(grpc_channel),
               self._data_buffer_time_limit_ms)
 
-    return self._data_channel_cache[url]
+    return self._data_channel_cache[cache_key]
 
-  def create_data_channel(self, remote_grpc_port):
-    # type: (beam_fn_api_pb2.RemoteGrpcPort) -> GrpcClientDataChannel
+  def create_data_channel(
+      self,
+      remote_grpc_port: beam_fn_api_pb2.RemoteGrpcPort,
+      data_stream_id: Optional[str] = None) -> GrpcClientDataChannel:
     url = remote_grpc_port.api_service_descriptor.url
     # TODO(https://github.com/apache/beam/issues/19737): this can return None
     #  if url is falsey, but this seems incorrect, as code that calls this
     #  method seems to always expect non-Optional values.
-    return self.create_data_channel_from_url(url)  # type: ignore[return-value]
+    return self.create_data_channel_from_url(url, 
data_stream_id=data_stream_id)  # type: ignore[return-value]
 
   def close(self):
     # type: () -> None
@@ -930,15 +949,17 @@ class GrpcClientDataChannelFactory(DataChannelFactory):
 class InMemoryDataChannelFactory(DataChannelFactory):
   """A singleton factory for ``InMemoryDataChannel``."""
   def __init__(self, in_memory_data_channel):
-    # type: (GrpcClientDataChannel) -> None
+    # type: (DataChannel) -> None
     self._in_memory_data_channel = in_memory_data_channel
 
-  def create_data_channel(self, unused_remote_grpc_port):
-    # type: (beam_fn_api_pb2.RemoteGrpcPort) -> GrpcClientDataChannel
+  def create_data_channel(
+      self,
+      unused_remote_grpc_port: beam_fn_api_pb2.RemoteGrpcPort,
+      data_stream_id: Optional[str] = None) -> DataChannel:
     return self._in_memory_data_channel
 
-  def create_data_channel_from_url(self, url):
-    # type: (Any) -> GrpcClientDataChannel
+  def create_data_channel_from_url(
+      self, url: Any, data_stream_id: Optional[str] = None) -> DataChannel:
     return self._in_memory_data_channel
 
   def close(self):
diff --git a/sdks/python/apache_beam/runners/worker/data_plane_test.py 
b/sdks/python/apache_beam/runners/worker/data_plane_test.py
index 5124bb69e6c..16648f23f5f 100644
--- a/sdks/python/apache_beam/runners/worker/data_plane_test.py
+++ b/sdks/python/apache_beam/runners/worker/data_plane_test.py
@@ -29,6 +29,7 @@ import grpc
 from apache_beam.portability.api import beam_fn_api_pb2
 from apache_beam.portability.api import beam_fn_api_pb2_grpc
 from apache_beam.runners.worker import data_plane
+from apache_beam.runners.worker.worker_id_interceptor import 
DataStreamIdInterceptor
 from apache_beam.runners.worker.worker_id_interceptor import 
WorkerIdInterceptor
 from apache_beam.utils import thread_pool_executor
 
@@ -40,7 +41,10 @@ class DataChannelTest(unittest.TestCase):
   def test_time_based_flush_grpc_data_channel(self):
     self._grpc_data_channel_test(True)
 
-  def _grpc_data_channel_test(self, time_based_flush=False):
+  def test_named_grpc_data_channel(self):
+    self._grpc_data_channel_test(data_stream_id='stream_1')
+
+  def _grpc_data_channel_test(self, time_based_flush=False, data_stream_id=''):
     if time_based_flush:
       data_servicer = data_plane.BeamFnDataServicer(
           data_buffer_time_limit_ms=100)
@@ -48,7 +52,7 @@ class DataChannelTest(unittest.TestCase):
       data_servicer = data_plane.BeamFnDataServicer()
     worker_id = 'worker_0'
     data_channel_service = \
-      data_servicer.get_conn_by_worker_id(worker_id)
+      data_servicer.get_conn_by_worker_id(worker_id, 
data_stream_id=data_stream_id)
 
     server = grpc.server(thread_pool_executor.shared_unbounded_instance())
     beam_fn_api_pb2_grpc.add_BeamFnDataServicer_to_server(data_servicer, 
server)
@@ -56,9 +60,10 @@ class DataChannelTest(unittest.TestCase):
     server.start()
 
     grpc_channel = grpc.insecure_channel('localhost:%s' % test_port)
-    # Add workerId to the grpc channel
-    grpc_channel = grpc.intercept_channel(
-        grpc_channel, WorkerIdInterceptor(worker_id))
+    interceptors = [WorkerIdInterceptor(worker_id)]
+    if data_stream_id:
+      interceptors.append(DataStreamIdInterceptor(data_stream_id))
+    grpc_channel = grpc.intercept_channel(grpc_channel, *interceptors)
     data_channel_stub = beam_fn_api_pb2_grpc.BeamFnDataStub(grpc_channel)
     if time_based_flush:
       data_channel_client = data_plane.GrpcClientDataChannel(
@@ -75,6 +80,33 @@ class DataChannelTest(unittest.TestCase):
       data_channel_client.wait()
       data_channel_service.wait()
 
+  def test_grpc_client_data_channel_factory_named_streams(self):
+    data_servicer = data_plane.BeamFnDataServicer()
+    server = grpc.server(thread_pool_executor.shared_unbounded_instance())
+    beam_fn_api_pb2_grpc.add_BeamFnDataServicer_to_server(data_servicer, 
server)
+    test_port = server.add_insecure_port('[::]:0')
+    server.start()
+
+    try:
+      factory = data_plane.GrpcClientDataChannelFactory(worker_id='worker_0')
+      url = 'localhost:%s' % test_port
+      ch_default = factory.create_data_channel_from_url(url)
+      ch_default2 = factory.create_data_channel_from_url(url, 
data_stream_id='')
+      ch_stream1 = factory.create_data_channel_from_url(
+          url, data_stream_id='stream_1')
+      ch_stream1_dup = factory.create_data_channel_from_url(
+          url, data_stream_id='stream_1')
+      ch_stream2 = factory.create_data_channel_from_url(
+          url, data_stream_id='stream_2')
+
+      self.assertIs(ch_default, ch_default2)
+      self.assertIs(ch_stream1, ch_stream1_dup)
+      self.assertIsNot(ch_default, ch_stream1)
+      self.assertIsNot(ch_stream1, ch_stream2)
+      factory.close()
+    finally:
+      server.stop(0)
+
   def test_in_memory_data_channel(self):
     channel = data_plane.InMemoryDataChannel()
     self._data_channel_test(channel, channel.inverse())
diff --git a/sdks/python/apache_beam/runners/worker/sdk_worker.py 
b/sdks/python/apache_beam/runners/worker/sdk_worker.py
index db25a40a405..84b4ae1bf23 100644
--- a/sdks/python/apache_beam/runners/worker/sdk_worker.py
+++ b/sdks/python/apache_beam/runners/worker/sdk_worker.py
@@ -706,7 +706,8 @@ class SdkWorker(object):
           instruction_id, request.cache_tokens):
         with self.maybe_profile(instruction_id):
           delayed_applications, requests_finalization = (
-              bundle_processor.process_bundle(instruction_id))
+              bundle_processor.process_bundle(
+                  instruction_id, getattr(request, 'data_stream_id', '')))
           monitoring_infos = bundle_processor.monitoring_infos()
           response = beam_fn_api_pb2.InstructionResponse(
               instruction_id=instruction_id,
diff --git a/sdks/python/apache_beam/runners/worker/sdk_worker_test.py 
b/sdks/python/apache_beam/runners/worker/sdk_worker_test.py
index bea313a4d2f..4b00051a337 100644
--- a/sdks/python/apache_beam/runners/worker/sdk_worker_test.py
+++ b/sdks/python/apache_beam/runners/worker/sdk_worker_test.py
@@ -412,6 +412,26 @@ class SdkWorkerTest(unittest.TestCase):
 
     self.assertIn(instruction_id, channel._cleaned_instruction_ids)
 
+  def test_process_bundle_passes_data_stream_id(self):
+    mock_bundle_processor = mock.MagicMock()
+    mock_bundle_processor.process_bundle.return_value = ([], False)
+    mock_bundle_processor.monitoring_infos.return_value = []
+    mock_bundle_processor.state_handler.process_instruction_id.return_value = 
contextlib.nullcontext(
+    )
+
+    bundle_processor_cache = mock.MagicMock()
+    bundle_processor_cache.get.return_value = mock_bundle_processor
+
+    worker = SdkWorker(bundle_processor_cache)
+    instruction_id = 'instruction_id'
+    request = beam_fn_api_pb2.ProcessBundleRequest(
+        process_bundle_descriptor_id='descriptor_id',
+        data_stream_id='stream_xyz')
+
+    worker.process_bundle(request, instruction_id)
+    mock_bundle_processor.process_bundle.assert_called_once_with(
+        instruction_id, 'stream_xyz')
+
 
 class CachingStateHandlerTest(unittest.TestCase):
   def test_caching(self):
diff --git a/sdks/python/apache_beam/runners/worker/worker_id_interceptor.py 
b/sdks/python/apache_beam/runners/worker/worker_id_interceptor.py
index 1db2b5f4a15..ca9476f0c2f 100644
--- a/sdks/python/apache_beam/runners/worker/worker_id_interceptor.py
+++ b/sdks/python/apache_beam/runners/worker/worker_id_interceptor.py
@@ -61,12 +61,49 @@ class WorkerIdInterceptor(grpc.UnaryUnaryClientInterceptor,
     metadata = []
     if client_call_details.metadata is not None:
       metadata = list(client_call_details.metadata)
-    if 'worker_id' in metadata:
+    if any(k == 'worker_id' for k, _ in metadata):
       raise RuntimeError('Header metadata already has a worker_id.')
     metadata.append(('worker_id', self._worker_id))
+
     new_client_details = _ClientCallDetails(
         client_call_details.method,
         client_call_details.timeout,
         metadata,
         client_call_details.credentials)
     return continuation(new_client_details, request)
+
+
+class DataStreamIdInterceptor(grpc.UnaryUnaryClientInterceptor,
+                              grpc.StreamStreamClientInterceptor):
+  """Client Interceptor to inject data_stream_id into metadata."""
+  def __init__(self, data_stream_id: Optional[str] = None) -> None:
+    self._data_stream_id = data_stream_id
+
+  def intercept_unary_unary(self, continuation, client_call_details, request):
+    return self._intercept(continuation, client_call_details, request)
+
+  def intercept_unary_stream(self, continuation, client_call_details, request):
+    return self._intercept(continuation, client_call_details, request)
+
+  def intercept_stream_unary(self, continuation, client_call_details, request):
+    return self._intercept(continuation, client_call_details, request)
+
+  def intercept_stream_stream(
+      self, continuation, client_call_details, request_iterator):
+    return self._intercept(continuation, client_call_details, request_iterator)
+
+  def _intercept(self, continuation, client_call_details, request):
+    if self._data_stream_id:
+      metadata = []
+      if client_call_details.metadata is not None:
+        metadata = list(client_call_details.metadata)
+      if any(k == 'data_stream_id' for k, _ in metadata):
+        raise RuntimeError('Header metadata already has a data_stream_id.')
+      metadata.append(('data_stream_id', self._data_stream_id))
+      new_client_details = _ClientCallDetails(
+          client_call_details.method,
+          client_call_details.timeout,
+          metadata,
+          client_call_details.credentials)
+      return continuation(new_client_details, request)
+    return continuation(client_call_details, request)
diff --git 
a/sdks/python/apache_beam/runners/worker/worker_id_interceptor_test.py 
b/sdks/python/apache_beam/runners/worker/worker_id_interceptor_test.py
index 0db9c1b4ddc..06546ab3ce2 100644
--- a/sdks/python/apache_beam/runners/worker/worker_id_interceptor_test.py
+++ b/sdks/python/apache_beam/runners/worker/worker_id_interceptor_test.py
@@ -24,6 +24,7 @@ import unittest
 
 import grpc
 
+from apache_beam.runners.worker.worker_id_interceptor import 
DataStreamIdInterceptor
 from apache_beam.runners.worker.worker_id_interceptor import 
WorkerIdInterceptor
 
 
@@ -62,7 +63,52 @@ class WorkerIdInterceptorTest(unittest.TestCase):
     with self.assertRaises(RuntimeError):
       WorkerIdInterceptor().intercept_stream_stream(
           continuation,
-          _ClientCallDetails(None, None, {'worker_id': '1'}, None), [])
+          _ClientCallDetails(None, None, [('worker_id', '1')], None), [])
+
+
+class DataStreamIdInterceptorTest(unittest.TestCase):
+  def test_data_stream_id_insertion(self):
+    data_stream_id_key = 'data_stream_id'
+    headers_holder = {}
+
+    def continuation(client_details, request_iterator):
+      headers_holder.update({
+          data_stream_id_key: dict(
+              client_details.metadata).get(data_stream_id_key)
+      })
+
+    DataStreamIdInterceptor('stream_123').intercept_stream_stream(
+        continuation, _ClientCallDetails(None, None, None, None), [])
+    self.assertEqual(
+        headers_holder[data_stream_id_key],
+        'stream_123',
+        'data_stream_id not set')
+
+  def test_no_data_stream_id_when_empty(self):
+    headers_holder = {}
+
+    def continuation(client_details, request_iterator):
+      headers_holder['metadata'] = client_details.metadata
+
+    DataStreamIdInterceptor('').intercept_stream_stream(
+        continuation, _ClientCallDetails(None, None, None, None), [])
+    self.assertIsNone(headers_holder['metadata'])
+
+  def test_failure_when_data_stream_id_exists(self):
+    data_stream_id_key = 'data_stream_id'
+    headers_holder = {}
+
+    def continuation(client_details, request_iterator):
+      headers_holder.update({
+          data_stream_id_key: dict(
+              client_details.metadata).get(data_stream_id_key)
+      })
+
+    with self.assertRaises(RuntimeError):
+      DataStreamIdInterceptor('stream_123').intercept_stream_stream(
+          continuation,
+          _ClientCallDetails(
+              None, None, [('data_stream_id', 'existing')], None), [])
 
 
 if __name__ == '__main__':
diff --git a/sdks/python/apache_beam/transforms/environments.py 
b/sdks/python/apache_beam/transforms/environments.py
index 6f17ca270eb..9d5e6a48cc0 100644
--- a/sdks/python/apache_beam/transforms/environments.py
+++ b/sdks/python/apache_beam/transforms/environments.py
@@ -917,6 +917,7 @@ def _python_sdk_capabilities_iter():
   yield common_urns.protocols.DATA_SAMPLING.urn
   yield common_urns.protocols.SDK_CONSUMING_RECEIVED_DATA.urn
   yield common_urns.protocols.ORDERED_LIST_STATE.urn
+  yield common_urns.protocols.NAMED_DATA_STREAMS.urn
 
 
 def python_sdk_dependencies(options, tmp_dir=None):
diff --git a/sdks/python/apache_beam/transforms/environments_test.py 
b/sdks/python/apache_beam/transforms/environments_test.py
index c32a85579fc..21ede9e0d07 100644
--- a/sdks/python/apache_beam/transforms/environments_test.py
+++ b/sdks/python/apache_beam/transforms/environments_test.py
@@ -72,6 +72,8 @@ class RunnerApiTest(unittest.TestCase):
         common_urns.sdf_components.TRUNCATE_SIZED_RESTRICTION.urn,
         sdk_capabilities)
     self.assertIn(common_urns.primitives.TO_STRING.urn, sdk_capabilities)
+    self.assertIn(
+        common_urns.protocols.NAMED_DATA_STREAMS.urn, sdk_capabilities)
 
   def test_default_capabilities(self):
     environment = DockerEnvironment.from_options(

Reply via email to