Skip to content

Commit baa48f0

Browse files
committed
refactor: centralize response header filtering
1 parent d0e6b9a commit baa48f0

5 files changed

Lines changed: 48 additions & 73 deletions

File tree

route_helpers.py

Lines changed: 24 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
import logging
22
import os
33
from functools import wraps
4-
from typing import Any, Callable, Dict, Optional
4+
from typing import Any, Callable, Dict, Mapping, Optional
55
from urllib.parse import urlparse
66

77
from flask import Response, g, jsonify, redirect, request, url_for
@@ -16,6 +16,29 @@
1616
CORS_ALLOWED_METHODS = "GET, POST, PUT, DELETE, PATCH, OPTIONS"
1717
CORS_DEFAULT_HEADERS = "Authorization, Content-Type, Accept, Origin, X-Requested-With"
1818
LOCAL_DEVELOPMENT_HOSTS = {"localhost", "127.0.0.1", "::1"}
19+
HOP_BY_HOP_RESPONSE_HEADERS = frozenset(
20+
{
21+
"connection",
22+
"content-encoding",
23+
"content-length",
24+
"keep-alive",
25+
"proxy-authenticate",
26+
"proxy-authorization",
27+
"te",
28+
"trailer",
29+
"transfer-encoding",
30+
"upgrade",
31+
}
32+
)
33+
34+
35+
def copy_upstream_response_headers(upstream_headers: Mapping[str, Any]) -> Dict[str, Any]:
36+
"""Copy only response headers that are safe for Flask to emit downstream."""
37+
return {
38+
key: value
39+
for key, value in upstream_headers.items()
40+
if key.lower() not in HOP_BY_HOP_RESPONSE_HEADERS
41+
}
1942

2043

2144
def mask_secret(value: Optional[str]) -> str:

routes/core.py

Lines changed: 1 addition & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -12,27 +12,13 @@
1212
from config import Config
1313
from error_handlers import APIError, INTERNAL_ERROR_MESSAGE, get_request_id, internal_error_payload
1414
from proxy import PROVIDER_DETAILS
15-
from route_helpers import apply_cors_headers, check_provider, login_required
15+
from route_helpers import apply_cors_headers, check_provider, copy_upstream_response_headers, login_required
1616
from services.auth_service import AuthService
1717
from services.metrics_service import MetricsService
1818

1919
logger = logging.getLogger(__name__)
2020

2121

22-
HOP_BY_HOP_RESPONSE_HEADERS = {
23-
"connection",
24-
"content-encoding",
25-
"content-length",
26-
"keep-alive",
27-
"proxy-authenticate",
28-
"proxy-authorization",
29-
"te",
30-
"trailer",
31-
"transfer-encoding",
32-
"upgrade",
33-
}
34-
35-
3622
def is_safe_redirect_target(target: str | None) -> bool:
3723
"""Allow only local, absolute-path redirects after login."""
3824
if not target:
@@ -114,15 +100,6 @@ def build_openrouter_dashboard_headers() -> dict:
114100
return headers
115101

116102

117-
def copy_upstream_response_headers(upstream_headers: requests.structures.CaseInsensitiveDict) -> dict:
118-
"""Copy only response headers that are safe for Flask to emit downstream."""
119-
return {
120-
key: value
121-
for key, value in upstream_headers.items()
122-
if key.lower() not in HOP_BY_HOP_RESPONSE_HEADERS
123-
}
124-
125-
126103
def register_core_routes(app) -> None:
127104
@app.errorhandler(CSRFError)
128105
def handle_csrf_error(error: CSRFError):

routes/proxy.py

Lines changed: 2 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -10,32 +10,10 @@
1010
from error_handlers import APIError
1111
from providers.registry import get_adapter
1212
from proxy import PROVIDER_DETAILS
13-
from route_helpers import api_auth_required, login_required
13+
from route_helpers import api_auth_required, copy_upstream_response_headers, login_required
1414

1515
logger = logging.getLogger(__name__)
1616

17-
HOP_BY_HOP_RESPONSE_HEADERS = {
18-
"connection",
19-
"content-encoding",
20-
"content-length",
21-
"keep-alive",
22-
"proxy-authenticate",
23-
"proxy-authorization",
24-
"te",
25-
"trailer",
26-
"transfer-encoding",
27-
"upgrade",
28-
}
29-
30-
31-
def _safe_response_headers(headers):
32-
return {
33-
key: value
34-
for key, value in headers.items()
35-
if key.lower() not in HOP_BY_HOP_RESPONSE_HEADERS
36-
}
37-
38-
3917
def _dashboard_chat_completions_url(app, provider):
4018
if provider == "googleai":
4119
project_id = os.environ.get("PROJECT_ID")
@@ -286,7 +264,7 @@ def proxy_chat_completions():
286264
response.content,
287265
status=status_code,
288266
content_type=response.headers.get("content-type", "application/json"),
289-
headers=_safe_response_headers(response.headers),
267+
headers=copy_upstream_response_headers(response.headers),
290268
)
291269

292270
except APIError as error:

routes/unified.py

Lines changed: 2 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -7,33 +7,11 @@
77

88
from error_handlers import APIError
99
from providers.registry import get_adapter
10-
from route_helpers import api_auth_required, login_required
10+
from route_helpers import api_auth_required, copy_upstream_response_headers, login_required
1111
from services.auth_service import AuthService
1212
from services.model_registry import ModelRegistry
1313

1414

15-
HOP_BY_HOP_RESPONSE_HEADERS = {
16-
"connection",
17-
"content-encoding",
18-
"content-length",
19-
"keep-alive",
20-
"proxy-authenticate",
21-
"proxy-authorization",
22-
"te",
23-
"trailer",
24-
"transfer-encoding",
25-
"upgrade",
26-
}
27-
28-
29-
def _safe_response_headers(headers):
30-
return {
31-
key: value
32-
for key, value in headers.items()
33-
if key.lower() not in HOP_BY_HOP_RESPONSE_HEADERS
34-
}
35-
36-
3715
def _provider_token(auth_service_cls, provider: str) -> str:
3816
token = auth_service_cls.get_google_token() if provider == "googleai" else auth_service_cls.get_api_key(provider)
3917
if not token:
@@ -184,7 +162,7 @@ def unified_chat_completions():
184162
response.content,
185163
status=response.status_code,
186164
content_type=response.headers.get("content-type", "application/json"),
187-
headers=_safe_response_headers(response.headers),
165+
headers=copy_upstream_response_headers(response.headers),
188166
)
189167
except ValueError as error:
190168
raise APIError(str(error), status_code=400) from error

test_route_helpers.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
api_auth_required,
88
apply_cors_headers,
99
build_cors_preflight_response,
10+
copy_upstream_response_headers,
1011
extract_bearer_token,
1112
is_cors_origin_allowed,
1213
mask_authorization_header,
@@ -37,6 +38,24 @@ def test_extract_bearer_token_rejects_missing_or_wrong_scheme(self):
3738
self.assertIsNone(extract_bearer_token("Basic token-value"))
3839
self.assertIsNone(extract_bearer_token("Bearer "))
3940

41+
def test_copy_upstream_response_headers_drops_hop_by_hop_values(self):
42+
headers = copy_upstream_response_headers(
43+
{
44+
"Content-Type": "application/json",
45+
"Connection": "keep-alive",
46+
"Transfer-Encoding": "chunked",
47+
"X-Request-ID": "req_123",
48+
}
49+
)
50+
51+
self.assertEqual(
52+
headers,
53+
{
54+
"Content-Type": "application/json",
55+
"X-Request-ID": "req_123",
56+
},
57+
)
58+
4059

4160
class RouteHelperCorsTest(unittest.TestCase):
4261
def setUp(self):

0 commit comments

Comments
 (0)