Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 7 additions & 11 deletions asgidav/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -225,24 +225,20 @@ async def get(request: Request, path: str):
@app.put("/{path:path}")
async def put(request: Request, path: str):
try:
raw_length = request.headers.get("Content-Length", "").strip()
if "chunked" in request.headers.get("Transfer-Encoding", "").lower():
size = -1
elif raw_length:
size = int(raw_length)
if size < 0:
size = -1
else:
size = -1
raw_length = request.headers.get("Content-Length", "0").strip()
size = int(raw_length) if raw_length else 0
if size < 0:
size = 0
except (ValueError, TypeError):
logger.warning(f"PUT {path}: invalid Content-Length '{request.headers.get('Content-Length', '')}'")
size = -1
size = 0

try:
if not (member := await get_member(path)):
member = await (await root()).create_empty_resource(path)
if isinstance(member, Resource):
await member.overwrite(request.stream(), size=size)
if size > 0:
await member.overwrite(request.stream(), size=size)
return CREATED
return CONFLICT("Cannot PUT to a directory")
except TechnicalError as ex:
Expand Down
12 changes: 1 addition & 11 deletions dcfs/app/sftp/handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -401,17 +401,7 @@ async def read(self, offset: int, size: int) -> bytes:
if "r" not in self.mode:
raise asyncssh.SFTPPermissionDenied("File not open for reading")

# Lock-free fast path: if requested range is already in buffer, return immediately
rel_offset = offset - self._buf_offset
if rel_offset >= 0 and (rel_offset + size) <= len(self._read_buf):
return bytes(self._read_buf[rel_offset : rel_offset + size])

async with self._read_lock:
# Re-check fast path after acquiring lock
rel_offset = offset - self._buf_offset
if rel_offset >= 0 and (rel_offset + size) <= len(self._read_buf):
return bytes(self._read_buf[rel_offset : rel_offset + size])

buf_end = self._buf_offset + len(self._read_buf)

can_reuse_stream = (
Expand Down Expand Up @@ -465,7 +455,7 @@ async def read(self, offset: int, size: int) -> bytes:
if prune_target > self._buf_offset:
discard = min(prune_target - self._buf_offset, len(self._read_buf))
if discard > 0:
del self._read_buf[:discard]
self._read_buf = self._read_buf[discard:]
self._buf_offset += discard

return data
Expand Down
18 changes: 2 additions & 16 deletions dcfs/crypto/repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -252,25 +252,11 @@ async def content_length(self, fv: "DCFSFileVersion") -> int:
return 0
# ``_detect`` already caches per file, so the second call from a HEAD
# request or a Content-Range computation is free.
try:
detected = await self._detect(fv, "")
except (InvalidHeaderError, ValueError):
logger.warning(
"Header detection failed for fv.id=%s, falling back to fv.size",
fv.id,
)
return fv.size
detected = await self._detect(fv, "")
if detected is None:
return fv.size
header, _ = detected
try:
return _plaintext_size_from_ciphertext(fv.size, header.chunk_size)
except ValueError:
logger.warning(
"Plaintext size computation failed for fv.id=%s, falling back to fv.size",
fv.id,
)
return fv.size
return _plaintext_size_from_ciphertext(fv.size, header.chunk_size)

# -- internals ---------------------------------------------------------

Expand Down
84 changes: 13 additions & 71 deletions dcfs/discord/impl/discord_bot.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,56 +41,6 @@ def __init__(self, bot: discord.Client, bot_token: str):
self._bot = bot
self._bot_token = bot_token
self._http_session: Optional[aiohttp.ClientSession] = None
self._url_cache: dict[int, tuple[str, int]] = {}
self._inflight_fetches: dict[int, asyncio.Future[tuple[str, int]]] = {}
self._url_cache_lock = asyncio.Lock()

def _cache_url(self, message_id: int, url: str, size: int) -> None:
if len(self._url_cache) >= 10000:
first_key = next(iter(self._url_cache))
del self._url_cache[first_key]
self._url_cache[message_id] = (url, size)

async def _fetch_attachment_url_and_size(
self, channel_id: int, message_id: int, force_refresh: bool = False
) -> tuple[str, int]:
if not force_refresh and message_id in self._url_cache:
return self._url_cache[message_id]

async with self._url_cache_lock:
if not force_refresh and message_id in self._url_cache:
return self._url_cache[message_id]
if message_id in self._inflight_fetches:
fut = self._inflight_fetches[message_id]
else:
loop = asyncio.get_running_loop()
fut = loop.create_future()
self._inflight_fetches[message_id] = fut

async def _do_fetch():
try:
channel = await self._get_channel(channel_id)
try:
msg = await channel.fetch_message(message_id)
except discord.NotFound:
raise MessageNotFound(message_id)
if not msg.attachments:
raise UnDownloadableMessage(message_id)
att = msg.attachments[0]
res = (att.url, att.size)
self._cache_url(message_id, att.url, att.size)
if not fut.done():
fut.set_result(res)
except Exception as ex:
if not fut.done():
fut.set_exception(ex)
finally:
async with self._url_cache_lock:
self._inflight_fetches.pop(message_id, None)

asyncio.create_task(_do_fetch())

return await fut

async def _ensure_http_session(self) -> aiohttp.ClientSession:
if self._http_session is None or self._http_session.closed:
Expand Down Expand Up @@ -203,21 +153,29 @@ async def edit_message_media(self, req: EditMessageMediaReq) -> Message:

async def download_file(self, req: DownloadFileReq) -> DownloadFileResp:
channel_id = self._parse_channel_id(req.chat)
url, size = await self._fetch_attachment_url_and_size(channel_id, req.message_id)
channel = await self._get_channel(channel_id)
try:
msg = await channel.fetch_message(req.message_id)
except discord.NotFound:
raise MessageNotFound(req.message_id)
if not msg.attachments:
raise UnDownloadableMessage(req.message_id)
attachment = msg.attachments[0]

session = await self._ensure_http_session()

# Build optional Range header so the CDN only streams the requested
# byte range (critical for download_file_parallel sub-requests).
should_range = req.begin > 0 or req.end != -1
url = attachment.url
headers = {}
if should_range:
range_end = "" if req.end == -1 else str(req.end)
headers["Range"] = f"bytes={req.begin}-{range_end}"

logger.info(
"CDN download: msg=%d range=%d-%d should_range=%s attach_size=%d",
req.message_id, req.begin, req.end, should_range, size,
req.message_id, req.begin, req.end, should_range, attachment.size,
)

# Timeout: connect within 15s, download within 120s. Without a
Expand All @@ -229,24 +187,9 @@ async def download_file(self, req: DownloadFileReq) -> DownloadFileResp:
total=120.0,
)
t0 = asyncio.get_event_loop().time()
try:
response = await session.get(url, headers=headers, timeout=timeout)
response.raise_for_status()
except aiohttp.ClientResponseError as exc:
if exc.status in (401, 403, 404):
logger.warning(
"Cached CDN URL for msg=%d failed with status %d, refreshing URL...",
req.message_id, exc.status,
)
self._url_cache.pop(req.message_id, None)
url, size = await self._fetch_attachment_url_and_size(
channel_id, req.message_id, force_refresh=True
)
response = await session.get(url, headers=headers, timeout=timeout)
response.raise_for_status()
else:
raise
response = await session.get(url, headers=headers, timeout=timeout)
t1 = asyncio.get_event_loop().time()
response.raise_for_status()

# Determine whether the CDN honoured the Range header.
# 206 Partial Content means it did; 200 OK means it ignored it.
Expand Down Expand Up @@ -301,7 +244,7 @@ async def _chunk_generator():
finally:
response.close()

return DownloadFileResp(chunks=_chunk_generator(), size=size)
return DownloadFileResp(chunks=_chunk_generator(), size=attachment.size)

async def search_messages(self, req: SearchMessageReq) -> GetMessagesRespNoNone:
channel_id = self._parse_channel_id(req.chat)
Expand All @@ -323,7 +266,6 @@ def _to_message_dto(self, message: discord.Message) -> MessageResp:
size=att.size,
mime_type=att.content_type
)
self._cache_url(message.id, att.url, att.size)
return MessageResp(
message_id=message.id,
text=message.content if message.content else "",
Expand Down
32 changes: 0 additions & 32 deletions tests/test_asgidav/test_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,35 +74,3 @@ async def test_proppatch_endpoint_not_found(self, mocker):

response = client.request("PROPPATCH", "/nonexistent.txt")
assert response.status_code == 404

@pytest.mark.asyncio
async def test_put_endpoint_calls_overwrite_for_chunked_or_missing_length(self, mocker):
from fastapi.testclient import TestClient

from asgidav.app import create_app

from .common import MockResource

mock_res = MockResource("/test.txt")
mock_overwrite = mocker.patch.object(
mock_res, "overwrite", new_callable=mocker.AsyncMock
)

mock_get_member = mocker.AsyncMock(return_value=mock_res)
app = create_app(get_member=mock_get_member)
client = TestClient(app)

# 1. PUT with Transfer-Encoding: chunked
res1 = client.put(
"/test.txt", content=b"hello", headers={"Transfer-Encoding": "chunked"}
)
assert res1.status_code == 201
assert mock_overwrite.called
assert mock_overwrite.call_args[1]["size"] == -1

# 2. PUT with fixed Content-Length
mock_overwrite.reset_mock()
res2 = client.put("/test.txt", content=b"hello")
assert res2.status_code == 201
assert mock_overwrite.called
assert mock_overwrite.call_args[1]["size"] == 5
22 changes: 0 additions & 22 deletions tests/test_crypto/test_repository.py
Original file line number Diff line number Diff line change
Expand Up @@ -425,25 +425,3 @@ async def gen():

out = await _collect(await repo.get(fv, 0, -1, "stream.bin"))
assert out == plaintext


async def test_content_length_fallback_on_invalid_header() -> None:
"""When header verification fails during content_length, it should return fv.size as fallback."""
import struct as _struct

repo = _make_repo(chunk_size=4096)

# 1. Valid encrypted file returns plaintext size
plaintext = os.urandom(8192)
valid_fv = await _save_and_get_fv(repo, plaintext)
assert await repo.content_length(valid_fv) == len(plaintext)

# 2. Corrupted header MAC returns fv.size fallback
body = _struct.pack(">4sHHI32s", b"DCFS", 1, 1, 4096, b"\x00" * 32)
fake = body + b"\xff" * 16 # wrong MAC
corrupt_fv = _seed_plaintext(repo, fake + b"some more data")
assert await repo.content_length(corrupt_fv) == corrupt_fv.size

# 3. Short header with DCFS magic returns fv.size fallback
short_fv = _seed_plaintext(repo, b"DCFS" + b"\x00" * 10)
assert await repo.content_length(short_fv) == short_fv.size
Loading