1212from dataclasses import dataclass
1313from typing import Any
1414
15- from google .protobuf .message import Message
16-
1715import temporalio .api .common .v1
1816import temporalio .common
1917import temporalio .converter
@@ -93,27 +91,31 @@ class _SystemNexusPayloadConverter(temporalio.converter.PayloadConverter):
9391 """Payload converter for system Nexus outer envelopes."""
9492
9593 _user_converters : _SystemNexusUserConverters
96- _outer_payload_converter : temporalio .converter .PayloadConverter
94+ _outer_payload_converter : _TemporalTransferTypePayloadConverter
95+ _headers : Mapping [str , temporalio .api .common .v1 .Payload ] | None
9796
9897 def __init__ (
9998 self ,
10099 user_payload_converter : temporalio .converter .PayloadConverter ,
101100 user_failure_converter : temporalio .converter .FailureConverter ,
101+ headers : Mapping [str , temporalio .api .common .v1 .Payload ] | None = None ,
102102 ) -> None :
103103 """Create a payload converter for system Nexus outer envelopes."""
104104 self ._user_converters = _SystemNexusUserConverters (
105105 user_payload_converter , user_failure_converter
106106 )
107- self ._outer_payload_converter = _TemporalTransferTypePayloadConverter .wrap (
107+
108+ self ._outer_payload_converter = _TemporalTransferTypePayloadConverter (
108109 _SystemNexusOuterPayloadConverter ()
109110 )
111+ self ._headers = headers
110112
111113 def to_payloads (
112114 self , values : Sequence [Any ]
113115 ) -> list [temporalio .api .common .v1 .Payload ]:
114116 """See base class."""
115117 with _user_converter_context (self ._user_converters ):
116- return self ._outer_payload_converter .to_payloads (values )
118+ return self ._outer_payload_converter .to_payloads (values , self . _headers )
117119
118120 def from_payloads (
119121 self ,
@@ -134,23 +136,13 @@ def is_system_endpoint(endpoint: str) -> bool:
134136 return endpoint == TEMPORAL_SYSTEM_ENDPOINT
135137
136138
137- def _apply_headers_to_request (
138- request : Message ,
139- headers : Mapping [str , temporalio .api .common .v1 .Payload ],
140- ) -> None :
141- """Apply headers to a system request when it supports Temporal headers."""
142- if not headers or "header" not in request .DESCRIPTOR .fields_by_name :
143- return
144- request_header = getattr (request , "header" )
145- for key , payload in headers .items ():
146- request_header .fields [key ].CopyFrom (payload )
147-
148-
149139def _is_system_payload (payload : temporalio .api .common .v1 .Payload ) -> bool :
150140 return (
151141 payload .metadata .get (_SYSTEM_PAYLOAD_METADATA_KEY )
152142 == _SYSTEM_PAYLOAD_METADATA_VALUE
153143 )
144+
145+
154146async def maybe_visit_payload (
155147 payload : temporalio .api .common .v1 .Payload ,
156148 visitor_functions : VisitorFunctions ,
@@ -175,9 +167,12 @@ async def maybe_visit_payload(
175167def _get_payload_converter ( # pyright: ignore[reportUnusedFunction]
176168 user_payload_converter : temporalio .converter .PayloadConverter ,
177169 user_failure_converter : temporalio .converter .FailureConverter ,
170+ headers : Mapping [str , temporalio .api .common .v1 .Payload ] | None = None ,
178171) -> temporalio .converter .PayloadConverter :
179172 """Return the fixed payload converter for system Nexus outer envelopes."""
180- return _SystemNexusPayloadConverter (user_payload_converter , user_failure_converter )
173+ return _SystemNexusPayloadConverter (
174+ user_payload_converter , user_failure_converter , headers
175+ )
181176
182177
183178def _get_serialization_context ( # pyright: ignore[reportUnusedFunction]
0 commit comments