mirror of
https://github.com/vastsa/FileCodeBox.git
synced 2026-08-31 09:43:50 +08:00
+2
-1
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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": "",
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"]
|
||||
)
|
||||
|
||||
@@ -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))
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user