Add configurable max_file_size to reject oversized files before ingestion

Large files buffered entirely in RAM can OOM workers. Add
max_file_size to source config (default None = no limit).

FS checks stat().st_size before read_bytes(). HTTP and WebDAV issue
a HEAD request before GET when a limit is configured. S3 checks the
size from the existing head_async() call before get_async().

FileTooLargeError is classified as PermanentError so oversized files
go straight to the DLQ instead of retrying.
This commit is contained in:
Chris McDonough 2026-06-01 18:40:46 +03:00 committed by Yiorgis Gozadinos
parent 249d6c85c5
commit 7e4aa8c71f
No known key found for this signature in database
16 changed files with 262 additions and 28 deletions

View file

@ -155,6 +155,24 @@ Bearer-token auth can replace HTTP Basic via the standard `headers` map:
Authorization: Bearer ${KDRIVE_TOKEN} Authorization: Bearer ${KDRIVE_TOKEN}
``` ```
### File size limits
Any source can set `max_file_size` (bytes) to reject oversized files
before they are read into memory. Files exceeding the limit go
straight to the DLQ without retrying.
```yaml
- type: fs
root: /data/docs
max_file_size: 104857600 # 100 MB
```
FS and S3 sources know the size before downloading (`stat`, object
metadata), so the limit is always enforced. For HTTP and WebDAV the check
relies on a `Content-Length` response header; a server that omits it (for
example a chunked response) is fetched in full and the limit does not
apply.
## Workers and retry ## Workers and retry
```yaml ```yaml

View file

@ -356,6 +356,11 @@ class _SourceBase(BaseModel):
"this source. None = inherit from WorkerConfig.retry.", "this source. None = inherit from WorkerConfig.retry.",
) )
circuit_breaker: CircuitBreakerConfig = Field(default_factory=CircuitBreakerConfig) circuit_breaker: CircuitBreakerConfig = Field(default_factory=CircuitBreakerConfig)
max_file_size: int | None = Field(
default=None,
description="Maximum file size in bytes to fetch. Files larger than "
"this are rejected with a PermanentError. None = no limit.",
)
class FSSourceConfig(_SourceBase): class FSSourceConfig(_SourceBase):

View file

@ -32,9 +32,15 @@ def build_source(
include_patterns=cfg.include_patterns or None, include_patterns=cfg.include_patterns or None,
supported_extensions=supported_extensions, supported_extensions=supported_extensions,
source_id=cfg.id, source_id=cfg.id,
max_file_size=cfg.max_file_size,
) )
if isinstance(cfg, HTTPSourceConfig): if isinstance(cfg, HTTPSourceConfig):
return HTTPSource(source_id=cfg.id, urls=cfg.urls, headers=cfg.headers) return HTTPSource(
source_id=cfg.id,
urls=cfg.urls,
headers=cfg.headers,
max_file_size=cfg.max_file_size,
)
if isinstance(cfg, S3SourceConfig): if isinstance(cfg, S3SourceConfig):
return S3Source( return S3Source(
uri=cfg.uri, uri=cfg.uri,
@ -43,6 +49,7 @@ def build_source(
include_patterns=cfg.include_patterns or None, include_patterns=cfg.include_patterns or None,
supported_extensions=supported_extensions, supported_extensions=supported_extensions,
source_id=cfg.id, source_id=cfg.id,
max_file_size=cfg.max_file_size,
) )
if isinstance(cfg, WebDAVSourceConfig): if isinstance(cfg, WebDAVSourceConfig):
return WebDAVSource( return WebDAVSource(
@ -54,6 +61,7 @@ def build_source(
ignore_patterns=cfg.ignore_patterns or None, ignore_patterns=cfg.ignore_patterns or None,
include_patterns=cfg.include_patterns or None, include_patterns=cfg.include_patterns or None,
supported_extensions=supported_extensions, supported_extensions=supported_extensions,
max_file_size=cfg.max_file_size,
) )
raise TypeError( # pragma: no cover - discriminator union exhausts all cases raise TypeError( # pragma: no cover - discriminator union exhausts all cases
f"Unsupported source config: {type(cfg).__name__}" f"Unsupported source config: {type(cfg).__name__}"

View file

@ -47,6 +47,18 @@ class FetchResult(BaseModel):
disk_path: Path | None = None disk_path: Path | None = None
class FileTooLargeError(Exception):
"""Raised when a file exceeds the configured max_file_size."""
def check_file_size(size: int, max_file_size: int | None, uri: str) -> None:
"""Raise FileTooLargeError if size exceeds the limit."""
if max_file_size is not None and size > max_file_size:
raise FileTooLargeError(
f"{uri}: file size {size} bytes exceeds limit of {max_file_size} bytes"
)
@runtime_checkable @runtime_checkable
class Source(Protocol): class Source(Protocol):
source_id: str source_id: str

View file

@ -12,6 +12,7 @@ from haiku.rag.ingester.sources.base import (
RevisionSnapshot, RevisionSnapshot,
SourceEvent, SourceEvent,
SourceEventKind, SourceEventKind,
check_file_size,
) )
from haiku.rag.ingester.sources.filter import ( from haiku.rag.ingester.sources.filter import (
FileFilter, FileFilter,
@ -36,6 +37,7 @@ class FSSource:
include_patterns: list[str] | None = None, include_patterns: list[str] | None = None,
supported_extensions: list[str] | None = None, supported_extensions: list[str] | None = None,
source_id: str | None = None, source_id: str | None = None,
max_file_size: int | None = None,
) -> None: ) -> None:
# Resolve so symlinks and relative paths collapse to one canonical # Resolve so symlinks and relative paths collapse to one canonical
# root. The queue uses source_id as a foreign key — two paths for # root. The queue uses source_id as a foreign key — two paths for
@ -52,6 +54,7 @@ class FSSource:
include_patterns=include_patterns, include_patterns=include_patterns,
supported_extensions=self.supported_extensions, supported_extensions=self.supported_extensions,
) )
self._max_file_size = max_file_size
def _resolve_within_root(self, uri: str) -> Path | None: def _resolve_within_root(self, uri: str) -> Path | None:
"""Resolve a URI to a real path guaranteed to live under ``self.root``. """Resolve a URI to a real path guaranteed to live under ``self.root``.
@ -88,6 +91,7 @@ class FSSource:
path = self._resolve_within_root(uri) path = self._resolve_within_root(uri)
if path is None: if path is None:
raise UnsupportedSourceError(f"Path escapes FS root ({self.root}): {uri}") raise UnsupportedSourceError(f"Path escapes FS root ({self.root}): {uri}")
check_file_size(path.stat().st_size, self._max_file_size, uri)
body = path.read_bytes() body = path.read_bytes()
content_type, _ = mimetypes.guess_type(path.name) content_type, _ = mimetypes.guess_type(path.name)
if content_type is None: if content_type is None:

View file

@ -11,6 +11,7 @@ from haiku.rag.ingester.sources.base import (
RevisionSnapshot, RevisionSnapshot,
SourceEvent, SourceEvent,
SourceEventKind, SourceEventKind,
check_file_size,
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@ -38,11 +39,13 @@ class HTTPSource:
urls: list[str] | None = None, urls: list[str] | None = None,
headers: dict[str, str] | None = None, headers: dict[str, str] | None = None,
transport: httpx.AsyncBaseTransport | None = None, transport: httpx.AsyncBaseTransport | None = None,
max_file_size: int | None = None,
) -> None: ) -> None:
self.source_id = source_id self.source_id = source_id
self.urls = list(urls or []) self.urls = list(urls or [])
self.headers = dict(headers or {}) self.headers = dict(headers or {})
self._http = httpx.AsyncClient(headers=self.headers, transport=transport) self._http = httpx.AsyncClient(headers=self.headers, transport=transport)
self._max_file_size = max_file_size
def supports(self, uri: str) -> bool: def supports(self, uri: str) -> bool:
return urlparse(uri).scheme in ("http", "https") return urlparse(uri).scheme in ("http", "https")
@ -62,6 +65,11 @@ class HTTPSource:
return revision return revision
async def fetch(self, uri: str) -> FetchResult: async def fetch(self, uri: str) -> FetchResult:
if self._max_file_size is not None:
head = await self._http.head(uri)
content_length = head.headers.get("content-length")
if content_length is not None:
check_file_size(int(content_length), self._max_file_size, uri)
response = await self._http.get(uri) response = await self._http.get(uri)
response.raise_for_status() response.raise_for_status()
body = response.content body = response.content

View file

@ -10,6 +10,7 @@ from haiku.rag.ingester.sources.base import (
RevisionSnapshot, RevisionSnapshot,
SourceEvent, SourceEvent,
SourceEventKind, SourceEventKind,
check_file_size,
) )
from haiku.rag.ingester.sources.filter import ( from haiku.rag.ingester.sources.filter import (
FileFilter, FileFilter,
@ -42,6 +43,7 @@ class S3Source:
include_patterns: list[str] | None = None, include_patterns: list[str] | None = None,
supported_extensions: list[str] | None = None, supported_extensions: list[str] | None = None,
source_id: str | None = None, source_id: str | None = None,
max_file_size: int | None = None,
) -> None: ) -> None:
self.bucket, self.prefix = _parse_s3_uri(uri) self.bucket, self.prefix = _parse_s3_uri(uri)
# uri_prefix is the canonical "everything I own" — used by supports() # uri_prefix is the canonical "everything I own" — used by supports()
@ -59,6 +61,7 @@ class S3Source:
include_patterns=include_patterns, include_patterns=include_patterns,
supported_extensions=self.supported_extensions, supported_extensions=self.supported_extensions,
) )
self._max_file_size = max_file_size
def supports(self, uri: str) -> bool: def supports(self, uri: str) -> bool:
return uri.startswith(self.uri_prefix) return uri.startswith(self.uri_prefix)
@ -86,6 +89,9 @@ class S3Source:
head = await obstore.head_async(store, key) head = await obstore.head_async(store, key)
etag = (head.get("e_tag") or "").strip('"').strip() or None etag = (head.get("e_tag") or "").strip('"').strip() or None
size = head.get("size") or head.get("content_length")
if size is not None:
check_file_size(int(size), self._max_file_size, uri)
resp = await obstore.get_async(store, key) resp = await obstore.get_async(store, key)
body = await resp.bytes_async() body = await resp.bytes_async()

View file

@ -12,6 +12,7 @@ from haiku.rag.ingester.sources.base import (
RevisionSnapshot, RevisionSnapshot,
SourceEvent, SourceEvent,
SourceEventKind, SourceEventKind,
check_file_size,
) )
from haiku.rag.ingester.sources.filter import ( from haiku.rag.ingester.sources.filter import (
FileFilter, FileFilter,
@ -170,6 +171,7 @@ class WebDAVSource:
include_patterns: list[str] | None = None, include_patterns: list[str] | None = None,
supported_extensions: list[str] | None = None, supported_extensions: list[str] | None = None,
transport: httpx.AsyncBaseTransport | None = None, transport: httpx.AsyncBaseTransport | None = None,
max_file_size: int | None = None,
) -> None: ) -> None:
self.source_id = source_id self.source_id = source_id
# Trailing slash matters: urljoin treats path-without-slash as a sibling # Trailing slash matters: urljoin treats path-without-slash as a sibling
@ -189,6 +191,7 @@ class WebDAVSource:
supported_extensions=self.supported_extensions, supported_extensions=self.supported_extensions,
) )
# transport is for testing — production callers leave it None. # transport is for testing — production callers leave it None.
self._max_file_size = max_file_size
auth = ( auth = (
(self.username, self.password) (self.username, self.password)
if self.username is not None and self.password is not None if self.username is not None and self.password is not None
@ -219,6 +222,11 @@ class WebDAVSource:
return entries[0].revision return entries[0].revision
async def fetch(self, uri: str) -> FetchResult: async def fetch(self, uri: str) -> FetchResult:
if self._max_file_size is not None:
head = await self._http.head(uri)
content_length = head.headers.get("content-length")
if content_length is not None:
check_file_size(int(content_length), self._max_file_size, uri)
response = await self._http.get(uri) response = await self._http.get(uri)
response.raise_for_status() response.raise_for_status()
body = response.content body = response.content

View file

@ -8,6 +8,7 @@ from pydantic import BaseModel
from haiku.rag.client.exceptions import UnsupportedSourceError from haiku.rag.client.exceptions import UnsupportedSourceError
from haiku.rag.ingester.exceptions import PermanentError, TransientError from haiku.rag.ingester.exceptions import PermanentError, TransientError
from haiku.rag.ingester.queue.models import Job, JobOp from haiku.rag.ingester.queue.models import Job, JobOp
from haiku.rag.ingester.sources.base import FileTooLargeError
from haiku.rag.telemetry import attach_context, logfire from haiku.rag.telemetry import attach_context, logfire
if TYPE_CHECKING: if TYPE_CHECKING:
@ -34,7 +35,7 @@ def _classify(exc: BaseException) -> Exception:
# UnsupportedSourceError is the typed signal from client/* that the # UnsupportedSourceError is the typed signal from client/* that the
# source will never ingest successfully on a retry (bad URI scheme, # source will never ingest successfully on a retry (bad URI scheme,
# missing file, unsupported extension, etc.). # missing file, unsupported extension, etc.).
if isinstance(exc, UnsupportedSourceError): if isinstance(exc, UnsupportedSourceError | FileTooLargeError):
return PermanentError(str(exc)) return PermanentError(str(exc))
if isinstance(exc, ValueError): if isinstance(exc, ValueError):

View file

@ -4,7 +4,7 @@ from pathlib import Path
import pytest import pytest
from haiku.rag.client.exceptions import UnsupportedSourceError from haiku.rag.client.exceptions import UnsupportedSourceError
from haiku.rag.ingester.sources.base import SourceEventKind from haiku.rag.ingester.sources.base import FileTooLargeError, SourceEventKind
from haiku.rag.ingester.sources.fs import FSSource from haiku.rag.ingester.sources.fs import FSSource
@ -279,3 +279,24 @@ async def test_fs_source_discover_skips_symlinked_directories(
uris = {e.uri async for e in src.discover(since=None)} uris = {e.uri async for e in src.discover(since=None)}
# No URI under /escape/* should appear. # No URI under /escape/* should appear.
assert not any("escape" in u for u in uris) assert not any("escape" in u for u in uris)
@pytest.mark.asyncio
async def test_fs_source_fetch_rejects_file_exceeding_max_size(fs_root: Path):
src = FSSource(root=fs_root, max_file_size=3)
with pytest.raises(FileTooLargeError):
await src.fetch((fs_root / "a.md").as_uri()) # "alpha" = 5 bytes
@pytest.mark.asyncio
async def test_fs_source_fetch_allows_file_within_max_size(fs_root: Path):
src = FSSource(root=fs_root, max_file_size=100)
result = await src.fetch((fs_root / "a.md").as_uri())
assert result.body == b"alpha"
@pytest.mark.asyncio
async def test_fs_source_fetch_no_limit_when_max_size_is_none(fs_root: Path):
src = FSSource(root=fs_root, max_file_size=None)
result = await src.fetch((fs_root / "a.md").as_uri())
assert result.body == b"alpha"

View file

@ -3,7 +3,7 @@ import hashlib
import httpx import httpx
import pytest import pytest
from haiku.rag.ingester.sources.base import SourceEventKind from haiku.rag.ingester.sources.base import FileTooLargeError, SourceEventKind
from haiku.rag.ingester.sources.http import HTTPSource from haiku.rag.ingester.sources.http import HTTPSource
@ -348,18 +348,11 @@ async def test_discover_propagates_non_transport_errors():
async def test_discover_emits_unchanged_for_known_url_without_revision(): async def test_discover_emits_unchanged_for_known_url_without_revision():
"""A server that returns no ETag or Last-Modified should not cause """A server that returns no ETag or Last-Modified should not cause
re-ingestion every sweep once the URL has been ingested.""" re-ingestion every sweep once the URL has been ingested."""
transport = _transport( transport = _transport({("HEAD", "https://example.com/a.md"): httpx.Response(200)})
{("HEAD", "https://example.com/a.md"): httpx.Response(200)}
)
src = HTTPSource( src = HTTPSource(
source_id="x", urls=["https://example.com/a.md"], transport=transport source_id="x", urls=["https://example.com/a.md"], transport=transport
) )
events = [ events = [e async for e in src.discover(known_uris={"https://example.com/a.md"})]
e
async for e in src.discover(
known_uris={"https://example.com/a.md"}
)
]
assert len(events) == 1 assert len(events) == 1
assert events[0].kind is SourceEventKind.UNCHANGED assert events[0].kind is SourceEventKind.UNCHANGED
assert events[0].revision is None assert events[0].revision is None
@ -377,3 +370,55 @@ async def test_discover_emits_upsert_for_unknown_url_without_revision():
events = [e async for e in src.discover()] events = [e async for e in src.discover()]
assert len(events) == 1 assert len(events) == 1
assert events[0].kind is SourceEventKind.UPSERT assert events[0].kind is SourceEventKind.UPSERT
@pytest.mark.asyncio
async def test_fetch_rejects_file_exceeding_max_size():
transport = _transport(
{
("HEAD", "https://example.com/big.bin"): httpx.Response(
200, headers={"content-length": "5000"}
),
}
)
src = HTTPSource(source_id="default", transport=transport, max_file_size=1000)
with pytest.raises(FileTooLargeError):
await src.fetch("https://example.com/big.bin")
@pytest.mark.asyncio
async def test_fetch_allows_file_within_max_size():
body = b"small"
transport = _transport(
{
("HEAD", "https://example.com/a.md"): httpx.Response(
200, headers={"content-length": str(len(body))}
),
("GET", "https://example.com/a.md"): httpx.Response(
200, content=body, headers={"content-type": "text/markdown"}
),
}
)
src = HTTPSource(source_id="default", transport=transport, max_file_size=1000)
result = await src.fetch("https://example.com/a.md")
assert result.body == body
@pytest.mark.asyncio
async def test_fetch_skips_head_when_no_max_size():
"""When max_file_size is None, no HEAD request is made."""
calls = []
def handler(request: httpx.Request) -> httpx.Response:
calls.append(request.method)
if request.method == "GET":
return httpx.Response(200, content=b"ok")
return httpx.Response(200)
src = HTTPSource(
source_id="default",
transport=httpx.MockTransport(handler),
max_file_size=None,
)
await src.fetch("https://example.com/a.md")
assert calls == ["GET"]

View file

@ -7,6 +7,7 @@ import pytest
from haiku.rag.client import HaikuRAG from haiku.rag.client import HaikuRAG
from haiku.rag.ingester.exceptions import PermanentError, TransientError from haiku.rag.ingester.exceptions import PermanentError, TransientError
from haiku.rag.ingester.queue.models import Job, JobOp, JobStatus from haiku.rag.ingester.queue.models import Job, JobOp, JobStatus
from haiku.rag.ingester.sources.base import FileTooLargeError
from haiku.rag.ingester.workers.pipeline import run_job from haiku.rag.ingester.workers.pipeline import run_job
from haiku.rag.store.models.document import Document from haiku.rag.store.models.document import Document
@ -289,3 +290,11 @@ async def test_directory_errors_classified_as_permanent(exc_class):
client.create_document_from_source.side_effect = exc_class("not a file") client.create_document_from_source.side_effect = exc_class("not a file")
with pytest.raises(PermanentError, match="path error"): with pytest.raises(PermanentError, match="path error"):
await run_job(client, _job()) await run_job(client, _job())
@pytest.mark.asyncio
async def test_file_too_large_classified_as_permanent():
client = _mock_client()
client.create_document_from_source.side_effect = FileTooLargeError("too big")
with pytest.raises(PermanentError, match="too big"):
await run_job(client, _job())

View file

@ -118,7 +118,9 @@ def _periodic(source, config, jobs, sync, **kwargs):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_stagger_start_sleeps_fraction_of_interval(jobs, sync, fs_config, monkeypatch): async def test_stagger_start_sleeps_fraction_of_interval(
jobs, sync, fs_config, monkeypatch
):
"""_stagger_start should sleep for a random fraction of poll_interval_s """_stagger_start should sleep for a random fraction of poll_interval_s
and return False (not stopped).""" and return False (not stopped)."""
monkeypatch.setattr("random.uniform", lambda a, b: b) # max jitter monkeypatch.setattr("random.uniform", lambda a, b: b) # max jitter
@ -130,7 +132,9 @@ async def test_stagger_start_sleeps_fraction_of_interval(jobs, sync, fs_config,
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_stagger_start_returns_true_when_stopped(jobs, sync, fs_config, monkeypatch): async def test_stagger_start_returns_true_when_stopped(
jobs, sync, fs_config, monkeypatch
):
"""If _stop is set before the jitter elapses, _stagger_start returns True.""" """If _stop is set before the jitter elapses, _stagger_start returns True."""
monkeypatch.setattr("random.uniform", lambda a, b: 10.0) # long jitter monkeypatch.setattr("random.uniform", lambda a, b: 10.0) # long jitter
source = _StubSource("src", []) source = _StubSource("src", [])
@ -511,7 +515,9 @@ async def test_watch_deleted_then_added_enqueues_upsert(tmp_path, jobs, sync):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_watch_added_file_deleted_before_stat_does_not_crash(tmp_path, jobs, sync): async def test_watch_added_file_deleted_before_stat_does_not_crash(
tmp_path, jobs, sync
):
"""If a file is deleted between the watchfiles event and the stat() """If a file is deleted between the watchfiles event and the stat()
call, the handler should return silently instead of raising call, the handler should return silently instead of raising
FileNotFoundError and killing the watch loop.""" FileNotFoundError and killing the watch loop."""

View file

@ -3,7 +3,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from haiku.rag.ingester.sources.base import SourceEventKind from haiku.rag.ingester.sources.base import FileTooLargeError, SourceEventKind
from haiku.rag.ingester.sources.s3 import S3Source from haiku.rag.ingester.sources.s3 import S3Source
@ -221,10 +221,7 @@ async def test_discover_emits_unchanged_for_known_key_without_etag(fake_s3_listi
once the key has been ingested.""" once the key has been ingested."""
fake_s3_listing([[{"path": "file.md", "size": 0, "last_modified": None}]]) fake_s3_listing([[{"path": "file.md", "size": 0, "last_modified": None}]])
src = S3Source(uri="s3://bucket/", supported_extensions=[".md"]) src = S3Source(uri="s3://bucket/", supported_extensions=[".md"])
events = [ events = [e async for e in src.discover(known_uris={"s3://bucket/file.md"})]
e
async for e in src.discover(known_uris={"s3://bucket/file.md"})
]
non_delete = [e for e in events if e.kind is not SourceEventKind.DELETE] non_delete = [e for e in events if e.kind is not SourceEventKind.DELETE]
assert len(non_delete) == 1 assert len(non_delete) == 1
assert non_delete[0].kind is SourceEventKind.UNCHANGED assert non_delete[0].kind is SourceEventKind.UNCHANGED
@ -238,3 +235,38 @@ async def test_discover_emits_upsert_for_unknown_key_without_etag(fake_s3_listin
events = [e async for e in src.discover()] events = [e async for e in src.discover()]
assert len(events) == 1 assert len(events) == 1
assert events[0].kind is SourceEventKind.UPSERT assert events[0].kind is SourceEventKind.UPSERT
@pytest.mark.asyncio
async def test_fetch_rejects_file_exceeding_max_size(fake_obstore_io):
head_async, get_async = fake_obstore_io
head_async.return_value = {"e_tag": '"abc"', "size": 5000}
src = S3Source(uri="s3://bucket/", max_file_size=1000)
with pytest.raises(FileTooLargeError):
await src.fetch("s3://bucket/file.txt")
get_async.assert_not_awaited()
@pytest.mark.asyncio
async def test_fetch_allows_file_within_max_size(fake_obstore_io):
head_async, get_async = fake_obstore_io
body = b"small"
head_async.return_value = {"e_tag": '"abc"', "size": len(body)}
get_async.return_value = _get_result(body)
src = S3Source(uri="s3://bucket/", max_file_size=1000)
result = await src.fetch("s3://bucket/file.txt")
assert result.body == body
@pytest.mark.asyncio
async def test_fetch_no_limit_when_max_size_is_none(fake_obstore_io):
head_async, get_async = fake_obstore_io
body = b"any size"
head_async.return_value = {"e_tag": '"abc"'}
get_async.return_value = _get_result(body)
src = S3Source(uri="s3://bucket/", max_file_size=None)
result = await src.fetch("s3://bucket/file.txt")
assert result.body == body

View file

@ -3,7 +3,7 @@ import hashlib
import httpx import httpx
import pytest import pytest
from haiku.rag.ingester.sources.base import SourceEventKind from haiku.rag.ingester.sources.base import FileTooLargeError, SourceEventKind
from haiku.rag.ingester.sources.webdav import WebDAVSource, _strip_etag from haiku.rag.ingester.sources.webdav import WebDAVSource, _strip_etag
@ -405,9 +405,7 @@ async def test_discover_emits_unchanged_for_known_uri_without_revision():
) )
events = [ events = [
e e
async for e in src.discover( async for e in src.discover(known_uris={"https://nc.example.com/dav/norev.md"})
known_uris={"https://nc.example.com/dav/norev.md"}
)
] ]
non_delete = [e for e in events if e.kind is not SourceEventKind.DELETE] non_delete = [e for e in events if e.kind is not SourceEventKind.DELETE]
assert len(non_delete) == 1 assert len(non_delete) == 1
@ -434,3 +432,58 @@ async def test_discover_emits_upsert_for_unknown_uri_without_revision():
non_delete = [e for e in events if e.kind is not SourceEventKind.DELETE] non_delete = [e for e in events if e.kind is not SourceEventKind.DELETE]
assert len(non_delete) == 1 assert len(non_delete) == 1
assert non_delete[0].kind is SourceEventKind.UPSERT assert non_delete[0].kind is SourceEventKind.UPSERT
@pytest.mark.asyncio
async def test_fetch_rejects_file_exceeding_max_size():
def handler(request: httpx.Request) -> httpx.Response:
if request.method == "HEAD":
return httpx.Response(200, headers={"content-length": "5000"})
return httpx.Response(200, content=b"big")
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
max_file_size=1000,
)
with pytest.raises(FileTooLargeError):
await src.fetch("https://nc.example.com/dav/big.bin")
@pytest.mark.asyncio
async def test_fetch_allows_file_within_max_size():
body = b"small"
def handler(request: httpx.Request) -> httpx.Response:
if request.method == "HEAD":
return httpx.Response(200, headers={"content-length": str(len(body))})
return httpx.Response(200, content=body, headers={"content-type": "text/plain"})
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
max_file_size=1000,
)
result = await src.fetch("https://nc.example.com/dav/a.txt")
assert result.body == body
@pytest.mark.asyncio
async def test_fetch_skips_head_when_no_max_size():
"""When max_file_size is None, no HEAD request is made."""
calls = []
def handler(request: httpx.Request) -> httpx.Response:
calls.append(request.method)
return httpx.Response(200, content=b"ok")
src = WebDAVSource(
source_id="nc",
base_url="https://nc.example.com/dav/",
transport=_transport(handler),
max_file_size=None,
)
await src.fetch("https://nc.example.com/dav/a.txt")
assert calls == ["GET"]

View file

@ -93,9 +93,7 @@ async def test_stop_completes_with_idle_workers(client, jobs, sync):
try: try:
await asyncio.wait_for(pool.stop(), timeout=2.0) await asyncio.wait_for(pool.stop(), timeout=2.0)
except TimeoutError: except TimeoutError:
pytest.fail( pytest.fail("stop() did not complete within 2s — idle workers were not woken")
"stop() did not complete within 2s — idle workers were not woken"
)
assert pool.live_workers == 0 assert pool.live_workers == 0