diff --git a/mediaflow_proxy/mpd_processor.py b/mediaflow_proxy/mpd_processor.py index f1eb64d..ad460b2 100644 --- a/mediaflow_proxy/mpd_processor.py +++ b/mediaflow_proxy/mpd_processor.py @@ -7,7 +7,7 @@ from fastapi import Request, Response, HTTPException from mediaflow_proxy.configs import settings from mediaflow_proxy.drm.decrypter import decrypt_segment -from mediaflow_proxy.utils.http_utils import encode_mediaflow_proxy_url +from mediaflow_proxy.utils.http_utils import encode_mediaflow_proxy_url, get_original_scheme logger = logging.getLogger(__name__) @@ -103,10 +103,14 @@ def build_hls(mpd_dict: dict, request: Request, key_id: str = None, key: str = N video_profiles = {} audio_profiles = {} + # Get the base URL for the playlist_endpoint endpoint + proxy_url = request.url_for("playlist_endpoint") + proxy_url = str(proxy_url.replace(scheme=get_original_scheme(request))) + for profile in mpd_dict["profiles"]: query_params.update({"profile_id": profile["id"], "key_id": key_id or "", "key": key or ""}) playlist_url = encode_mediaflow_proxy_url( - str(request.url_for("playlist_endpoint")), + proxy_url, query_params=query_params, ) @@ -150,6 +154,10 @@ def build_hls_playlist(mpd_dict: dict, profiles: list[dict], request: Request) - current_time = datetime.now(timezone.utc) live_stream_delay = timedelta(seconds=settings.mpd_live_stream_delay) target_end_time = current_time - live_stream_delay + + proxy_url = request.url_for("segment_endpoint") + proxy_url = str(proxy_url.replace(scheme=get_original_scheme(request))) + for index, profile in enumerate(profiles): segments = profile["segments"] if not segments: @@ -189,7 +197,7 @@ def build_hls_playlist(mpd_dict: dict, profiles: list[dict], request: Request) - ) hls.append( encode_mediaflow_proxy_url( - str(request.url_for("segment_endpoint")), + proxy_url, query_params=query_params, ) ) diff --git a/mediaflow_proxy/utils/http_utils.py b/mediaflow_proxy/utils/http_utils.py index d06ab97..89d1fc2 100644 --- a/mediaflow_proxy/utils/http_utils.py +++ b/mediaflow_proxy/utils/http_utils.py @@ -218,3 +218,34 @@ def encode_mediaflow_proxy_url( base_url = parse.urljoin(mediaflow_proxy_url, endpoint) return f"{base_url}?{encoded_params}" + + +def get_original_scheme(request) -> str: + """ + Determines the original scheme (http or https) of the request. + + Args: + request (Request): The incoming HTTP request. + + Returns: + str: The original scheme ('http' or 'https') + """ + # Check the X-Forwarded-Proto header first + forwarded_proto = request.headers.get("X-Forwarded-Proto") + if forwarded_proto: + return forwarded_proto + + # Check if the request is secure + if request.url.scheme == "https" or request.headers.get("X-Forwarded-Ssl") == "on": + return "https" + + # Check for other common headers that might indicate HTTPS + if ( + request.headers.get("X-Forwarded-Ssl") == "on" + or request.headers.get("X-Forwarded-Protocol") == "https" + or request.headers.get("X-Url-Scheme") == "https" + ): + return "https" + + # Default to http if no indicators of https are found + return "http" diff --git a/mediaflow_proxy/utils/m3u8_processor.py b/mediaflow_proxy/utils/m3u8_processor.py index 0438fb8..5738d34 100644 --- a/mediaflow_proxy/utils/m3u8_processor.py +++ b/mediaflow_proxy/utils/m3u8_processor.py @@ -3,7 +3,7 @@ from urllib import parse from pydantic import HttpUrl -from mediaflow_proxy.utils.http_utils import encode_mediaflow_proxy_url +from mediaflow_proxy.utils.http_utils import encode_mediaflow_proxy_url, get_original_scheme class M3U8Processor: @@ -17,6 +17,7 @@ class M3U8Processor: """ self.request = request self.key_url = key_url + self.mediaflow_proxy_url = str(request.url_for("hls_stream_proxy").replace(scheme=get_original_scheme(request))) async def process_m3u8(self, content: str, base_url: str) -> str: """ @@ -75,7 +76,7 @@ class M3U8Processor: full_url = parse.urljoin(base_url, url) return encode_mediaflow_proxy_url( - str(self.request.url_for("hls_stream_proxy")), + self.mediaflow_proxy_url, "", full_url, query_params=dict(self.request.query_params),