from __future__ import annotations from ipaddress import IPv4Address, IPv6Address, ip_address from urllib.parse import urlsplit from starlette.requests import Request from app.core.config import AppSettings IPAddress = IPv4Address | IPv6Address def _parse_ip(value: str | None) -> IPAddress | None: if not value: return None candidate = value.strip().strip('"') if candidate.startswith("[") and "]" in candidate: candidate = candidate[1 : candidate.index("]")] if "%" in candidate: candidate = candidate.split("%", 1)[0] try: return ip_address(candidate) except ValueError: return None def _is_trusted_proxy(address: IPAddress | None, settings: AppSettings) -> bool: return address is not None and any(address in network for network in settings.proxy_networks) def client_ip(request: Request, settings: AppSettings) -> str: """Return the nearest untrusted address without trusting client-supplied proxy headers.""" peer = _parse_ip(request.client.host if request.client else None) if not _is_trusted_proxy(peer, settings): return str(peer) if peer is not None else "unknown" forwarded = request.headers.get("X-Forwarded-For") if not forwarded: return str(peer) chain = [_parse_ip(item) for item in forwarded.split(",")] if not chain or any(address is None for address in chain): return str(peer) addresses = [address for address in chain if address is not None] for address in reversed(addresses): if not _is_trusted_proxy(address, settings): return str(address) return str(addresses[0]) def is_loopback_client(request: Request, settings: AppSettings) -> bool: address = _parse_ip(client_ip(request, settings)) return address is not None and address.is_loopback def request_is_https(request: Request, settings: AppSettings) -> bool: if request.url.scheme.lower() == "https": return True peer = _parse_ip(request.client.host if request.client else None) if not _is_trusted_proxy(peer, settings): return False forwarded_proto = request.headers.get("X-Forwarded-Proto", "") values = [value.strip().lower() for value in forwarded_proto.split(",") if value.strip()] return len(values) == 1 and values[0] == "https" def origin_is_https(request: Request) -> bool: origin = request.headers.get("Origin") return bool(origin and urlsplit(origin).scheme.lower() == "https") def use_secure_cookie(request: Request, settings: AppSettings) -> bool: return settings.cookie_secure or request_is_https(request, settings) or origin_is_https(request)