Compare commits

...
5 Commits
4 changed files with 58 additions and 43 deletions
+5
View File
@@ -20,6 +20,11 @@ Experience Cista by visiting [Cista Demo](https://drop.zi.fi) for a test run and
We recommend using [UV](https://docs.astral.sh/uv/getting-started/installation/) to directly run Cista:
Try it out locally at http://localhost:8000 (serves the current directory):
```fish
uvx cista
```
Create an account: (otherwise the server is public for all)
```fish
uvx cista --user yourname --privileged
+18 -36
View File
@@ -15,7 +15,8 @@ import re
import httpx
import websockets
from sanic import Blueprint
from sanic import Blueprint, json
from sanic import raw as raw_response
from sanic.exceptions import Forbidden, SanicException, Unauthorized
from sanic.log import logger
@@ -220,8 +221,6 @@ async def proxy_auth_request(request):
if key.lower() not in resp_hop_by_hop
]
from sanic import raw as raw_response
return raw_response(
raw_content,
status=response.status_code,
@@ -231,35 +230,31 @@ async def proxy_auth_request(request):
except httpx.RequestError as e:
logger.error(f"Auth proxy request failed: {e}")
from sanic import json
return json(
{"detail": "Authentication service unavailable", "error": str(e)},
{"detail": "Authentication service unavailable"},
status=503,
)
async def proxy_auth_websocket(request, ws):
"""Proxy a WebSocket connection to the auth backend."""
path = request.path
query_string = request.query_string
ws_backend = PASKIA_BACKEND_URL.replace("http://", "ws://").replace(
"https://", "wss://"
)
url = f"{ws_backend}{path}"
if query_string:
url = f"{url}?{query_string}"
url = f"ws{PASKIA_BACKEND_URL.removeprefix('http')}{request.path}"
if request.query_string:
url = f"{url}?{request.query_string}"
additional_headers = {}
if "cookie" in request.headers:
additional_headers["cookie"] = request.headers["cookie"]
if "authorization" in request.headers:
additional_headers["authorization"] = request.headers["authorization"]
if "host" in request.headers:
additional_headers["host"] = request.headers["host"]
if "origin" in request.headers:
additional_headers["origin"] = request.headers["origin"]
if "user-agent" in request.headers:
additional_headers["user-agent"] = request.headers["user-agent"]
additional_headers["x-forwarded-for"] = request.ip
additional_headers["x-forwarded-for"] = request.client_ip.strip("[]")
additional_headers["x-forwarded-host"] = request.host
additional_headers["x-forwarded-proto"] = request.scheme
@@ -291,23 +286,20 @@ async def proxy_auth_websocket(request, ws):
logger.error(f"WebSocket proxy to {url} failed: {e}")
def _is_websocket_request(request) -> bool:
"""Check if the request is a WebSocket upgrade request."""
connection = request.headers.get("connection", "").lower()
upgrade = request.headers.get("upgrade", "").lower()
connection_tokens = [t.strip() for t in connection.split(",")]
return "upgrade" in connection_tokens and upgrade == "websocket"
# Blueprint for auth proxy routes (only registered when paskia_enabled())
bp = Blueprint("sso", url_prefix="/auth")
async def _handle_websocket_upgrade(request):
"""Handle WebSocket upgrade and proxy the connection."""
protocol = request.transport.get_protocol()
ws = await protocol.websocket_handshake(request, subprotocols=None)
@bp.websocket("/ws/<path:path>")
async def auth_websocket_proxy(request, ws, path=""):
"""Proxy WebSocket connections to the auth backend."""
await proxy_auth_websocket(request, ws)
# Blueprint for auth proxy routes (only registered when paskia_enabled())
bp = Blueprint("sso", url_prefix="/auth")
@bp.websocket("/ws/")
async def auth_websocket_proxy_root(request, ws):
"""Proxy root WebSocket connections to the auth backend."""
await proxy_auth_websocket(request, ws)
@bp.route(
@@ -315,20 +307,10 @@ bp = Blueprint("sso", url_prefix="/auth")
)
async def auth_proxy(request, path=""):
"""Proxy all auth requests to the auth backend."""
if _is_websocket_request(request):
await _handle_websocket_upgrade(request)
from sanic import empty
return empty()
return await proxy_auth_request(request)
@bp.route("/", methods=["GET", "POST", "PUT", "DELETE", "PATCH", "OPTIONS", "HEAD"])
async def auth_proxy_root(request):
"""Proxy root auth requests to the auth backend."""
if _is_websocket_request(request):
await _handle_websocket_upgrade(request)
from sanic import empty
return empty()
return await proxy_auth_request(request)
+32 -2
View File
@@ -17,6 +17,33 @@ from cista import config
from cista.fileio import fuid
from cista.protocol import FileEntry, Space, UpdDel, UpdIns, UpdKeep
# Platform-specific allocated size calculation
if sys.platform == "win32":
import ctypes
from ctypes import wintypes
kernel32 = ctypes.windll.kernel32
GetCompressedFileSizeW = kernel32.GetCompressedFileSizeW
GetCompressedFileSizeW.argtypes = [wintypes.LPCWSTR, ctypes.POINTER(wintypes.DWORD)]
GetCompressedFileSizeW.restype = wintypes.DWORD
INVALID_FILE_SIZE = 0xFFFFFFFF
def get_allocated_size(path: Path, st: stat_result) -> int:
"""Get actual disk allocation on Windows using GetCompressedFileSizeW."""
high = wintypes.DWORD()
low = GetCompressedFileSizeW(str(path), ctypes.byref(high))
if low == INVALID_FILE_SIZE and ctypes.get_last_error() != 0:
raise OSError(f"GetCompressedFileSizeW failed for {path}")
return (high.value << 32) + low
else:
def get_allocated_size(path: Path, st: stat_result) -> int:
"""Get actual disk allocation on Unix using st_blocks."""
# st_blocks is in 512-byte units
return st.st_blocks * 512
pubsub = {}
sortkey = natsort_keygen(alg=ns.LOCALE)
@@ -148,8 +175,11 @@ def walk(rel: PurePosixPath, stat: stat_result | None = None) -> list[FileEntry]
try:
st = stat or path.stat()
isfile = int(not S_ISDIR(st.st_mode))
# st_blocks is in 512-byte units
allocated = st.st_blocks * 512 if isfile else 0
try:
allocated = get_allocated_size(path, st) if isfile else 0
except Exception:
logger.exception(f"get_allocated_size failed for {path}")
allocated = st.st_size if isfile else 0
entry = FileEntry(
level=len(rel.parts),
name=rel.name,
+3 -5
View File
@@ -114,6 +114,7 @@ filterwarnings = [
]
[tool.ruff.lint]
extend-select = ["E402"]
isort.known-first-party = ["cista"]
per-file-ignores."tests/*" = ["S", "ANN", "D", "INP", "PLR2004"]
per-file-ignores."scripts/*" = ["T20"]
@@ -121,16 +122,13 @@ per-file-ignores."scripts/*" = ["T20"]
[dependency-groups]
dev = [
"pytest>=8.4.1",
"pytest-asyncio>=0.25.0",
"pytest-cov>=7.0.0",
"ruff>=0.8.0",
"mypy>=1.13.0",
"pre-commit>=4.0.0",
"httpx>=0.28.1",
]
test = [
"pytest>=8.4.1",
"pytest-cov>=6.0.0",
"pytest-asyncio>=0.25.0",
]
[tool.coverage.run]
source = ["cista"]