|
1 | | -from typing import Dict, Union |
| 1 | +from ipaddress import ip_address, ip_network |
| 2 | +from typing import Dict, Iterable, Union |
2 | 3 | from datetime import datetime, timedelta |
3 | 4 | from fastapi import HTTPException, Request |
4 | 5 |
|
| 6 | +from core.settings import settings |
| 7 | + |
| 8 | + |
| 9 | +def _iter_trusted_proxies() -> Iterable[str]: |
| 10 | + trusted_proxies = getattr(settings, "trustedProxies", []) |
| 11 | + if isinstance(trusted_proxies, str): |
| 12 | + trusted_proxies = [item.strip() for item in trusted_proxies.split(",")] |
| 13 | + return [item for item in trusted_proxies if item] |
| 14 | + |
| 15 | + |
| 16 | +def _is_trusted_proxy(host: str) -> bool: |
| 17 | + try: |
| 18 | + remote_addr = ip_address(host) |
| 19 | + except ValueError: |
| 20 | + return False |
| 21 | + |
| 22 | + for proxy in _iter_trusted_proxies(): |
| 23 | + try: |
| 24 | + if remote_addr in ip_network(proxy, strict=False): |
| 25 | + return True |
| 26 | + except ValueError: |
| 27 | + continue |
| 28 | + return False |
| 29 | + |
| 30 | + |
| 31 | +def _get_forwarded_for_ip(header_value: str, fallback_ip: str) -> str: |
| 32 | + forwarded_chain = [ |
| 33 | + item.strip() for item in header_value.split(",") if item.strip() |
| 34 | + ] |
| 35 | + if not forwarded_chain: |
| 36 | + return fallback_ip |
| 37 | + |
| 38 | + for candidate in reversed(forwarded_chain): |
| 39 | + try: |
| 40 | + ip_address(candidate) |
| 41 | + except ValueError: |
| 42 | + return fallback_ip |
| 43 | + if not _is_trusted_proxy(candidate): |
| 44 | + return candidate |
| 45 | + return forwarded_chain[0] |
| 46 | + |
| 47 | + |
| 48 | +def get_client_ip(request: Request) -> str: |
| 49 | + client_host = request.client.host if request.client else "unknown" |
| 50 | + if not _is_trusted_proxy(client_host): |
| 51 | + return client_host |
| 52 | + |
| 53 | + forwarded_for = request.headers.get("X-Forwarded-For") |
| 54 | + if forwarded_for: |
| 55 | + return _get_forwarded_for_ip(forwarded_for, client_host) |
| 56 | + |
| 57 | + real_ip = request.headers.get("X-Real-IP") |
| 58 | + if real_ip: |
| 59 | + try: |
| 60 | + ip_address(real_ip) |
| 61 | + except ValueError: |
| 62 | + return client_host |
| 63 | + return real_ip |
| 64 | + |
| 65 | + return client_host |
| 66 | + |
5 | 67 |
|
6 | 68 | class IPRateLimit: |
7 | 69 | def __init__(self, count: int, minutes: int): |
@@ -35,11 +97,7 @@ async def remove_expired_ip(self) -> None: |
35 | 97 | } |
36 | 98 |
|
37 | 99 | def __call__(self, request: Request) -> str: |
38 | | - ip = ( |
39 | | - request.headers.get("X-Real-IP") |
40 | | - or request.headers.get("X-Forwarded-For") |
41 | | - or request.client.host |
42 | | - ) |
| 100 | + ip = get_client_ip(request) |
43 | 101 | if not self.check_ip(ip): |
44 | 102 | raise HTTPException(status_code=423, detail="请求次数过多,请稍后再试") |
45 | 103 | return ip |
0 commit comments