fix: harden share retrieval and admin visibility (#482, #480)

This commit is contained in:
Lan
2026-07-10 17:39:53 +08:00
parent 36300ef54e
commit 8d7d856c62
8 changed files with 296 additions and 24 deletions
+2 -1
View File
@@ -90,7 +90,7 @@ async def get_expire_info(
def get_code_generate_type() -> str:
code_generate_type = getattr(settings, "code_generate_type", "number")
code_generate_type = getattr(settings, "code_generate_type", "secret")
if code_generate_type in {"secret", "string"}:
return "secret"
return "number"
@@ -127,5 +127,6 @@ async def calculate_file_hash(file: UploadFile, chunk_size=1024 * 1024) -> str:
ip_limit = {
"error": IPRateLimit(count=settings.errorCount, minutes=settings.errorMinute),
"metadata": IPRateLimit(count=settings.errorCount, minutes=settings.errorMinute),
"upload": IPRateLimit(count=settings.uploadCount, minutes=settings.uploadMinute),
}
+57 -18
View File
@@ -11,6 +11,7 @@ from typing import Optional, Tuple, Union
from fastapi import APIRouter, Form, Request, UploadFile, File, Depends, HTTPException
from pydantic import BaseModel, ValidationError
from starlette import status
from tortoise.expressions import Case, F, Q, When
from apps.admin.dependencies import share_required_login
from apps.base.models import FileCodes, UploadChunk, PresignUploadSession
@@ -29,7 +30,12 @@ from apps.base.utils import (
from core.response import APIResponse
from core.settings import settings
from core.storage import storages, FileStorageInterface
from core.utils import get_select_token, get_now, sanitize_filename
from core.utils import (
get_file_url as get_proxy_file_url,
get_select_token,
get_now,
sanitize_filename,
)
share_api = APIRouter(prefix="/share", tags=["分享"])
@@ -211,11 +217,25 @@ async def get_code_file_by_code(
return True, file_code
async def update_file_usage(file_code: FileCodes) -> None:
file_code.used_count += 1
if file_code.expired_count > 0:
file_code.expired_count -= 1
await file_code.save()
async def consume_file_usage(file_code: FileCodes) -> bool:
"""原子校验分享状态并记录一次实际领取。"""
now = await get_now()
eligible = (
Q(expired_count__gt=0)
| Q(expired_count__lt=0, expired_at__gt=now)
| Q(expired_count__lt=0, expired_at=None)
)
updated = await FileCodes.filter(Q(id=file_code.id) & eligible).update(
expired_count=Case(
When(expired_count__gt=0, then=F("expired_count") - 1),
default=F("expired_count"),
),
used_count=F("used_count") + 1,
)
if not updated:
return False
await file_code.refresh_from_db()
return True
def build_file_metadata(file_code: FileCodes) -> dict:
@@ -242,9 +262,13 @@ async def build_select_detail(
file_code: FileCodes, file_storage: FileStorageInterface
) -> dict:
metadata = build_file_metadata(file_code)
download_url = (
None if file_code.text is not None else await file_storage.get_file_url(file_code)
)
if file_code.text is not None:
download_url = None
elif file_code.expired_count >= 0:
# 有次数限制的文件必须经过下载接口,第三方直链无法阻止重复使用。
download_url = await get_proxy_file_url(file_code.code)
else:
download_url = await file_storage.get_file_url(file_code)
content = file_code.text if file_code.text is not None else None
return {
**metadata,
@@ -255,24 +279,28 @@ async def build_select_detail(
@share_api.get("/metadata/")
async def get_file_metadata(code: str, ip: str = Depends(ip_limit["error"])):
async def get_file_metadata(code: str, ip: str = Depends(ip_limit["metadata"])):
has, file_code = await get_code_file_by_code(code)
if not has:
ip_limit["error"].add_ip(ip)
ip_limit["metadata"].add_ip(ip)
return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
ip_limit["metadata"].add_ip(ip)
return APIResponse(detail=build_file_metadata(file_code))
@share_api.post("/metadata/")
async def post_file_metadata(data: SelectFileModel, ip: str = Depends(ip_limit["error"])):
async def post_file_metadata(
data: SelectFileModel, ip: str = Depends(ip_limit["metadata"])
):
has, file_code = await get_code_file_by_code(data.code)
if not has:
ip_limit["error"].add_ip(ip)
ip_limit["metadata"].add_ip(ip)
return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
ip_limit["metadata"].add_ip(ip)
return APIResponse(detail=build_file_metadata(file_code))
@@ -285,7 +313,8 @@ async def get_code_file(code: str, ip: str = Depends(ip_limit["error"])):
return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
await update_file_usage(file_code)
if not await consume_file_usage(file_code):
return APIResponse(code=404, detail="文件已过期")
return await file_storage.get_file_response(file_code)
@@ -298,8 +327,16 @@ async def select_file(data: SelectFileModel, ip: str = Depends(ip_limit["error"]
return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
await update_file_usage(file_code)
return APIResponse(detail=await build_select_detail(file_code, file_storage))
detail = await build_select_detail(file_code, file_storage)
download_url = detail.get("download_url")
consumes_on_download = isinstance(download_url, str) and download_url.startswith(
"/share/download?"
)
if not consumes_on_download and not await consume_file_usage(file_code):
return APIResponse(code=404, detail="文件已过期")
if not consumes_on_download:
detail.update(build_file_metadata(file_code))
return APIResponse(detail=detail)
@share_api.get("/download")
@@ -309,10 +346,12 @@ async def download_file(key: str, code: str, ip: str = Depends(ip_limit["error"]
if await get_select_token(normalized_code) != key:
ip_limit["error"].add_ip(ip)
raise HTTPException(status_code=403, detail="下载鉴权失败")
has, file_code = await get_code_file_by_code(normalized_code, False)
has, file_code = await get_code_file_by_code(normalized_code)
if not has:
return APIResponse(code=404, detail="文件不存在")
return APIResponse(code=404, detail=file_code)
assert isinstance(file_code, FileCodes)
if not await consume_file_usage(file_code):
return APIResponse(code=404, detail="文件已过期")
return (
APIResponse(detail=file_code.text)
if file_code.text
+2
View File
@@ -64,6 +64,8 @@ async def ensure_security_settings() -> None:
def _sync_ip_limits() -> None:
ip_limit["error"].minutes = settings.errorMinute
ip_limit["error"].count = settings.errorCount
ip_limit["metadata"].minutes = settings.errorMinute
ip_limit["metadata"].count = settings.errorCount
ip_limit["upload"].minutes = settings.uploadMinute
ip_limit["upload"].count = settings.uploadCount
+1 -1
View File
@@ -44,7 +44,7 @@ DEFAULT_CONFIG = {
"uploadSize": 1024 * 1024 * 10,
"allowed_file_types": ["*"],
"expireStyle": ["day", "hour", "minute", "forever", "count"],
"code_generate_type": "number",
"code_generate_type": "secret",
"uploadMinute": 1,
"enableChunk": 0,
"webdav_url": "",
+1
View File
@@ -28,6 +28,7 @@ async def delete_expire_files():
if not dirs and not files:
os.rmdir(root)
await ip_limit["error"].remove_expired_ip()
await ip_limit["metadata"].remove_expired_ip()
await ip_limit["upload"].remove_expired_ip()
expire_data = await FileCodes.filter(
Q(expired_at__lt=await get_now()) | Q(expired_count=0)
+14 -4
View File
@@ -33,6 +33,12 @@ from core.tasks import delete_expire_files, clean_incomplete_uploads
from core.version import APP_VERSION
def normalize_public_flag(value) -> int:
if isinstance(value, str):
return int(value.strip().lower() in {"1", "true", "on", "yes"})
return int(bool(value))
def build_public_config() -> dict:
return {
"name": settings.name,
@@ -45,7 +51,7 @@ def build_public_config() -> dict:
"openUpload": settings.openUpload,
"notify_title": settings.notify_title,
"notify_content": settings.notify_content,
"show_admin_address": settings.showAdminAddr,
"show_admin_address": normalize_public_flag(settings.showAdminAddr),
"max_save_seconds": settings.max_save_seconds,
}
@@ -61,7 +67,7 @@ def build_public_meta() -> dict:
"features": {
"chunkUpload": bool(settings.enableChunk),
"guestUpload": bool(settings.openUpload),
"adminAddressVisible": bool(settings.showAdminAddr),
"adminAddressVisible": bool(normalize_public_flag(settings.showAdminAddr)),
"expirationModes": settings.expireStyle,
},
"limits": {
@@ -153,7 +159,9 @@ def parse_setup_options(data: dict) -> dict:
if not expire_styles:
raise ValueError("至少需要选择一种过期方式")
code_generate_type = get_form_value(data, "code_generate_type", "number")
code_generate_type = get_form_value(
data, "code_generate_type", DEFAULT_CONFIG["code_generate_type"]
)
if code_generate_type not in {"number", "secret"}:
raise ValueError("提取码类型不正确")
@@ -222,7 +230,9 @@ def build_setup_page(error: str = "", form: dict | None = None) -> str:
chunk_checked = (
" checked" if normalize_bool_field(form, "enableChunk", False) else ""
)
code_generate_type = get_form_value(form, "code_generate_type", "number")
code_generate_type = get_form_value(
form, "code_generate_type", DEFAULT_CONFIG["code_generate_type"]
)
selected_expire_styles = get_form_list(form, "expireStyle") or list(
DEFAULT_CONFIG["expireStyle"]
)
+26
View File
@@ -0,0 +1,26 @@
import unittest
from core.settings import settings
from main import build_public_config, build_public_meta
class AdminAddressPublicConfigTests(unittest.TestCase):
def setUp(self):
self._original_user_config = dict(settings.user_config)
def tearDown(self):
settings.user_config = self._original_user_config
def test_admin_address_is_exposed_as_strict_binary_flag(self):
cases = ((0, 0), (1, 1), ("0", 0), ("1", 1))
for configured_value, expected_value in cases:
with self.subTest(configured_value=configured_value):
settings.showAdminAddr = configured_value
public_value = build_public_config()["show_admin_address"]
feature_value = build_public_meta()["features"]["adminAddressVisible"]
self.assertIs(type(public_value), int)
self.assertEqual(public_value, expected_value)
self.assertIs(feature_value, bool(expected_value))
+193
View File
@@ -0,0 +1,193 @@
import asyncio
import datetime
import unittest
from unittest.mock import patch
from tortoise import Tortoise
from apps.base.models import FileCodes
from apps.base import views
from core.settings import DEFAULT_CONFIG, settings
from core.utils import get_now, get_select_token
from main import build_setup_page, parse_setup_options
class FakeStorage:
direct_url_calls = 0
async def get_file_url(self, file_code):
self.direct_url_calls += 1
return "https://storage.example/reusable-url"
async def get_file_response(self, file_code):
return {"downloaded": file_code.code}
class FakeRateLimit:
def __init__(self):
self.seen_ips = []
def add_ip(self, ip):
self.seen_ips.append(ip)
class ShareUsageSecurityTests(unittest.TestCase):
def test_new_installations_default_to_secret_codes(self):
self.assertEqual(DEFAULT_CONFIG["code_generate_type"], "secret")
self.assertEqual(
parse_setup_options({"expireStyle": ["day"]})["code_generate_type"],
"secret",
)
self.assertIn('<option value="secret" selected>', build_setup_page())
def test_atomic_usage_download_and_metadata_limits(self):
asyncio.run(self._run_scenario())
async def _run_scenario(self):
original_config = dict(settings.user_config)
original_metadata_limit = views.ip_limit["metadata"]
settings.file_storage = "local"
settings.jwt_secret = "test-download-secret"
await Tortoise.init(
config={
"connections": {
"default": {
"engine": "tortoise.backends.sqlite",
"credentials": {"file_path": ":memory:"},
}
},
"apps": {
"models": {
"models": ["apps.base.models"],
"default_connection": "default",
}
},
"use_tz": False,
"timezone": "Asia/Shanghai",
}
)
await Tortoise.generate_schemas()
try:
await self._assert_count_consumption_is_atomic()
await self._assert_time_expiration_and_usage_are_atomic()
await self._assert_download_url_is_single_use()
await self._assert_limited_files_do_not_expose_direct_urls()
await self._assert_metadata_counts_all_attempts()
finally:
views.ip_limit["metadata"] = original_metadata_limit
settings.user_config = original_config
await Tortoise.close_connections()
async def _assert_count_consumption_is_atomic(self):
record = await FileCodes.create(
code="atomic",
prefix="atomic",
suffix=".txt",
expired_count=1,
expired_at=await get_now() + datetime.timedelta(days=1),
)
stale_copies = await asyncio.gather(
*(FileCodes.get(id=record.id) for _ in range(20))
)
consumed = await asyncio.gather(
*(views.consume_file_usage(item) for item in stale_copies)
)
await record.refresh_from_db()
self.assertEqual(sum(consumed), 1)
self.assertEqual(record.expired_count, 0)
self.assertEqual(record.used_count, 1)
async def _assert_time_expiration_and_usage_are_atomic(self):
valid = await FileCodes.create(
code="time-valid",
text="valid",
expired_count=-1,
expired_at=await get_now() + datetime.timedelta(minutes=5),
)
expired = await FileCodes.create(
code="time-expired",
text="expired",
expired_count=-1,
expired_at=await get_now() - datetime.timedelta(seconds=1),
)
valid_copies = await asyncio.gather(
*(FileCodes.get(id=valid.id) for _ in range(20))
)
consumed = await asyncio.gather(
*(views.consume_file_usage(item) for item in valid_copies)
)
self.assertTrue(all(consumed))
self.assertFalse(await views.consume_file_usage(expired))
await valid.refresh_from_db()
await expired.refresh_from_db()
self.assertEqual(valid.expired_count, -1)
self.assertEqual(valid.used_count, 20)
self.assertEqual(expired.used_count, 0)
async def _assert_download_url_is_single_use(self):
record = await FileCodes.create(
code="single-use",
prefix="report",
suffix=".pdf",
expired_count=1,
expired_at=await get_now() + datetime.timedelta(days=1),
)
key = await get_select_token(record.code)
with patch.dict(views.storages, {"local": FakeStorage}):
selected = await views.select_file(
data=views.SelectFileModel(code=record.code), ip="127.0.0.1"
)
await record.refresh_from_db()
self.assertEqual(selected.detail["expired_count"], 1)
self.assertEqual(record.expired_count, 1)
self.assertEqual(record.used_count, 0)
first = await views.download_file(key=key, code=record.code, ip="127.0.0.1")
second = await views.download_file(key=key, code=record.code, ip="127.0.0.1")
self.assertEqual(first, {"downloaded": record.code})
self.assertEqual(second.code, 404)
await record.refresh_from_db()
self.assertEqual(record.expired_count, 0)
self.assertEqual(record.used_count, 1)
async def _assert_limited_files_do_not_expose_direct_urls(self):
record = await FileCodes.create(
code="proxy-only",
prefix="archive",
suffix=".zip",
expired_count=2,
expired_at=await get_now() + datetime.timedelta(days=1),
)
storage = FakeStorage()
detail = await views.build_select_detail(record, storage)
self.assertTrue(detail["download_url"].startswith("/share/download?"))
self.assertEqual(storage.direct_url_calls, 0)
async def _assert_metadata_counts_all_attempts(self):
await FileCodes.create(
code="metadata",
text="hello",
expired_count=-1,
expired_at=None,
)
limiter = FakeRateLimit()
views.ip_limit["metadata"] = limiter
success = await views.get_file_metadata(code="metadata", ip="192.0.2.1")
missing = await views.get_file_metadata(code="missing", ip="192.0.2.1")
self.assertEqual(success.code, 200)
self.assertEqual(missing.code, 404)
self.assertEqual(limiter.seen_ips, ["192.0.2.1", "192.0.2.1"])
if __name__ == "__main__":
unittest.main()