77import threading
88
99from datetime import datetime , timedelta
10- from typing import Any , Optional , Union
10+ from typing import Any , Mapping , Optional , Union , cast
1111from warnings import deprecated
1212import azure .functions as func
1313from urllib .parse import urlparse , quote
2525 AzureFunctionsDefaultClientInterceptorImpl ,
2626)
2727from .internal .serialization import DEFAULT_FUNCTIONS_DATA_CONVERTER
28- from .http import HttpManagementPayload
28+ from .http . http_management_payload import HttpManagementPayload , replace_url_origin
2929from .internal .compat .durable_orchestration_status import DurableOrchestrationStatus
3030from .internal .compat .entity_state_response import EntityStateResponse
3131from .internal .compat .orchestration_runtime_status import OrchestrationRuntimeStatus , to_durabletask_statuses
3636_sync_client_cache_lock = threading .Lock ()
3737
3838
39+ def _first_forwarded_value (value : str ) -> str :
40+ return value .split ("," , 1 )[0 ].strip ().strip ('"' )
41+
42+
43+ def _get_request_origin (
44+ request : func .HttpRequest ,
45+ use_forwarded_host : bool ) -> str :
46+ request_url = urlparse (request .url )
47+ proto = request_url .scheme
48+ host = request_url .netloc
49+ if not use_forwarded_host :
50+ return f"{ proto } ://{ host } "
51+
52+ request_headers = cast (Mapping [str , str ], request .headers )
53+ headers = {
54+ name .lower (): value for name , value in request_headers .items ()
55+ }
56+
57+ forwarded = headers .get ("forwarded" )
58+ if forwarded :
59+ forwarded_values : dict [str , str ] = {}
60+ for pair in forwarded .split ("," , 1 )[0 ].split (";" ):
61+ name , separator , value = pair .partition ("=" )
62+ if separator :
63+ forwarded_values [name .strip ().lower ()] = value .strip ().strip ('"' )
64+
65+ forwarded_proto = forwarded_values .get ("proto" )
66+ if forwarded_proto :
67+ proto = forwarded_proto
68+ forwarded_host = forwarded_values .get ("host" )
69+ if forwarded_host :
70+ return f"{ proto } ://{ forwarded_host } "
71+
72+ forwarded_proto = headers .get ("x-forwarded-proto" )
73+ if forwarded_proto :
74+ first_proto = _first_forwarded_value (forwarded_proto )
75+ if first_proto :
76+ proto = first_proto
77+
78+ forwarded_host = headers .get ("x-forwarded-host" )
79+ if forwarded_host :
80+ first_host = _first_forwarded_value (forwarded_host )
81+ if first_host :
82+ host = first_host
83+
84+ return f"{ proto } ://{ host } "
85+
86+
87+ def _build_http_management_payload (
88+ instance_id : str ,
89+ management_urls : dict [str , str ],
90+ base_url : str ,
91+ http_base_url : str ,
92+ required_query_string_parameters : str ,
93+ use_forwarded_host : bool ,
94+ request : func .HttpRequest | None ) -> HttpManagementPayload :
95+ encoded_instance_id = quote (instance_id , safe = "" )
96+ configured_base_url = http_base_url or base_url
97+ request_origin : str | None = None
98+ if request is not None :
99+ request_origin = _get_request_origin (request , use_forwarded_host )
100+ if configured_base_url :
101+ management_base_url = replace_url_origin (
102+ configured_base_url .rstrip ("/" ), request_origin )
103+ else :
104+ management_base_url = (
105+ f"{ request_origin } /runtime/webhooks/durabletask" )
106+ else :
107+ management_base_url = configured_base_url .rstrip ("/" )
108+ instance_status_url = (
109+ f"{ management_base_url } /instances/{ encoded_instance_id } " )
110+
111+ return HttpManagementPayload (
112+ instance_id ,
113+ instance_status_url ,
114+ required_query_string_parameters ,
115+ management_urls = management_urls ,
116+ request_origin = request_origin )
117+
118+
39119# Client class used for Durable Functions
40120class DurableFunctionsClient (AsyncTaskHubGrpcClient ):
41121 """A gRPC client passed to Durable Functions durable client bindings.
@@ -52,6 +132,7 @@ class DurableFunctionsClient(AsyncTaskHubGrpcClient):
52132 requiredQueryStringParameters : str
53133 rpcBaseUrl : str
54134 httpBaseUrl : str
135+ useForwardedHost : bool
55136 maxGrpcMessageSizeInBytes : int
56137 # The host sends this as a .NET TimeSpan string; it is currently stored
57138 # as-received (see _parse_client_configuration) and is unused, so the raw
@@ -161,6 +242,7 @@ def _parse_client_configuration(self, client_as_string: str) -> None:
161242 self .requiredQueryStringParameters = client .get ("requiredQueryStringParameters" ) or ""
162243 self .rpcBaseUrl = client .get ("rpcBaseUrl" ) or ""
163244 self .httpBaseUrl = client .get ("httpBaseUrl" ) or ""
245+ self .useForwardedHost = client .get ("useForwardedHost" ) or False
164246 self .maxGrpcMessageSizeInBytes = client .get ("maxGrpcMessageSizeInBytes" ) or 0
165247 # TODO: convert the string value back to timedelta - annoying regex?
166248 self .grpcHttpClientTimeout = client .get ("grpcHttpClientTimeout" ) or timedelta (seconds = 30 )
@@ -215,21 +297,14 @@ def create_http_management_payload(
215297 return self ._get_client_response_links (resolved_request , instance_id )
216298
217299 def _get_client_response_links (self , request : func .HttpRequest | None , instance_id : str ) -> HttpManagementPayload :
218- instance_status_url = self ._get_instance_status_url (request , instance_id )
219- return HttpManagementPayload (instance_id , instance_status_url , self .requiredQueryStringParameters )
220-
221- def _get_instance_status_url (self , request : func .HttpRequest | None , instance_id : str ) -> str :
222- encoded_instance_id = quote (instance_id )
223- if request is not None :
224- request_url = urlparse (request .url )
225- location_url = f"{ request_url .scheme } ://{ request_url .netloc } "
226- location_url = location_url + "/runtime/webhooks/durabletask/instances/" + encoded_instance_id
227- else :
228- # No request available (v1-style call): fall back to the base URL
229- # supplied in the client binding configuration.
230- base_url = self .baseUrl .rstrip ("/" ) if self .baseUrl else ""
231- location_url = base_url + "/instances/" + encoded_instance_id
232- return location_url
300+ return _build_http_management_payload (
301+ instance_id ,
302+ self .managementUrls ,
303+ self .baseUrl ,
304+ self .httpBaseUrl ,
305+ self .requiredQueryStringParameters ,
306+ self .useForwardedHost ,
307+ request )
233308
234309 # ------------------------------------------------------------------
235310 # Backwards-compatibility shims for the v1 azure-functions-durable
@@ -519,6 +594,7 @@ class SyncDurableFunctionsClient(TaskHubGrpcClient):
519594 requiredQueryStringParameters : str
520595 rpcBaseUrl : str
521596 httpBaseUrl : str
597+ useForwardedHost : bool
522598 maxGrpcMessageSizeInBytes : int
523599 grpcHttpClientTimeout : timedelta | str
524600
@@ -566,6 +642,7 @@ def _parse_client_configuration(self, client_as_string: str) -> None:
566642 "requiredQueryStringParameters" ) or ""
567643 self .rpcBaseUrl = client .get ("rpcBaseUrl" ) or ""
568644 self .httpBaseUrl = client .get ("httpBaseUrl" ) or ""
645+ self .useForwardedHost = client .get ("useForwardedHost" ) or False
569646 self .maxGrpcMessageSizeInBytes = client .get (
570647 "maxGrpcMessageSizeInBytes" ) or 0
571648 self .grpcHttpClientTimeout = client .get (
@@ -598,21 +675,14 @@ def create_http_management_payload(
598675 def _get_client_response_links (
599676 self , request : func .HttpRequest | None ,
600677 instance_id : str ) -> HttpManagementPayload :
601- return HttpManagementPayload (
678+ return _build_http_management_payload (
602679 instance_id ,
603- self ._get_instance_status_url (request , instance_id ),
604- self .requiredQueryStringParameters )
605-
606- def _get_instance_status_url (
607- self , request : func .HttpRequest | None , instance_id : str ) -> str :
608- encoded_instance_id = quote (instance_id )
609- if request is not None :
610- request_url = urlparse (request .url )
611- return (
612- f"{ request_url .scheme } ://{ request_url .netloc } "
613- f"/runtime/webhooks/durabletask/instances/{ encoded_instance_id } " )
614- base_url = self .baseUrl .rstrip ("/" ) if self .baseUrl else ""
615- return f"{ base_url } /instances/{ encoded_instance_id } "
680+ self .managementUrls ,
681+ self .baseUrl ,
682+ self .httpBaseUrl ,
683+ self .requiredQueryStringParameters ,
684+ self .useForwardedHost ,
685+ request )
616686
617687
618688def _close_cached_sync_clients () -> None :
0 commit comments