Skip to content

Commit 5acfe36

Browse files
committed
fix: handle_response_headers
1 parent 016c3fe commit 5acfe36

2 files changed

Lines changed: 38 additions & 29 deletions

File tree

projects/fal_client/src/fal_client/_headers.py

Lines changed: 26 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,23 @@
11
from __future__ import annotations
22

3-
from typing import Literal, Union, get_args
3+
from typing import Literal, Union, get_args, Optional, Any
4+
5+
from httpx import Headers
6+
7+
try:
8+
from fal.ref import get_current_app
9+
except ImportError:
10+
11+
def get_current_app() -> Optional[Any]:
12+
return None
13+
14+
15+
def _current_fal_app_request() -> Optional[Any]:
16+
"""Get the current request if we are running in a fal app."""
17+
if (app := get_current_app()) is not None and app.current_request is not None:
18+
return app.current_request
19+
return None
20+
421

522
MIN_REQUEST_TIMEOUT_SECONDS = 1 # Minimum allowed request timeout in seconds
623

@@ -54,15 +71,13 @@ def add_priority_header(priority: Priority, headers: dict[str, str]) -> None:
5471
headers[QUEUE_PRIORITY_HEADER] = priority
5572

5673

57-
def add_forwarded_headers(request_headers, _headers: dict[str, str]) -> None:
58-
if cdn_token := request_headers.get("x-fal-cdn-token"):
59-
_headers["x-fal-forwarded-cdn-token"] = cdn_token
74+
def add_fal_app_context_headers(headers: dict[str, str]) -> None:
75+
if request := _current_fal_app_request():
76+
if cdn_token := request.headers.get("x-fal-cdn-token"):
77+
headers["x-fal-cdn-token"] = cdn_token
6078

61-
forwarded_request_ids = []
62-
if request_ids := request_headers.get("x-fal-forwarded-request-id"):
63-
forwarded_request_ids.extend(request_ids.split(","))
64-
if request_id := request_headers.get("x-fal-request-id"):
65-
forwarded_request_ids.append(request_id)
6679

67-
if forwarded_request_ids:
68-
_headers["x-fal-forwarded-request-id"] = ",".join(forwarded_request_ids)
80+
def handle_response_headers(response_headers: Headers) -> None:
81+
if request := _current_fal_app_request():
82+
if cdn_token := response_headers.get("x-fal-cdn-token"):
83+
request.headers["x-fal-cdn-token"] = cdn_token

projects/fal_client/src/fal_client/client.py

Lines changed: 12 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -49,9 +49,10 @@
4949
add_priority_header,
5050
add_timeout_header,
5151
add_hint_header,
52+
add_fal_app_context_headers,
53+
handle_response_headers,
5254
REQUEST_TIMEOUT_TYPE_HEADER,
5355
REQUEST_TIMEOUT_HEADER,
54-
add_forwarded_headers,
5556
)
5657

5758
if TYPE_CHECKING:
@@ -63,14 +64,6 @@
6364
if TYPE_CHECKING:
6465
from PIL import Image
6566

66-
try:
67-
from fal.ref import get_current_app
68-
except ImportError:
69-
70-
def get_current_app() -> Optional[Any]:
71-
return None
72-
73-
7467
AnyJSON = Dict[str, Any]
7568
UploadRepositoryId = Literal["fal_v3", "cdn", "fal"]
7669

@@ -1609,8 +1602,7 @@ async def run(
16091602
if start_timeout is not None:
16101603
add_timeout_header(start_timeout, _headers)
16111604

1612-
if (app := get_current_app()) is not None and app.current_request is not None:
1613-
add_forwarded_headers(app.current_request.headers, _headers)
1605+
add_fal_app_context_headers(_headers)
16141606

16151607
response = await _async_maybe_retry_request(
16161608
self._client,
@@ -1620,8 +1612,9 @@ async def run(
16201612
timeout=timeout,
16211613
headers=_headers,
16221614
)
1623-
16241615
_raise_for_status(response)
1616+
handle_response_headers(response.headers)
1617+
16251618
return response.json()
16261619

16271620
async def submit(
@@ -1665,8 +1658,7 @@ async def submit(
16651658
if start_timeout is not None:
16661659
add_timeout_header(start_timeout, _headers)
16671660

1668-
if (app := get_current_app()) is not None and app.current_request is not None:
1669-
add_forwarded_headers(app.current_request.headers, _headers)
1661+
add_fal_app_context_headers(_headers)
16701662

16711663
response = await _async_maybe_retry_request(
16721664
self._client,
@@ -1677,6 +1669,7 @@ async def submit(
16771669
headers=_headers,
16781670
)
16791671
_raise_for_status(response)
1672+
handle_response_headers(response.headers)
16801673

16811674
data = response.json()
16821675
return AsyncRequestHandle(
@@ -2111,8 +2104,7 @@ def run(
21112104
if start_timeout is not None:
21122105
add_timeout_header(start_timeout, _headers)
21132106

2114-
if (app := get_current_app()) is not None and app.current_request is not None:
2115-
add_forwarded_headers(app.current_request.headers, _headers)
2107+
add_fal_app_context_headers(_headers)
21162108

21172109
response = _maybe_retry_request(
21182110
self._client,
@@ -2123,6 +2115,8 @@ def run(
21232115
headers=_headers,
21242116
)
21252117
_raise_for_status(response)
2118+
handle_response_headers(response.headers)
2119+
21262120
return response.json()
21272121

21282122
def submit(
@@ -2163,8 +2157,7 @@ def submit(
21632157
if start_timeout is not None:
21642158
add_timeout_header(start_timeout, _headers)
21652159

2166-
if (app := get_current_app()) is not None and app.current_request is not None:
2167-
add_forwarded_headers(app.current_request.headers, _headers)
2160+
add_fal_app_context_headers(_headers)
21682161

21692162
response = _maybe_retry_request(
21702163
self._client,
@@ -2175,6 +2168,7 @@ def submit(
21752168
headers=_headers,
21762169
)
21772170
_raise_for_status(response)
2171+
handle_response_headers(response.headers)
21782172

21792173
data = response.json()
21802174
return SyncRequestHandle(

0 commit comments

Comments
 (0)