Skip to content

Commit 6d929cd

Browse files
andystaplesCopilotberndverst
authored
Use canonical management URLs for HTTP payloads (#237)
* Fix canonical management payload URLs Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 3ec5c077-6e8b-4cb3-916a-12bd3cb13937 * Address management URL review feedback Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 3ec5c077-6e8b-4cb3-916a-12bd3cb13937 * Validate forwarded management URL origins Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: 3ec5c077-6e8b-4cb3-916a-12bd3cb13937 --------- Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Co-authored-by: Bernd Verst <github@bernd.dev> Copilot-Session: 3ec5c077-6e8b-4cb3-916a-12bd3cb13937
1 parent c595295 commit 6d929cd

4 files changed

Lines changed: 364 additions & 37 deletions

File tree

azure-functions-durable/CHANGELOG.md

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,9 @@ avoiding repeated allocation of unused worker resources.
3232

3333
FIXED
3434

35+
- HTTP management payloads now preserve the host-provided management URL
36+
templates and configured HTTP base paths, include `rewindPostUri`, encode
37+
instance IDs, and use forwarded request origins when enabled by the host.
3538
- Fixed deprecated v1 status-query methods omitting orchestration output,
3639
custom status, and failure details when input display was disabled.
3740
- Fixed asynchronous durable-client construction failing after an application

azure-functions-durable/azure/durable_functions/client.py

Lines changed: 101 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
import threading
88

99
from datetime import datetime, timedelta
10-
from typing import Any, Optional, Union
10+
from typing import Any, Mapping, Optional, Union, cast
1111
from warnings import deprecated
1212
import azure.functions as func
1313
from urllib.parse import urlparse, quote
@@ -25,7 +25,7 @@
2525
AzureFunctionsDefaultClientInterceptorImpl,
2626
)
2727
from .internal.serialization import DEFAULT_FUNCTIONS_DATA_CONVERTER
28-
from .http import HttpManagementPayload
28+
from .http.http_management_payload import HttpManagementPayload, replace_url_origin
2929
from .internal.compat.durable_orchestration_status import DurableOrchestrationStatus
3030
from .internal.compat.entity_state_response import EntityStateResponse
3131
from .internal.compat.orchestration_runtime_status import OrchestrationRuntimeStatus, to_durabletask_statuses
@@ -36,6 +36,86 @@
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
40120
class 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

618688
def _close_cached_sync_clients() -> None:

azure-functions-durable/azure/durable_functions/http/http_management_payload.py

Lines changed: 59 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,11 @@
22
# Licensed under the MIT License.
33

44
import json
5-
from typing import Any
5+
from typing import Any, Mapping
6+
from urllib.parse import quote, urlsplit, urlunsplit
7+
8+
9+
_INSTANCE_ID_PLACEHOLDER = "INSTANCEID"
610

711

812
class HttpManagementPayload(dict[str, str]):
@@ -18,24 +22,52 @@ class HttpManagementPayload(dict[str, str]):
1822
via ``json.dumps(payload)``.
1923
"""
2024

21-
def __init__(self, instance_id: str, instance_status_url: str, required_query_string_parameters: str):
25+
def __init__(
26+
self,
27+
instance_id: str,
28+
instance_status_url: str,
29+
required_query_string_parameters: str,
30+
*,
31+
management_urls: Mapping[str, str] | None = None,
32+
request_origin: str | None = None):
2233
"""Initializes the HttpManagementPayload with the necessary URLs.
2334
2435
Args:
2536
instance_id (str): The ID of the Durable Function instance.
2637
instance_status_url (str): The base URL for the instance status.
2738
required_query_string_parameters (str): The required URL parameters provided by the Durable extension.
39+
management_urls (Mapping[str, str] | None): Canonical URL templates
40+
provided by the Durable extension.
41+
request_origin (str | None): Externally visible request origin used
42+
to replace the templates' internal origin.
2843
"""
29-
super().__init__({
30-
'id': instance_id,
44+
fallback_urls = {
3145
'purgeHistoryDeleteUri': instance_status_url + "?" + required_query_string_parameters,
3246
'restartPostUri': instance_status_url + "/restart?" + required_query_string_parameters,
3347
'sendEventPostUri': instance_status_url + "/raiseEvent/{eventName}?" + required_query_string_parameters,
3448
'statusQueryGetUri': instance_status_url + "?" + required_query_string_parameters,
3549
'terminatePostUri': instance_status_url + "/terminate?reason={text}&" + required_query_string_parameters,
50+
'rewindPostUri': instance_status_url + "/rewind?reason={text}&" + required_query_string_parameters,
3651
'resumePostUri': instance_status_url + "/resume?reason={text}&" + required_query_string_parameters,
37-
'suspendPostUri': instance_status_url + "/suspend?reason={text}&" + required_query_string_parameters
38-
})
52+
'suspendPostUri': instance_status_url + "/suspend?reason={text}&" + required_query_string_parameters,
53+
}
54+
templates = management_urls or {}
55+
placeholder = templates.get("id") or _INSTANCE_ID_PLACEHOLDER
56+
encoded_instance_id = quote(instance_id, safe="")
57+
58+
urls = {'id': instance_id}
59+
for name, fallback_url in fallback_urls.items():
60+
template = templates.get(name)
61+
if not template:
62+
urls[name] = fallback_url
63+
continue
64+
65+
url = template.replace(placeholder, encoded_instance_id)
66+
if placeholder != _INSTANCE_ID_PLACEHOLDER:
67+
url = url.replace(_INSTANCE_ID_PLACEHOLDER, encoded_instance_id)
68+
urls[name] = replace_url_origin(url, request_origin)
69+
70+
super().__init__(urls)
3971

4072
def __str__(self) -> str:
4173
return json.dumps(self)
@@ -48,3 +80,24 @@ def urls(self) -> dict[str, Any]:
4880
def to_json(self) -> dict[str, Any]:
4981
"""Return the management URLs as a plain ``dict``."""
5082
return dict(self)
83+
84+
85+
def replace_url_origin(url: str, request_origin: str | None) -> str:
86+
if request_origin is None:
87+
return url
88+
89+
parsed_url = urlsplit(url)
90+
parsed_origin = urlsplit(request_origin)
91+
if not parsed_origin.scheme or not parsed_origin.netloc:
92+
raise ValueError(
93+
"request_origin must include both a scheme and an authority")
94+
if not parsed_url.scheme or not parsed_url.netloc:
95+
return url
96+
97+
return urlunsplit((
98+
parsed_origin.scheme,
99+
parsed_origin.netloc,
100+
parsed_url.path,
101+
parsed_url.query,
102+
parsed_url.fragment,
103+
))

0 commit comments

Comments
 (0)