From 07c5a97929a79243f367ada682661f574cbf7bd2 Mon Sep 17 00:00:00 2001 From: Chris McDonough Date: Mon, 1 Jun 2026 07:00:15 -0400 Subject: [PATCH] Reuse httpx.AsyncClient across requests in HTTP and WebDAV sources MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit HTTPSource and WebDAVSource previously created a new AsyncClient for every head(), fetch(), and discover() call — no connection reuse, TLS renegotiation on every request, and connection pool churn at scale. Create the client once in __init__ and reuse it for the lifetime of the source. Add aclose() to both sources, called by PollerManager on shutdown to cleanly close the connection pool. --- .../haiku/rag/ingester/pollers/manager.py | 3 + .../haiku/rag/ingester/sources/http.py | 135 +++++++++--------- .../haiku/rag/ingester/sources/webdav.py | 59 ++++---- 3 files changed, 96 insertions(+), 101 deletions(-) diff --git a/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py b/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py index 6339a5c8..bd7638d5 100644 --- a/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py +++ b/haiku_rag_slim/haiku/rag/ingester/pollers/manager.py @@ -96,6 +96,9 @@ class PollerManager: await asyncio.gather(*self._tasks, return_exceptions=True) self._tasks.clear() self._started = False + for source in self.sources: + if hasattr(source, "aclose"): + await source.aclose() @property def pollers(self) -> list[BasePoller]: diff --git a/haiku_rag_slim/haiku/rag/ingester/sources/http.py b/haiku_rag_slim/haiku/rag/ingester/sources/http.py index 83c0bada..b6ff4afa 100644 --- a/haiku_rag_slim/haiku/rag/ingester/sources/http.py +++ b/haiku_rag_slim/haiku/rag/ingester/sources/http.py @@ -39,48 +39,44 @@ class HTTPSource: self.source_id = source_id self.urls = list(urls or []) self.headers = dict(headers or {}) - # transport is for testing — production callers leave it None and httpx - # uses its real transport. - self._transport = transport + self._http = httpx.AsyncClient(headers=self.headers, transport=transport) def supports(self, uri: str) -> bool: return urlparse(uri).scheme in ("http", "https") - def _client(self) -> httpx.AsyncClient: - return httpx.AsyncClient(headers=self.headers, transport=self._transport) + async def aclose(self) -> None: + await self._http.aclose() async def head(self, uri: str) -> str | None: """HEAD probe for the cheap revision short-circuit. Returns the ETag (or Last-Modified) so an unchanged remote URL can skip the full GET. None on HTTP error so the caller falls back to fetch(); network errors propagate and the worker's classifier handles them.""" - async with self._client() as http: - response = await http.head(uri) - if response.is_error: - return None + response = await self._http.head(uri) + if response.is_error: + return None revision, _ = _extract_revision(response.headers) return revision async def fetch(self, uri: str) -> FetchResult: - async with self._client() as http: - response = await http.get(uri) - response.raise_for_status() - body = response.content - content_type = ( - response.headers.get("content-type", "application/octet-stream") - .split(";")[0] - .strip() - .lower() - ) - revision, extra = _extract_revision(response.headers) - return FetchResult( - uri=uri, - body=body, - content_type=content_type, - content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(), - revision=revision, - extra_metadata=extra, - ) + response = await self._http.get(uri) + response.raise_for_status() + body = response.content + content_type = ( + response.headers.get("content-type", "application/octet-stream") + .split(";")[0] + .strip() + .lower() + ) + revision, extra = _extract_revision(response.headers) + return FetchResult( + uri=uri, + body=body, + content_type=content_type, + content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(), + revision=revision, + extra_metadata=extra, + ) async def discover( self, @@ -104,53 +100,52 @@ class HTTPSource: now = datetime.now(UTC) configured = set(self.urls) - async with self._client() as http: - for url in self.urls: - try: - head = await http.head(url) - except Exception: - yield SourceEvent( - source_id=self.source_id, - uri=url, - kind=SourceEventKind.UPSERT, - revision=None, - discovered_at=now, - ) - continue - - if head.status_code == 410: - yield SourceEvent( - source_id=self.source_id, - uri=url, - kind=SourceEventKind.DELETE, - revision=None, - discovered_at=now, - ) - continue - - if head.is_error: - yield SourceEvent( - source_id=self.source_id, - uri=url, - kind=SourceEventKind.UPSERT, - revision=None, - discovered_at=now, - ) - continue - - revision, _ = _extract_revision(head.headers) - if revision is not None and snapshot.get(url) == revision: - kind = SourceEventKind.UNCHANGED - else: - kind = SourceEventKind.UPSERT - + for url in self.urls: + try: + head = await self._http.head(url) + except Exception: yield SourceEvent( source_id=self.source_id, uri=url, - kind=kind, - revision=revision, + kind=SourceEventKind.UPSERT, + revision=None, discovered_at=now, ) + continue + + if head.status_code == 410: + yield SourceEvent( + source_id=self.source_id, + uri=url, + kind=SourceEventKind.DELETE, + revision=None, + discovered_at=now, + ) + continue + + if head.is_error: + yield SourceEvent( + source_id=self.source_id, + uri=url, + kind=SourceEventKind.UPSERT, + revision=None, + discovered_at=now, + ) + continue + + revision, _ = _extract_revision(head.headers) + if revision is not None and snapshot.get(url) == revision: + kind = SourceEventKind.UNCHANGED + else: + kind = SourceEventKind.UPSERT + + yield SourceEvent( + source_id=self.source_id, + uri=url, + kind=kind, + revision=revision, + discovered_at=now, + ) # Anything previously known to this source that's no longer in # config emits DELETE so delete_orphans can clean up. Without this, diff --git a/haiku_rag_slim/haiku/rag/ingester/sources/webdav.py b/haiku_rag_slim/haiku/rag/ingester/sources/webdav.py index 7dd6a15e..0bb512f8 100644 --- a/haiku_rag_slim/haiku/rag/ingester/sources/webdav.py +++ b/haiku_rag_slim/haiku/rag/ingester/sources/webdav.py @@ -189,41 +189,39 @@ class WebDAVSource: supported_extensions=self.supported_extensions, ) # transport is for testing — production callers leave it None. - self._transport = transport - - def supports(self, uri: str) -> bool: - return uri.startswith(self.base_url) - - def _client(self) -> httpx.AsyncClient: auth = ( (self.username, self.password) if self.username is not None and self.password is not None else None ) - return httpx.AsyncClient( - auth=auth, headers=self.headers, transport=self._transport + self._http = httpx.AsyncClient( + auth=auth, headers=self.headers, transport=transport ) + def supports(self, uri: str) -> bool: + return uri.startswith(self.base_url) + + async def aclose(self) -> None: + await self._http.aclose() + async def head(self, uri: str) -> str | None: - async with self._client() as http: - response = await http.request( - "PROPFIND", - uri, - headers={"Depth": "0", "Content-Type": "application/xml"}, - content=_PROPFIND_BODY, - ) - if response.is_error: - return None - entries = _parse_multistatus(response.content) + response = await self._http.request( + "PROPFIND", + uri, + headers={"Depth": "0", "Content-Type": "application/xml"}, + content=_PROPFIND_BODY, + ) + if response.is_error: + return None + entries = _parse_multistatus(response.content) if not entries: return None return entries[0].revision async def fetch(self, uri: str) -> FetchResult: - async with self._client() as http: - response = await http.get(uri) - response.raise_for_status() - body = response.content + response = await self._http.get(uri) + response.raise_for_status() + body = response.content content_type = ( response.headers.get("content-type", "application/octet-stream") .split(";")[0] @@ -261,15 +259,14 @@ class WebDAVSource: now = datetime.now(UTC) seen: set[str] = set() - async with self._client() as http: - response = await http.request( - "PROPFIND", - self.base_url, - headers={"Depth": "infinity", "Content-Type": "application/xml"}, - content=_PROPFIND_BODY, - ) - response.raise_for_status() - entries = _parse_multistatus(response.content) + response = await self._http.request( + "PROPFIND", + self.base_url, + headers={"Depth": "infinity", "Content-Type": "application/xml"}, + content=_PROPFIND_BODY, + ) + response.raise_for_status() + entries = _parse_multistatus(response.content) for entry in entries: if entry.is_collection: