Skip to content

Commit 1b6d8e7

Browse files
committed
fix: harden share rate limit and code generation (#479)
1 parent eb6fc62 commit 1b6d8e7

4 files changed

Lines changed: 88 additions & 13 deletions

File tree

‎apps/base/dependencies.py‎

Lines changed: 64 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,69 @@
1-
from typing import Dict, Union
1+
from ipaddress import ip_address, ip_network
2+
from typing import Dict, Iterable, Union
23
from datetime import datetime, timedelta
34
from fastapi import HTTPException, Request
45

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+
567

668
class IPRateLimit:
769
def __init__(self, count: int, minutes: int):
@@ -35,11 +97,7 @@ async def remove_expired_ip(self) -> None:
3597
}
3698

3799
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)
43101
if not self.check_ip(ip):
44102
raise HTTPException(status_code=423, detail="请求次数过多,请稍后再试")
45103
return ip

‎apps/base/utils.py‎

Lines changed: 19 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -76,8 +76,6 @@ async def get_expire_info(
7676
expired_at, extra = result
7777
if expire_style == "count":
7878
expired_count = extra
79-
elif expire_style == "forever":
80-
code = await get_random_code(style="string")
8179
else:
8280
expired_at = result
8381
if expired_at and expired_at - now > max_timedelta:
@@ -91,9 +89,26 @@ async def get_expire_info(
9189
return expired_at, expired_count, used_count, code
9290

9391

94-
async def get_random_code(style: str = "num") -> str:
92+
def get_code_generate_type() -> str:
93+
code_generate_type = getattr(settings, "code_generate_type", "number")
94+
if code_generate_type in {"secret", "string"}:
95+
return "secret"
96+
return "number"
97+
98+
99+
async def get_random_code(style: str | None = None) -> str:
100+
code_style = style or get_code_generate_type()
101+
if code_style == "num":
102+
code_style = "number"
103+
if code_style == "string":
104+
code_style = "secret"
105+
95106
while True:
96-
code = await get_random_num() if style == "num" else await get_random_string()
107+
code = (
108+
await get_random_num()
109+
if code_style == "number"
110+
else await get_random_string()
111+
)
97112
if not await FileCodes.filter(code=code).exists():
98113
return str(code)
99114

‎core/settings.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -43,6 +43,7 @@
4343
"uploadSize": 1024 * 1024 * 10,
4444
"allowed_file_types": ["*"],
4545
"expireStyle": ["day", "hour", "minute", "forever", "count"],
46+
"code_generate_type": "number",
4647
"uploadMinute": 1,
4748
"enableChunk": 0,
4849
"webdav_url": "",
@@ -73,6 +74,7 @@
7374
"serverPort": 12345,
7475
"showAdminAddr": 0,
7576
"robotsText": "User-agent: *\nDisallow: /",
77+
"trustedProxies": [],
7678
}
7779

7880

‎core/utils.py‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -5,8 +5,8 @@
55
import datetime
66
import hashlib
77
import os
8-
import random
98
import re
9+
import secrets
1010
import string
1111
import time
1212

@@ -18,7 +18,7 @@ async def get_random_num():
1818
获取随机数
1919
:return:
2020
"""
21-
return random.randint(10000, 99999)
21+
return secrets.randbelow(90000) + 10000
2222

2323

2424
r_s = string.ascii_uppercase + string.digits
@@ -29,7 +29,7 @@ async def get_random_string():
2929
获取随机字符串
3030
:return:
3131
"""
32-
return "".join(random.choice(r_s) for _ in range(5))
32+
return "".join(secrets.choice(r_s) for _ in range(5))
3333

3434

3535
async def get_now():

0 commit comments

Comments
 (0)