Skip to content

Commit ae77026

Browse files
authored
Merge pull request #496 from vastsa/fix/pr494-security-patches
fix: harden share expiry, download token, quota SQL, and CORS
2 parents 59ed733 + a41693f commit ae77026

7 files changed

Lines changed: 223 additions & 17 deletions

File tree

‎apps/base/quota.py‎

Lines changed: 43 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,43 @@
88
from core.utils import get_now
99

1010

11+
def _detect_sql_dialect() -> str:
12+
"""识别当前默认连接的 SQL 方言。
13+
14+
优先使用 Tortoise capabilities.dialect;
15+
回退到 db_config 中的 engine 路径字符串。
16+
"""
17+
conn = connections.get("default")
18+
dialect = getattr(getattr(conn, "capabilities", None), "dialect", "") or ""
19+
if dialect:
20+
return dialect.lower()
21+
22+
try:
23+
engine = str(connections.db_config.get("default", {}).get("engine", "")).lower()
24+
except Exception:
25+
engine = ""
26+
if "postgres" in engine or "asyncpg" in engine or "psycopg" in engine:
27+
return "postgres"
28+
if "mysql" in engine:
29+
return "mysql"
30+
if "sqlite" in engine:
31+
return "sqlite"
32+
return "sqlite"
33+
34+
35+
def _sql_placeholders(count: int) -> list[str]:
36+
"""根据数据库方言生成参数占位符(多数据库兼容)。
37+
38+
SQLite 用 ?,PostgreSQL 用 $1/$2/...,MySQL 用 %s。
39+
"""
40+
dialect = _detect_sql_dialect()
41+
if dialect in {"postgres", "postgresql"}:
42+
return [f"${i}" for i in range(1, count + 1)]
43+
if dialect == "mysql":
44+
return ["%s"] * count
45+
return ["?"] * count
46+
47+
1148
def get_storage_limit() -> int:
1249
try:
1350
return max(0, int(getattr(settings, "storageLimit", 0)))
@@ -49,15 +86,16 @@ async def reserve_storage(token: str, size: int, ttl_seconds: int) -> None:
4986
raise HTTPException(status_code=409, detail="上传容量预留信息不一致")
5087

5188
try:
89+
ph = _sql_placeholders(6)
5290
affected, _ = await conn.execute_query(
53-
"""
91+
f"""
5492
INSERT INTO storagereservation (token, size, expires_at)
55-
SELECT ?, ?, ?
93+
SELECT {ph[0]}, {ph[1]}, {ph[2]}
5694
WHERE (
5795
COALESCE((SELECT SUM(size) FROM filecodes), 0)
58-
+ COALESCE((SELECT SUM(size) FROM storagereservation WHERE expires_at > ?), 0)
59-
+ ?
60-
) <= ?
96+
+ COALESCE((SELECT SUM(size) FROM storagereservation WHERE expires_at > {ph[3]}), 0)
97+
+ {ph[4]}
98+
) <= {ph[5]}
6199
""",
62100
[token, requested_size, expires_at, now, requested_size, limit],
63101
)

‎apps/base/utils.py‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,13 @@
1919
)
2020

2121

22+
def validate_expire_style(expire_style: str) -> str:
23+
"""校验过期方式是否在管理员配置的白名单内。"""
24+
if expire_style not in settings.expireStyle:
25+
raise HTTPException(status_code=400, detail="过期时间类型错误")
26+
return expire_style
27+
28+
2229
async def get_file_path_name(file: UploadFile) -> Tuple[str, str, str, str, str]:
2330
today = await get_now()
2431
storage_path = settings.storage_path.strip("/")

‎apps/base/views.py‎

Lines changed: 11 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
get_file_path_name,
2929
ip_limit,
3030
get_chunk_file_path_name,
31+
validate_expire_style,
3132
)
3233
from core.response import APIResponse
3334
from core.settings import settings
@@ -152,6 +153,7 @@ async def share_text(
152153
expire_style: str = Form(default="day"),
153154
ip: str = Depends(ip_limit["upload"]),
154155
):
156+
validate_expire_style(expire_style)
155157
text_size = len(text.encode("utf-8"))
156158
max_txt_size = 222 * 1024
157159
if text_size > max_txt_size:
@@ -187,8 +189,7 @@ async def share_file(
187189
):
188190
file_size = await validate_file_size(file, settings.uploadSize)
189191
validate_file_type(file.filename or "", file.content_type)
190-
if expire_style not in settings.expireStyle:
191-
raise HTTPException(status_code=400, detail="过期时间类型错误")
192+
validate_expire_style(expire_style)
192193
path, suffix, prefix, uuid_file_name, save_path = await get_file_path_name(file)
193194
reservation_token = f"file:{uuid.uuid4().hex}"
194195
await reserve_storage(reservation_token, file_size, ttl_seconds=3600)
@@ -374,7 +375,12 @@ async def select_file(data: SelectFileModel, ip: str = Depends(ip_limit["error"]
374375
async def download_file(key: str, code: str, ip: str = Depends(ip_limit["error"])):
375376
file_storage: FileStorageInterface = storages[settings.file_storage]()
376377
normalized_code = normalize_share_code(code)
377-
if await get_select_token(normalized_code) != key:
378+
# 同时接受当前窗口与上一窗口 token,避免时间窗边界竞态导致偶发 403
379+
valid_keys = {
380+
await get_select_token(normalized_code, offset=0),
381+
await get_select_token(normalized_code, offset=1),
382+
}
383+
if key not in valid_keys:
378384
ip_limit["error"].add_ip(ip)
379385
raise HTTPException(status_code=403, detail="下载鉴权失败")
380386
has, file_code = await get_code_file_by_code(normalized_code)
@@ -659,6 +665,7 @@ async def complete_upload(
659665
chunk_info = await UploadChunk.filter(upload_id=upload_id, chunk_index=-1).first()
660666
if not chunk_info:
661667
raise HTTPException(status.HTTP_404_NOT_FOUND, detail="上传会话不存在")
668+
validate_expire_style(data.expire_style)
662669
await reserve_storage(
663670
f"chunk:{upload_id}",
664671
chunk_info.file_size,
@@ -775,8 +782,7 @@ async def presign_upload_init(
775782
403,
776783
f"文件大小超过限制,最大为 {settings.uploadSize / (1024 * 1024):.2f} MB",
777784
)
778-
if data.expire_style not in settings.expireStyle:
779-
raise HTTPException(400, "过期时间类型错误")
785+
validate_expire_style(data.expire_style)
780786

781787
upload_id = uuid.uuid4().hex
782788
reservation_token = f"presign:{upload_id}"

‎core/storage.py‎

Lines changed: 10 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,7 @@
2525
from apps.base.models import FileCodes, UploadChunk
2626
from core.utils import get_file_url, sanitize_filename
2727
from fastapi.responses import FileResponse, StreamingResponse
28+
from starlette.background import BackgroundTask
2829

2930

3031
class FileStorageInterface:
@@ -395,7 +396,9 @@ async def stream_generator():
395396
return StreamingResponse(
396397
stream_generator(),
397398
media_type="application/octet-stream",
398-
headers=headers
399+
headers=headers,
400+
# 兜底关闭会话:客户端中断时与 generator finally 双保险
401+
background=BackgroundTask(session.close),
399402
)
400403
except HTTPException:
401404
raise
@@ -733,7 +736,9 @@ async def stream_generator():
733736
return StreamingResponse(
734737
stream_generator(),
735738
media_type="application/octet-stream",
736-
headers=headers
739+
headers=headers,
740+
# 兜底关闭会话:客户端中断时与 generator finally 双保险
741+
background=BackgroundTask(session.close),
737742
)
738743
except HTTPException:
739744
raise
@@ -1209,7 +1214,9 @@ async def stream_generator():
12091214
return StreamingResponse(
12101215
stream_generator(),
12111216
media_type="application/octet-stream",
1212-
headers=headers
1217+
headers=headers,
1218+
# 兜底关闭会话:客户端中断时与 generator finally 双保险
1219+
background=BackgroundTask(session.close),
12131220
)
12141221
except aiohttp.ClientError as e:
12151222
raise HTTPException(

‎core/utils.py‎

Lines changed: 8 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -41,17 +41,22 @@ async def get_now():
4141
return datetime.datetime.now(datetime.timezone(datetime.timedelta(hours=8)))
4242

4343

44-
async def get_select_token(code: str):
44+
async def get_select_token(code: str, offset: int = 0):
4545
"""
4646
获取下载token
47-
:param code:
47+
:param code: 取件码
48+
:param offset: 时间窗口偏移(0=当前窗口,1=上一个窗口)。
49+
用于兼容窗口边界竞态:用户在某窗口末尾获取的 token,
50+
请求到达服务器时可能已进入下一窗口。
4851
:return:
4952
"""
5053
token = getattr(settings, "jwt_secret", "")
5154
if not token:
5255
raise RuntimeError("应用签名密钥未初始化")
56+
# 每个窗口约 1000 秒;offset 允许校验上一窗口,避免边界竞态
57+
time_factor = int(time.time() / 1000) - max(0, int(offset))
5358
return hashlib.sha256(
54-
f"{code}{int(time.time() / 1000)}000{token}".encode()
59+
f"{code}{time_factor}000{token}".encode()
5560
).hexdigest()
5661

5762

‎main.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -767,7 +767,9 @@ async def refresh_settings_middleware(request, call_next):
767767
app.add_middleware(
768768
CORSMiddleware,
769769
allow_origins=["*"],
770-
allow_credentials=True,
770+
# 前端使用 Bearer Token,不依赖 Cookie credentials。
771+
# allow_origins=["*"] 与 allow_credentials=True 组合不符合 CORS 规范。
772+
allow_credentials=False,
771773
allow_methods=["*"],
772774
allow_headers=["*"],
773775
)
Lines changed: 141 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,141 @@
1+
import asyncio
2+
import hashlib
3+
import time
4+
import unittest
5+
from unittest.mock import patch
6+
7+
from fastapi import HTTPException
8+
from tortoise import Tortoise
9+
10+
from apps.base import views
11+
from apps.base.models import FileCodes
12+
from apps.base.quota import _detect_sql_dialect, _sql_placeholders, reserve_storage
13+
from apps.base.utils import validate_expire_style
14+
from core.settings import settings
15+
from core.utils import get_select_token
16+
17+
18+
class FakeStorage:
19+
async def get_file_url(self, file_code):
20+
return f"https://example.invalid/{file_code.code}"
21+
22+
async def get_file_response(self, file_code):
23+
return {"downloaded": file_code.code}
24+
25+
26+
class SecurityPatchTests(unittest.TestCase):
27+
def test_validate_expire_style_rejects_unknown_mode(self):
28+
original = dict(settings.user_config)
29+
try:
30+
settings.expireStyle = ["day", "count"]
31+
self.assertEqual(validate_expire_style("day"), "day")
32+
with self.assertRaises(HTTPException) as ctx:
33+
validate_expire_style("forever")
34+
self.assertEqual(ctx.exception.status_code, 400)
35+
finally:
36+
settings.user_config = original
37+
38+
def test_sql_placeholders_follow_dialect(self):
39+
asyncio.run(self._assert_sql_placeholders())
40+
41+
async def _assert_sql_placeholders(self):
42+
await Tortoise.init(
43+
config={
44+
"connections": {
45+
"default": {
46+
"engine": "tortoise.backends.sqlite",
47+
"credentials": {"file_path": ":memory:"},
48+
}
49+
},
50+
"apps": {
51+
"models": {
52+
"models": ["apps.base.models"],
53+
"default_connection": "default",
54+
}
55+
},
56+
"use_tz": False,
57+
"timezone": "Asia/Shanghai",
58+
}
59+
)
60+
try:
61+
self.assertEqual(_detect_sql_dialect(), "sqlite")
62+
self.assertEqual(_sql_placeholders(3), ["?", "?", "?"])
63+
64+
with patch(
65+
"apps.base.quota._detect_sql_dialect", return_value="postgres"
66+
):
67+
self.assertEqual(_sql_placeholders(3), ["$1", "$2", "$3"])
68+
with patch("apps.base.quota._detect_sql_dialect", return_value="mysql"):
69+
self.assertEqual(_sql_placeholders(2), ["%s", "%s"])
70+
71+
# sqlite 下 reserve_storage 仍可正常工作
72+
settings.storageLimit = 100
73+
await Tortoise.generate_schemas()
74+
await reserve_storage("patch-token", 10, 300)
75+
finally:
76+
settings.storageLimit = 0
77+
await Tortoise.close_connections()
78+
79+
def test_download_token_accepts_previous_window(self):
80+
asyncio.run(self._assert_download_token_window())
81+
82+
async def _assert_download_token_window(self):
83+
original = dict(settings.user_config)
84+
settings.jwt_secret = "test-secret-key-for-token-window"
85+
now = 2_000_500 # 落在窗口 2000 内
86+
code = "ABCDE"
87+
with patch("core.utils.time.time", return_value=now):
88+
current = await get_select_token(code, offset=0)
89+
previous = await get_select_token(code, offset=1)
90+
91+
expected_current = hashlib.sha256(
92+
f"{code}{int(now / 1000)}000{settings.jwt_secret}".encode()
93+
).hexdigest()
94+
expected_previous = hashlib.sha256(
95+
f"{code}{int(now / 1000) - 1}000{settings.jwt_secret}".encode()
96+
).hexdigest()
97+
self.assertEqual(current, expected_current)
98+
self.assertEqual(previous, expected_previous)
99+
100+
await Tortoise.init(
101+
config={
102+
"connections": {
103+
"default": {
104+
"engine": "tortoise.backends.sqlite",
105+
"credentials": {"file_path": ":memory:"},
106+
}
107+
},
108+
"apps": {
109+
"models": {
110+
"models": ["apps.base.models"],
111+
"default_connection": "default",
112+
}
113+
},
114+
"use_tz": False,
115+
"timezone": "Asia/Shanghai",
116+
}
117+
)
118+
await Tortoise.generate_schemas()
119+
try:
120+
await FileCodes.create(
121+
code=code,
122+
prefix="demo",
123+
suffix=".txt",
124+
expired_count=-1,
125+
size=1,
126+
)
127+
# 模拟请求进入下一窗口,但仍携带上一窗口 token
128+
next_window = now + 1000
129+
with patch("core.utils.time.time", return_value=next_window):
130+
with patch.dict(views.storages, {"local": FakeStorage}):
131+
result = await views.download_file(
132+
key=current, code=code, ip="127.0.0.1"
133+
)
134+
self.assertEqual(result, {"downloaded": code})
135+
finally:
136+
settings.user_config = original
137+
await Tortoise.close_connections()
138+
139+
140+
if __name__ == "__main__":
141+
unittest.main()

0 commit comments

Comments
 (0)