fix(v2): enforce repository protocol safety
This commit is contained in:
@@ -26,7 +26,12 @@ from backup_tool.db.models import (
|
|||||||
Session,
|
Session,
|
||||||
User,
|
User,
|
||||||
)
|
)
|
||||||
from backup_tool.repository import RepositoryError, initialize, inspect_repository
|
from backup_tool.repository import (
|
||||||
|
RepositoryError,
|
||||||
|
initialize,
|
||||||
|
inspect_repository,
|
||||||
|
remove_repository,
|
||||||
|
)
|
||||||
from backup_tool.security.auth import (
|
from backup_tool.security.auth import (
|
||||||
hash_password,
|
hash_password,
|
||||||
hash_token,
|
hash_token,
|
||||||
@@ -529,10 +534,17 @@ def create_app(settings: Settings) -> FastAPI:
|
|||||||
compression=initialized.compression,
|
compression=initialized.compression,
|
||||||
encryption=initialized.encryption,
|
encryption=initialized.encryption,
|
||||||
)
|
)
|
||||||
db.add(repository)
|
try:
|
||||||
await db.flush()
|
db.add(repository)
|
||||||
await audit(db, request, "create", "repository", repository.id, "success", user.id)
|
await db.flush()
|
||||||
await db.commit()
|
await audit(db, request, "create", "repository", repository.id, "success", user.id)
|
||||||
|
await db.commit()
|
||||||
|
except Exception as error:
|
||||||
|
await db.rollback()
|
||||||
|
remove_repository(initialized.root)
|
||||||
|
raise Problem(
|
||||||
|
409, "repository_create_failed", "Repository metadata could not be stored."
|
||||||
|
) from error
|
||||||
return {
|
return {
|
||||||
"id": repository.id,
|
"id": repository.id,
|
||||||
"name": repository.name,
|
"name": repository.name,
|
||||||
@@ -589,7 +601,7 @@ def create_app(settings: Settings) -> FastAPI:
|
|||||||
if repository is None:
|
if repository is None:
|
||||||
raise Problem(404, "resource_not_found", "Repository was not found.")
|
raise Problem(404, "resource_not_found", "Repository was not found.")
|
||||||
try:
|
try:
|
||||||
inspected = inspect_repository(Path(repository.root))
|
inspected = inspect_repository(settings, Path(repository.root))
|
||||||
except RepositoryError as error:
|
except RepositoryError as error:
|
||||||
raise Problem(409, "repository_invalid", str(error)) from error
|
raise Problem(409, "repository_invalid", str(error)) from error
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -5,9 +5,11 @@ import json
|
|||||||
import os
|
import os
|
||||||
import shutil
|
import shutil
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
|
from datetime import UTC, datetime
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
from backup_tool.config import Settings
|
from backup_tool.config import Settings
|
||||||
|
from backup_tool.ids import new_uuid7
|
||||||
|
|
||||||
|
|
||||||
class RepositoryError(ValueError):
|
class RepositoryError(ValueError):
|
||||||
@@ -38,7 +40,14 @@ def _contained(root: Path, relative_path: str) -> Path:
|
|||||||
def _canonical_payload(compression: str, encryption: str) -> dict[str, object]:
|
def _canonical_payload(compression: str, encryption: str) -> dict[str, object]:
|
||||||
if compression != "none" or encryption != "none":
|
if compression != "none" or encryption != "none":
|
||||||
raise RepositoryError("requested repository policy is unavailable")
|
raise RepositoryError("requested repository policy is unavailable")
|
||||||
return {"compression": compression, "encryption": encryption, "format_version": 1}
|
return {
|
||||||
|
"repository_id": str(new_uuid7()),
|
||||||
|
"format_version": 1,
|
||||||
|
"digest_algorithm": "sha256",
|
||||||
|
"compression": compression,
|
||||||
|
"encryption": {"mode": "none", "key_id": None},
|
||||||
|
"created_at": datetime.now(UTC).isoformat(timespec="microseconds").replace("+00:00", "Z"),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def _canonical_json(payload: dict[str, object]) -> str:
|
def _canonical_json(payload: dict[str, object]) -> str:
|
||||||
@@ -76,8 +85,24 @@ def initialize(
|
|||||||
return InitializedRepository(root=root, compression=compression, encryption=encryption)
|
return InitializedRepository(root=root, compression=compression, encryption=encryption)
|
||||||
|
|
||||||
|
|
||||||
def inspect_repository(root: Path) -> InitializedRepository:
|
def remove_repository(root: Path) -> None:
|
||||||
metadata = root / "repository.json"
|
try:
|
||||||
|
shutil.rmtree(root)
|
||||||
|
except FileNotFoundError:
|
||||||
|
return
|
||||||
|
|
||||||
|
|
||||||
|
def inspect_repository(settings: Settings, root: Path) -> InitializedRepository:
|
||||||
|
try:
|
||||||
|
resolved_root = root.resolve(strict=True)
|
||||||
|
except OSError as error:
|
||||||
|
raise RepositoryError("repository root is unreadable") from error
|
||||||
|
if root.is_symlink() or not any(
|
||||||
|
resolved_root != allowed.resolve() and allowed.resolve() in resolved_root.parents
|
||||||
|
for allowed in settings.repository_roots
|
||||||
|
):
|
||||||
|
raise RepositoryError("repository root escapes configured roots")
|
||||||
|
metadata = resolved_root / "repository.json"
|
||||||
try:
|
try:
|
||||||
raw = metadata.read_text(encoding="utf-8")
|
raw = metadata.read_text(encoding="utf-8")
|
||||||
payload = json.loads(raw)
|
payload = json.loads(raw)
|
||||||
@@ -85,12 +110,31 @@ def inspect_repository(root: Path) -> InitializedRepository:
|
|||||||
raise RepositoryError("repository metadata is unreadable") from error
|
raise RepositoryError("repository metadata is unreadable") from error
|
||||||
if not isinstance(payload, dict):
|
if not isinstance(payload, dict):
|
||||||
raise RepositoryError("repository metadata is invalid")
|
raise RepositoryError("repository metadata is invalid")
|
||||||
try:
|
required = {
|
||||||
expected = _canonical_json(
|
"repository_id",
|
||||||
_canonical_payload(payload["compression"], payload["encryption"])
|
"format_version",
|
||||||
)
|
"digest_algorithm",
|
||||||
except (KeyError, TypeError, RepositoryError) as error:
|
"compression",
|
||||||
raise RepositoryError("repository metadata is invalid") from error
|
"encryption",
|
||||||
if payload.get("format_version") != 1 or raw != expected:
|
"created_at",
|
||||||
|
}
|
||||||
|
if (
|
||||||
|
set(payload) != required
|
||||||
|
or payload.get("format_version") != 1
|
||||||
|
or payload.get("digest_algorithm") != "sha256"
|
||||||
|
):
|
||||||
raise RepositoryError("repository metadata is not canonical")
|
raise RepositoryError("repository metadata is not canonical")
|
||||||
return InitializedRepository(root=root, compression="none", encryption="none")
|
encryption = payload.get("encryption")
|
||||||
|
if payload.get("compression") != "none" or encryption != {"mode": "none", "key_id": None}:
|
||||||
|
raise RepositoryError("repository metadata is invalid")
|
||||||
|
try:
|
||||||
|
from uuid import UUID
|
||||||
|
|
||||||
|
UUID(str(payload["repository_id"]))
|
||||||
|
created_at = str(payload["created_at"])
|
||||||
|
if not created_at.endswith("Z"):
|
||||||
|
raise ValueError
|
||||||
|
datetime.fromisoformat(created_at[:-1] + "+00:00")
|
||||||
|
except (ValueError, TypeError) as error:
|
||||||
|
raise RepositoryError("repository metadata is invalid") from error
|
||||||
|
return InitializedRepository(root=resolved_root, compression="none", encryption="none")
|
||||||
|
|||||||
@@ -45,7 +45,7 @@ def test_repository_inspection_rejects_noncanonical_metadata(tmp_path: Path) ->
|
|||||||
'{"encryption":"none","compression":"none","format_version":1}\n'
|
'{"encryption":"none","compression":"none","format_version":1}\n'
|
||||||
)
|
)
|
||||||
with pytest.raises(RepositoryError, match="canonical"):
|
with pytest.raises(RepositoryError, match="canonical"):
|
||||||
inspect_repository(result.root)
|
inspect_repository(settings, result.root)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
|
|||||||
@@ -148,7 +148,8 @@ async def test_logout_revocation_survives_app_restart(tmp_path) -> None:
|
|||||||
transport=ASGITransport(restarted), base_url=settings.public_base_url
|
transport=ASGITransport(restarted), base_url=settings.public_base_url
|
||||||
) as client:
|
) as client:
|
||||||
rejected = await client.get(
|
rejected = await client.get(
|
||||||
"/api/v2/auth/session", headers={"Cookie": f"backup_tool_session={copied_cookie}"}
|
"/api/v2/auth/session",
|
||||||
|
headers={"Cookie": f"backup_tool_session={copied_cookie}"},
|
||||||
)
|
)
|
||||||
assert rejected.status_code == 401
|
assert rejected.status_code == 401
|
||||||
await restarted.state.engine.dispose()
|
await restarted.state.engine.dispose()
|
||||||
|
|||||||
Reference in New Issue
Block a user