Refactor: drop aiohttp serve_static in favor of our own to have more control

This commit is contained in:
arabcoders 2025-11-19 18:09:26 +03:00
parent ab56b875f3
commit 7569ff252a
7 changed files with 238 additions and 42 deletions

24
API.md
View file

@ -54,6 +54,7 @@ This document describes the available endpoints and their usage. All endpoints r
- [POST /api/file/actions](#post-apifileactions) - [POST /api/file/actions](#post-apifileactions)
- [POST /api/file/download](#post-apifiledownload) - [POST /api/file/download](#post-apifiledownload)
- [GET /api/file/download/{token}](#get-apifiledownloadtoken) - [GET /api/file/download/{token}](#get-apifiledownloadtoken)
- [GET /api/download/{filename}](#get-apidownloadfilename)
- [GET /api/random/background](#get-apirandombackground) - [GET /api/random/background](#get-apirandombackground)
- [GET /api/presets](#get-apipresets) - [GET /api/presets](#get-apipresets)
- [GET /api/dl\_fields](#get-apidl_fields) - [GET /api/dl\_fields](#get-apidl_fields)
@ -1327,6 +1328,27 @@ or an error:
--- ---
### GET /api/download/{filename}
**Purpose**: Serve downloaded files directly from the download path.
**Path Parameter**:
- `filename`: Relative path to the file within the download directory (URL-encoded).
**Response**:
- `200 OK` with file content and appropriate headers
- `302 Found` redirect if the file path needs normalization
- `404 Not Found` if the file doesn't exist
- **Error Responses**
When an error occurs, responses follow a structure similar to:
```json
{ "error": "Description of the error" }
```
with an appropriate HTTP status code.
---
### GET /api/random/background ### GET /api/random/background
**Purpose**: Get a random background image from configured backends. **Purpose**: Get a random background image from configured backends.
@ -1334,7 +1356,7 @@ or an error:
- `force=true` (optional) - Force fetch a new image instead of using cache. - `force=true` (optional) - Force fetch a new image instead of using cache.
**Response**: **Response**:
Binary image data with appropriate `Content-Type` header. Binary image data with appropriate headers
--- ---

View file

@ -17,9 +17,8 @@ from .config import Config
from .DownloadQueue import DownloadQueue from .DownloadQueue import DownloadQueue
from .encoder import Encoder from .encoder import Encoder
from .Events import EventBus from .Events import EventBus
from .ffprobe import ffprobe
from .router import RouteType, get_routes from .router import RouteType, get_routes
from .Utils import decrypt_data, encrypt_data, get_file, get_mime_type, load_modules from .Utils import decrypt_data, encrypt_data, get_file, load_modules
LOG: logging.Logger = logging.getLogger("http_api") LOG: logging.Logger = logging.getLogger("http_api")
@ -144,8 +143,6 @@ class HttpAPI:
app.router.add_route("OPTIONS", route.path, handler=options_handler, name=f"{route.name}_opts") app.router.add_route("OPTIONS", route.path, handler=options_handler, name=f"{route.name}_opts")
registered_options.append(route.path) registered_options.append(route.path)
app.router.add_static(f"{base_path}/api/download/", self.config.download_path, name="download_static")
@staticmethod @staticmethod
def basic_auth(username: str, password: str) -> Awaitable: def basic_auth(username: str, password: str) -> Awaitable:
""" """
@ -294,14 +291,6 @@ class HttpAPI:
return new_response return new_response
if request.path.startswith(static_path) and isinstance(response, web.FileResponse):
try:
ff_info = await ffprobe(response._path)
mime_type = get_mime_type(ff_info.get("metadata", {}), response._path)
response.content_type = mime_type
except Exception:
pass
return response return response
return middleware_handler return middleware_handler

View file

@ -697,10 +697,7 @@ class HandleTask:
s: dict[list[str]] = {"h": [], "d": [], "u": [], "f": []} s: dict[list[str]] = {"h": [], "d": [], "u": [], "f": []}
for task in self._tasks.get_all(): for task in self._tasks.get_all():
if not task.enabled: if not task.enabled or not task.handler_enabled:
continue
if not task.handler_enabled:
s["d"].append(task.name) s["d"].append(task.name)
continue continue
@ -725,7 +722,7 @@ class HandleTask:
if len(self._tasks.get_all()) > 0: if len(self._tasks.get_all()) > 0:
LOG.info( LOG.info(
f"Task Handler summary: Handled: {len(s['h'])}, Unhandled: {len(s['u'])}, Disabled: {len(s['d'])}, Failed: {len(s['f'])}." f"Tasks handler summary: Handled: {len(s['h'])}, Unhandled: {len(s['u'])}, Disabled: {len(s['d'])}, Failed: {len(s['f'])}."
) )
def _handle_exception(self, fut: asyncio.Task, task: Task) -> None: def _handle_exception(self, fut: asyncio.Task, task: Task) -> None:

View file

@ -55,12 +55,12 @@ def make_route_name(method: str, path: str) -> str:
return f"{method}:" + ".".join(segments or ["root"]) return f"{method}:" + ".".join(segments or ["root"])
def route(method: RouteType | str, path: str, name: str | None = None, **kwargs) -> Awaitable: def route(method: RouteType | str | list[str], path: str, name: str | None = None, **kwargs) -> Awaitable:
""" """
Decorator to mark a method as an HTTP route handler. Decorator to mark a method as an HTTP route handler.
Args: Args:
method (RouteType|str): The HTTP method. method (RouteType|str|list[str]): The HTTP method(s).
path (str): The path to the route. path (str): The path to the route.
name (str): The name of the route. name (str): The name of the route.
kwargs: Additional keyword arguments. kwargs: Additional keyword arguments.
@ -69,55 +69,64 @@ def route(method: RouteType | str, path: str, name: str | None = None, **kwargs)
Awaitable: The decorated function. Awaitable: The decorated function.
""" """
if not name: methods = [method] if isinstance(method, (str, RouteType)) else method
name = make_route_name(method, path)
def decorator(func): def decorator(func):
@wraps(func) @wraps(func)
async def wrapper(*args, **kwargs): async def wrapper(*args, **kwargs):
return await func(*args, **kwargs) return await func(*args, **kwargs)
route_type: str = RouteType.SOCKET if RouteType.SOCKET == method else RouteType.HTTP for m in methods:
if route_type not in ROUTES: route_name = name if name else make_route_name(m, path)
ROUTES[route_type] = {} route_type: str = RouteType.SOCKET if RouteType.SOCKET == m else RouteType.HTTP
ROUTES[route_type][name] = Route(method=method.upper(), path=path, name=name, handler=wrapper) if route_type not in ROUTES:
if "http" == route_type and path.endswith("/") and "/" != path and not kwargs.get("no_slash", False): ROUTES[route_type] = {}
ROUTES[route_type][f"{name}_no_slash"] = Route(
method=method.upper(), path=path[:-1], name=f"{name}_no_slash", handler=wrapper if route_name in ROUTES.get(route_type, {}):
) route_name = f"{route_name}_{len(ROUTES[route_type]) + 1}"
ROUTES[route_type][route_name] = Route(method=m.upper(), path=path, name=route_name, handler=wrapper)
if "http" == route_type and path.endswith("/") and "/" != path and not kwargs.get("no_slash", False):
ROUTES[route_type][f"{route_name}_no_slash"] = Route(
method=m.upper(), path=path[:-1], name=f"{route_name}_no_slash", handler=wrapper
)
return wrapper return wrapper
return decorator return decorator
def add_route(method: RouteType | str, path: str, handler: Awaitable, name: str | None = None, **kwargs): def add_route(method: RouteType | str | list[str], path: str, handler: Awaitable, name: str | None = None, **kwargs):
""" """
Decorator to mark a method as an HTTP route handler. Decorator to mark a method as an HTTP route handler.
Args: Args:
method (RouteType|str): The HTTP method. method (RouteType|str|list[str]): The HTTP method(s).
path (str): The path to the route. path (str): The path to the route.
name (str): The name of the route. name (str): The name of the route.
handler (Awaitable): The function that handles the route. handler (Awaitable): The function that handles the route.
kwargs: Additional keyword arguments. kwargs: Additional keyword arguments.
""" """
if not name: methods = [method] if isinstance(method, (str, RouteType)) else method
name = make_route_name(method, path)
route_type: str = RouteType.SOCKET if RouteType.SOCKET == method else RouteType.HTTP for m in methods:
route_name = name if name else make_route_name(m, path)
route_type: str = RouteType.SOCKET if RouteType.SOCKET == m else RouteType.HTTP
if route_type not in ROUTES: if route_type not in ROUTES:
ROUTES[route_type] = {} ROUTES[route_type] = {}
ROUTES[route_type][name] = Route(method=method.upper(), path=path, name=name, handler=handler) if route_name in ROUTES.get(route_type, {}):
route_name = f"{route_name}_{len(ROUTES[route_type]) + 1}"
if "http" == route_type and path.endswith("/") and "/" != path and not kwargs.get("no_slash", False): ROUTES[route_type][route_name] = Route(method=m.upper(), path=path, name=route_name, handler=handler)
ROUTES[route_type][f"{name}_no_slash"] = Route(
method=method.upper(), path=path[:-1], name=f"{name}_no_slash", handler=handler if "http" == route_type and path.endswith("/") and "/" != path and not kwargs.get("no_slash", False):
) ROUTES[route_type][f"{route_name}_no_slash"] = Route(
method=m.upper(), path=path[:-1], name=f"{route_name}_no_slash", handler=handler
)
def get_route(route_type: RouteType, name: str) -> dict[str, Route] | None: def get_route(route_type: RouteType, name: str) -> dict[str, Route] | None:

View file

@ -20,6 +20,7 @@ EXT_TO_MIME: dict = {
".json": "application/json", ".json": "application/json",
".ico": "image/x-icon", ".ico": "image/x-icon",
".webmanifest": "application/manifest+json", ".webmanifest": "application/manifest+json",
".m4a": "audio/mp4",
} }
FRONTEND_ROUTES: list[str] = [ FRONTEND_ROUTES: list[str] = [
@ -36,6 +37,7 @@ FRONTEND_ROUTES: list[str] = [
"/browser/{path:.*}", "/browser/{path:.*}",
] ]
async def serve_static_file(request: Request, config: Config) -> Response: async def serve_static_file(request: Request, config: Config) -> Response:
""" """
Preload static files from the ui/exported folder. Preload static files from the ui/exported folder.

View file

@ -0,0 +1,66 @@
import logging
from aiohttp import web
from aiohttp.web import Request, Response
from app.library.config import Config
from app.library.router import route
from app.library.Utils import get_file
LOG: logging.Logger = logging.getLogger(__name__)
@route(["GET", "HEAD"], "/api/download/{filename:.+}", "download_static")
async def download_file(request: Request, config: Config, app: web.Application) -> Response:
"""
Serve download files.
Args:
request (Request): The request object.
config (Config): The configuration instance.
app (web.Application): The aiohttp application instance.
Returns:
Response: The file response with correct MIME type.
"""
if not (filename := request.match_info.get("filename")):
return web.json_response({"error": "Filename required"}, status=web.HTTPBadRequest.status_code)
try:
realFile, status = get_file(download_path=config.download_path, file=filename)
if web.HTTPFound.status_code == status:
return Response(
status=web.HTTPFound.status_code,
headers={
"Location": str(
app.router["download_static"].url_for(
file=str(realFile).replace(config.download_path, "").strip("/")
)
),
},
)
if web.HTTPNotFound.status_code == status:
return web.json_response(data={"error": "File not found"}, status=status)
except Exception as e:
LOG.exception("Error retrieving file '%s': %s", filename, str(e))
return web.json_response(
data={"error": "Internal server error."},
status=web.HTTPInternalServerError.status_code,
)
if not realFile.is_file():
return web.json_response({"error": "File not found"}, status=web.HTTPNotFound.status_code)
from app.routes.api._static import EXT_TO_MIME
content_type = EXT_TO_MIME.get(realFile.suffix)
if not content_type:
import mimetypes
content_type, _ = mimetypes.guess_type(str(realFile))
if not content_type:
content_type = "application/octet-stream"
return web.FileResponse(path=str(realFile), headers={"Content-Type": content_type})

View file

@ -1,6 +1,7 @@
import asyncio import asyncio
import copy import copy
import re import re
import shutil
import tempfile import tempfile
import uuid import uuid
from dataclasses import dataclass from dataclasses import dataclass
@ -526,6 +527,116 @@ class TestCalcDownloadPath:
with pytest.raises(Exception, match="must resolve inside the base download folder"): with pytest.raises(Exception, match="must resolve inside the base download folder"):
calc_download_path(str(self.base_path), folder, create_path=False) calc_download_path(str(self.base_path), folder, create_path=False)
def test_calc_download_path_partial_match_attack(self):
"""
Test specific prefix matching vulnerability.
Base: /tmp/test
Target: /tmp/test_suffix (starts with base, but is a sibling)
This tests that the relative_to() check catches what startswith() alone would miss.
"""
# Create a sibling directory that starts with the base path name
sibling_dir = Path(str(self.temp_dir) + "_suffix")
sibling_dir.mkdir(exist_ok=True)
# Try to access the sibling directory
folder = f"../{sibling_dir.name}"
with pytest.raises(Exception, match="must resolve inside the base download folder"):
calc_download_path(str(self.base_path), folder, create_path=False)
# Clean up
shutil.rmtree(sibling_dir, ignore_errors=True)
def test_calc_download_path_symlink_attack_outside(self):
"""Test that symlinks pointing outside base directory are blocked."""
# Create a symlink pointing outside the base directory
outside_dir = Path(self.temp_dir).parent / "outside_target"
outside_dir.mkdir(exist_ok=True)
symlink_path = self.base_path / "evil_symlink"
try:
symlink_path.symlink_to(outside_dir, target_is_directory=True)
# Try to use the symlink
with pytest.raises(Exception, match="must resolve inside the base download folder"):
calc_download_path(str(self.base_path), "evil_symlink", create_path=False)
finally:
# Clean up
if symlink_path.exists():
symlink_path.unlink()
shutil.rmtree(outside_dir, ignore_errors=True)
def test_calc_download_path_symlink_attack_with_traversal(self):
"""Test symlink combined with path traversal."""
# Create a directory outside base
outside_dir = Path(self.temp_dir).parent / "target_dir"
outside_dir.mkdir(exist_ok=True)
# Create a symlink inside base pointing outside
safe_dir = self.base_path / "safe"
safe_dir.mkdir(exist_ok=True)
symlink_path = safe_dir / "link_to_outside"
try:
symlink_path.symlink_to(outside_dir, target_is_directory=True)
# Try to traverse through the symlink
with pytest.raises(Exception, match="must resolve inside the base download folder"):
calc_download_path(str(self.base_path), "safe/link_to_outside", create_path=False)
finally:
# Clean up
if symlink_path.exists():
symlink_path.unlink()
shutil.rmtree(safe_dir, ignore_errors=True)
shutil.rmtree(outside_dir, ignore_errors=True)
def test_calc_download_path_symlink_safe_internal(self):
"""Test that symlinks pointing inside base directory are allowed."""
# Create target directory inside base
target_dir = self.base_path / "target"
target_dir.mkdir(exist_ok=True)
# Create a symlink inside base pointing to another location inside base
symlink_path = self.base_path / "link_to_target"
try:
symlink_path.symlink_to(target_dir, target_is_directory=True)
# This should succeed since symlink resolves inside base
result = calc_download_path(str(self.base_path), "link_to_target", create_path=False)
assert str(target_dir) == result, "Internal symlinks should be allowed"
finally:
# Clean up
if symlink_path.exists():
symlink_path.unlink()
shutil.rmtree(target_dir, ignore_errors=True)
def test_calc_download_path_extremely_long_path(self):
"""Test handling of extremely long paths."""
# Create a very long folder name (most filesystems limit to 255 chars per component)
long_component = "a" * 300
# This should raise an error during path creation or validation (OSError or ValueError)
with pytest.raises((OSError, ValueError)):
calc_download_path(str(self.base_path), long_component, create_path=True)
def test_calc_download_path_many_nested_levels(self):
"""Test deeply nested directory structure."""
# Create a path with many nested levels
deep_path = "/".join([f"level{i}" for i in range(100)])
# This should work but might be slow
result = calc_download_path(str(self.base_path), deep_path, create_path=False)
expected = str(self.base_path / deep_path)
assert result == expected, "Should handle deeply nested paths"
def test_calc_download_path_directory_with_spaces(self):
"""Test paths with multiple spaces and special spacing."""
folder = "folder with multiple spaces"
result = calc_download_path(str(self.base_path), folder, create_path=True)
expected_path = self.base_path / folder
assert result == str(expected_path), "Should handle spaces correctly"
assert expected_path.exists(), "Directory with spaces should be created"
class TestMergeDict: class TestMergeDict:
"""Test the merge_dict function.""" """Test the merge_dict function."""