Merge pull request #395 from mcdonc/perf/reuse-httpx-clients
perf: reuse httpx.AsyncClient in HTTP and WebDAV sources
This commit is contained in:
commit
e9875fc842
9 changed files with 174 additions and 101 deletions
|
|
@ -181,6 +181,7 @@ class IngesterApp:
|
||||||
await asyncio.gather(api_task, return_exceptions=True)
|
await asyncio.gather(api_task, return_exceptions=True)
|
||||||
await self._pollers.stop()
|
await self._pollers.stop()
|
||||||
await self._stop_pool()
|
await self._stop_pool()
|
||||||
|
await self._pollers.close_sources()
|
||||||
|
|
||||||
async def run_batch(self) -> BatchReport:
|
async def run_batch(self) -> BatchReport:
|
||||||
"""Run one discover() sweep across every configured source, drain the
|
"""Run one discover() sweep across every configured source, drain the
|
||||||
|
|
@ -213,6 +214,7 @@ class IngesterApp:
|
||||||
)
|
)
|
||||||
finally:
|
finally:
|
||||||
await self._stop_pool()
|
await self._stop_pool()
|
||||||
|
await self._pollers.close_sources()
|
||||||
|
|
||||||
async def _maybe_start_api(self, api: bool):
|
async def _maybe_start_api(self, api: bool):
|
||||||
"""Spin up the FastAPI control plane on an asyncio task. Returns
|
"""Spin up the FastAPI control plane on an asyncio task. Returns
|
||||||
|
|
|
||||||
|
|
@ -97,6 +97,13 @@ class PollerManager:
|
||||||
self._tasks.clear()
|
self._tasks.clear()
|
||||||
self._started = False
|
self._started = False
|
||||||
|
|
||||||
|
async def close_sources(self) -> None:
|
||||||
|
"""Close all source adapters (e.g. HTTP connection pools). Must be
|
||||||
|
called after the worker pool has fully stopped so in-flight fetches
|
||||||
|
don't hit a closed client."""
|
||||||
|
for source in self.sources:
|
||||||
|
await source.aclose()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def pollers(self) -> list[BasePoller]:
|
def pollers(self) -> list[BasePoller]:
|
||||||
return list(self._pollers)
|
return list(self._pollers)
|
||||||
|
|
|
||||||
|
|
@ -61,6 +61,11 @@ class Source(Protocol):
|
||||||
"""
|
"""
|
||||||
...
|
...
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
"""Release any resources held by the source (e.g. HTTP connection
|
||||||
|
pools). Called once during shutdown, after all workers have stopped."""
|
||||||
|
...
|
||||||
|
|
||||||
async def fetch(self, uri: str) -> FetchResult: ...
|
async def fetch(self, uri: str) -> FetchResult: ...
|
||||||
|
|
||||||
def discover(
|
def discover(
|
||||||
|
|
|
||||||
|
|
@ -75,6 +75,9 @@ class FSSource:
|
||||||
return False
|
return False
|
||||||
return self._resolve_within_root(uri) is not None
|
return self._resolve_within_root(uri) is not None
|
||||||
|
|
||||||
|
async def aclose(self) -> None: # pragma: no cover - no resources to release
|
||||||
|
pass
|
||||||
|
|
||||||
async def head(self, uri: str) -> str | None:
|
async def head(self, uri: str) -> str | None:
|
||||||
path = self._resolve_within_root(uri)
|
path = self._resolve_within_root(uri)
|
||||||
if path is None or not path.exists():
|
if path is None or not path.exists():
|
||||||
|
|
|
||||||
|
|
@ -39,48 +39,44 @@ class HTTPSource:
|
||||||
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 {})
|
||||||
# transport is for testing — production callers leave it None and httpx
|
self._http = httpx.AsyncClient(headers=self.headers, transport=transport)
|
||||||
# uses its real transport.
|
|
||||||
self._transport = transport
|
|
||||||
|
|
||||||
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")
|
||||||
|
|
||||||
def _client(self) -> httpx.AsyncClient:
|
async def aclose(self) -> None:
|
||||||
return httpx.AsyncClient(headers=self.headers, transport=self._transport)
|
await self._http.aclose()
|
||||||
|
|
||||||
async def head(self, uri: str) -> str | None:
|
async def head(self, uri: str) -> str | None:
|
||||||
"""HEAD probe for the cheap revision short-circuit. Returns the
|
"""HEAD probe for the cheap revision short-circuit. Returns the
|
||||||
ETag (or Last-Modified) so an unchanged remote URL can skip 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();
|
full GET. None on HTTP error so the caller falls back to fetch();
|
||||||
network errors propagate and the worker's classifier handles them."""
|
network errors propagate and the worker's classifier handles them."""
|
||||||
async with self._client() as http:
|
response = await self._http.head(uri)
|
||||||
response = await http.head(uri)
|
if response.is_error:
|
||||||
if response.is_error:
|
return None
|
||||||
return None
|
|
||||||
revision, _ = _extract_revision(response.headers)
|
revision, _ = _extract_revision(response.headers)
|
||||||
return revision
|
return revision
|
||||||
|
|
||||||
async def fetch(self, uri: str) -> FetchResult:
|
async def fetch(self, uri: str) -> FetchResult:
|
||||||
async with self._client() as http:
|
response = await self._http.get(uri)
|
||||||
response = await http.get(uri)
|
response.raise_for_status()
|
||||||
response.raise_for_status()
|
body = response.content
|
||||||
body = response.content
|
content_type = (
|
||||||
content_type = (
|
response.headers.get("content-type", "application/octet-stream")
|
||||||
response.headers.get("content-type", "application/octet-stream")
|
.split(";")[0]
|
||||||
.split(";")[0]
|
.strip()
|
||||||
.strip()
|
.lower()
|
||||||
.lower()
|
)
|
||||||
)
|
revision, extra = _extract_revision(response.headers)
|
||||||
revision, extra = _extract_revision(response.headers)
|
return FetchResult(
|
||||||
return FetchResult(
|
uri=uri,
|
||||||
uri=uri,
|
body=body,
|
||||||
body=body,
|
content_type=content_type,
|
||||||
content_type=content_type,
|
content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(),
|
||||||
content_hash=hashlib.md5(body, usedforsecurity=False).hexdigest(),
|
revision=revision,
|
||||||
revision=revision,
|
extra_metadata=extra,
|
||||||
extra_metadata=extra,
|
)
|
||||||
)
|
|
||||||
|
|
||||||
async def discover(
|
async def discover(
|
||||||
self,
|
self,
|
||||||
|
|
@ -104,53 +100,52 @@ class HTTPSource:
|
||||||
now = datetime.now(UTC)
|
now = datetime.now(UTC)
|
||||||
configured = set(self.urls)
|
configured = set(self.urls)
|
||||||
|
|
||||||
async with self._client() as http:
|
for url in self.urls:
|
||||||
for url in self.urls:
|
try:
|
||||||
try:
|
head = await self._http.head(url)
|
||||||
head = await http.head(url)
|
except Exception:
|
||||||
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
|
|
||||||
|
|
||||||
yield SourceEvent(
|
yield SourceEvent(
|
||||||
source_id=self.source_id,
|
source_id=self.source_id,
|
||||||
uri=url,
|
uri=url,
|
||||||
kind=kind,
|
kind=SourceEventKind.UPSERT,
|
||||||
revision=revision,
|
revision=None,
|
||||||
discovered_at=now,
|
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
|
# Anything previously known to this source that's no longer in
|
||||||
# config emits DELETE so delete_orphans can clean up. Without this,
|
# config emits DELETE so delete_orphans can clean up. Without this,
|
||||||
|
|
|
||||||
|
|
@ -63,6 +63,9 @@ class S3Source:
|
||||||
def supports(self, uri: str) -> bool:
|
def supports(self, uri: str) -> bool:
|
||||||
return uri.startswith(self.uri_prefix)
|
return uri.startswith(self.uri_prefix)
|
||||||
|
|
||||||
|
async def aclose(self) -> None: # pragma: no cover - no resources to release
|
||||||
|
pass
|
||||||
|
|
||||||
async def head(self, uri: str) -> str | None:
|
async def head(self, uri: str) -> str | None:
|
||||||
import obstore # type: ignore[import-not-found]
|
import obstore # type: ignore[import-not-found]
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -189,41 +189,39 @@ 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._transport = transport
|
|
||||||
|
|
||||||
def supports(self, uri: str) -> bool:
|
|
||||||
return uri.startswith(self.base_url)
|
|
||||||
|
|
||||||
def _client(self) -> httpx.AsyncClient:
|
|
||||||
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
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
return httpx.AsyncClient(
|
self._http = httpx.AsyncClient(
|
||||||
auth=auth, headers=self.headers, transport=self._transport
|
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 def head(self, uri: str) -> str | None:
|
||||||
async with self._client() as http:
|
response = await self._http.request(
|
||||||
response = await http.request(
|
"PROPFIND",
|
||||||
"PROPFIND",
|
uri,
|
||||||
uri,
|
headers={"Depth": "0", "Content-Type": "application/xml"},
|
||||||
headers={"Depth": "0", "Content-Type": "application/xml"},
|
content=_PROPFIND_BODY,
|
||||||
content=_PROPFIND_BODY,
|
)
|
||||||
)
|
if response.is_error:
|
||||||
if response.is_error:
|
return None
|
||||||
return None
|
entries = _parse_multistatus(response.content)
|
||||||
entries = _parse_multistatus(response.content)
|
|
||||||
if not entries:
|
if not entries:
|
||||||
return None
|
return None
|
||||||
return entries[0].revision
|
return entries[0].revision
|
||||||
|
|
||||||
async def fetch(self, uri: str) -> FetchResult:
|
async def fetch(self, uri: str) -> FetchResult:
|
||||||
async with self._client() as http:
|
response = await self._http.get(uri)
|
||||||
response = await http.get(uri)
|
response.raise_for_status()
|
||||||
response.raise_for_status()
|
body = response.content
|
||||||
body = response.content
|
|
||||||
content_type = (
|
content_type = (
|
||||||
response.headers.get("content-type", "application/octet-stream")
|
response.headers.get("content-type", "application/octet-stream")
|
||||||
.split(";")[0]
|
.split(";")[0]
|
||||||
|
|
@ -261,15 +259,14 @@ class WebDAVSource:
|
||||||
now = datetime.now(UTC)
|
now = datetime.now(UTC)
|
||||||
seen: set[str] = set()
|
seen: set[str] = set()
|
||||||
|
|
||||||
async with self._client() as http:
|
response = await self._http.request(
|
||||||
response = await http.request(
|
"PROPFIND",
|
||||||
"PROPFIND",
|
self.base_url,
|
||||||
self.base_url,
|
headers={"Depth": "infinity", "Content-Type": "application/xml"},
|
||||||
headers={"Depth": "infinity", "Content-Type": "application/xml"},
|
content=_PROPFIND_BODY,
|
||||||
content=_PROPFIND_BODY,
|
)
|
||||||
)
|
response.raise_for_status()
|
||||||
response.raise_for_status()
|
entries = _parse_multistatus(response.content)
|
||||||
entries = _parse_multistatus(response.content)
|
|
||||||
|
|
||||||
for entry in entries:
|
for entry in entries:
|
||||||
if entry.is_collection:
|
if entry.is_collection:
|
||||||
|
|
|
||||||
|
|
@ -23,6 +23,7 @@ from haiku.rag.config import (
|
||||||
WorkerConfig,
|
WorkerConfig,
|
||||||
)
|
)
|
||||||
from haiku.rag.ingester.app import IngesterApp
|
from haiku.rag.ingester.app import IngesterApp
|
||||||
|
from haiku.rag.ingester.pollers.manager import PollerManager
|
||||||
from haiku.rag.ingester.workers.pool import WorkerPool
|
from haiku.rag.ingester.workers.pool import WorkerPool
|
||||||
from haiku.rag.store.models.document import Document
|
from haiku.rag.store.models.document import Document
|
||||||
|
|
||||||
|
|
@ -305,3 +306,60 @@ async def test_stop_pool_warns_when_shutdown_grace_elapses(tmp_path, caplog):
|
||||||
|
|
||||||
assert pool.released == 1
|
assert pool.released == 1
|
||||||
assert "Shutdown grace" in caplog.text
|
assert "Shutdown grace" in caplog.text
|
||||||
|
|
||||||
|
|
||||||
|
def _record_close_order(monkeypatch) -> list[str]:
|
||||||
|
"""Record the order of _stop_pool and PollerManager.close_sources.
|
||||||
|
|
||||||
|
Workers share the pollers' Source instances for fetch(), so the source
|
||||||
|
clients must be closed only after the pool has stopped — otherwise an
|
||||||
|
in-flight fetch during the shutdown grace hits a closed client.
|
||||||
|
"""
|
||||||
|
order: list[str] = []
|
||||||
|
orig_stop_pool = IngesterApp._stop_pool
|
||||||
|
orig_close = PollerManager.close_sources
|
||||||
|
|
||||||
|
async def rec_stop(self):
|
||||||
|
order.append("stop_pool")
|
||||||
|
await orig_stop_pool(self)
|
||||||
|
|
||||||
|
async def rec_close(self):
|
||||||
|
order.append("close_sources")
|
||||||
|
await orig_close(self)
|
||||||
|
|
||||||
|
monkeypatch.setattr(IngesterApp, "_stop_pool", rec_stop)
|
||||||
|
monkeypatch.setattr(PollerManager, "close_sources", rec_close)
|
||||||
|
return order
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_run_batch_closes_sources_after_pool_stops(
|
||||||
|
tmp_path, use_client, monkeypatch
|
||||||
|
):
|
||||||
|
(tmp_path / "a.md").write_text("hello")
|
||||||
|
use_client(_mock_client())
|
||||||
|
app = IngesterApp(config=_config(tmp_path), db_path=tmp_path / "db.lancedb")
|
||||||
|
order = _record_close_order(monkeypatch)
|
||||||
|
|
||||||
|
await app.run_batch()
|
||||||
|
|
||||||
|
assert order == ["stop_pool", "close_sources"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_serve_closes_sources_after_pool_stops(tmp_path, use_client, monkeypatch):
|
||||||
|
use_client(_mock_client())
|
||||||
|
config = _config(tmp_path)
|
||||||
|
config.ingester.api = APIConfig(enabled=False)
|
||||||
|
app = IngesterApp(config=config, db_path=tmp_path / "db.lancedb")
|
||||||
|
order = _record_close_order(monkeypatch)
|
||||||
|
|
||||||
|
task = asyncio.create_task(app.serve(api=False))
|
||||||
|
try:
|
||||||
|
await _wait_until(lambda: app._pool is not None and app._pool.live_workers > 0)
|
||||||
|
finally:
|
||||||
|
task.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await asyncio.wait_for(task, timeout=5.0)
|
||||||
|
|
||||||
|
assert order == ["stop_pool", "close_sources"]
|
||||||
|
|
|
||||||
|
|
@ -59,6 +59,9 @@ def test_source_protocol_runtime_checkable():
|
||||||
def supports(self, uri: str) -> bool:
|
def supports(self, uri: str) -> bool:
|
||||||
return True
|
return True
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
async def head(self, uri: str):
|
async def head(self, uri: str):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue