requested updates

This commit is contained in:
bryan davis 2026-06-26 11:24:13 -05:00
parent 5a3483f0e2
commit eb8710d83a
No known key found for this signature in database
GPG key ID: D11B4A4C0C731E5E
5 changed files with 306 additions and 33 deletions

View file

@ -56,9 +56,15 @@ class DoclingServeChunker(DocumentChunker):
def __init__(self, config: AppConfig = Config): def __init__(self, config: AppConfig = Config):
self.config = config self.config = config
ds = config.providers.docling_serve
self.client = DoclingServeClient( self.client = DoclingServeClient(
base_urls=config.providers.docling_serve.base_urls, base_urls=ds.base_urls,
api_key=config.providers.docling_serve.api_key, api_key=ds.api_key,
max_attempts=ds.max_attempts,
retry_base_delay=ds.retry_base_delay,
retry_max_delay=ds.retry_max_delay,
breaker_failure_threshold=ds.breaker_failure_threshold,
breaker_cooldown_s=ds.breaker_cooldown_s,
) )
self.chunker_type = config.processing.chunker_type self.chunker_type = config.processing.chunker_type

View file

@ -227,6 +227,32 @@ class DoclingServeConfig(BaseModel):
base_url: str | list[str] = "http://localhost:5001" base_url: str | list[str] = "http://localhost:5001"
api_key: str = "" api_key: str = ""
max_attempts: int = Field(
default=3,
ge=1,
description="Total attempts per request before giving up; each retry "
"fails over to another instance (instances crash on memory leaks).",
)
retry_base_delay: float = Field(
default=0.5,
ge=0,
description="Base seconds for the exponential backoff between retries.",
)
retry_max_delay: float = Field(
default=8.0, ge=0, description="Cap on the retry backoff, in seconds."
)
breaker_failure_threshold: int = Field(
default=3,
ge=1,
description="Consecutive failures before an instance's circuit breaker "
"opens and routing skips it.",
)
breaker_cooldown_s: float = Field(
default=30.0,
ge=0,
description="How long an open breaker skips an instance before allowing "
"a probe.",
)
@property @property
def base_urls(self) -> list[str]: def base_urls(self) -> list[str]:

View file

@ -59,9 +59,15 @@ class DoclingServeConverter(DocumentConverter):
config: Application configuration containing docling-serve settings. config: Application configuration containing docling-serve settings.
""" """
self.config = config self.config = config
ds = config.providers.docling_serve
self.client = DoclingServeClient( self.client = DoclingServeClient(
base_urls=config.providers.docling_serve.base_urls, base_urls=ds.base_urls,
api_key=config.providers.docling_serve.api_key, api_key=ds.api_key,
max_attempts=ds.max_attempts,
retry_base_delay=ds.retry_base_delay,
retry_max_delay=ds.retry_max_delay,
breaker_failure_threshold=ds.breaker_failure_threshold,
breaker_cooldown_s=ds.breaker_cooldown_s,
) )
@property @property

View file

@ -3,6 +3,7 @@
import asyncio import asyncio
import itertools import itertools
import logging import logging
import time
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from typing import Any, TypeVar from typing import Any, TypeVar
@ -12,31 +13,75 @@ logger = logging.getLogger(__name__)
_T = TypeVar("_T") _T = TypeVar("_T")
# Statuses worth retrying on another instance: 408/429 are transient overload, # Statuses signalling an unhealthy instance: 408/429 are transient overload,
# 5xx covers a crashed or restarting docling-serve worker (the OOM-leak failure # 5xx covers a crashed or restarting docling-serve worker (the OOM-leak failure
# mode). Other 4xx are the caller's fault and won't succeed elsewhere. # mode). Other 4xx are the caller's fault and won't succeed elsewhere.
_RETRYABLE_STATUS = frozenset({408, 429}) _RETRYABLE_STATUS = frozenset({408, 429})
def _is_retryable(exc: BaseException) -> bool: def _is_retryable(exc: BaseException) -> bool:
"""Whether a failed docling-serve request should be retried on another """Whether a failure reflects an unhealthy docling-serve instance.
instance. True for transport-level failures (connection reset / timeout
an instance that crashed or went unresponsive) and server-side 5xx/overload; The single predicate behind both behaviours: such a failure is retried on
False for other 4xx and task-level ``ValueError``\\ s, which won't succeed on another instance *and* counted against the instance's circuit breaker. True
a retry (a deterministically bad document is handled upstream, not here).""" for transport-level failures (connection reset / timeout an instance that
crashed or went unresponsive) and server-side 5xx/overload; False for other
4xx and task-level ``ValueError``\\ s, which won't succeed on a retry and
aren't the instance's fault (a deterministically bad document is handled
upstream, not here)."""
if isinstance(exc, httpx.HTTPStatusError): if isinstance(exc, httpx.HTTPStatusError):
status = exc.response.status_code status = exc.response.status_code
return status in _RETRYABLE_STATUS or status >= 500 return status in _RETRYABLE_STATUS or status >= 500
return isinstance(exc, httpx.TransportError) return isinstance(exc, httpx.TransportError)
# Process-global registry of round-robin rotators over docling-serve class _InstanceBreaker:
# instances, keyed by the sorted tuple of base URLs. Clients are constructed """Per-instance circuit breaker (closed / open with cooldown).
# per-job by `get_converter` / `get_chunker`, so the rotation cursor has to
# live OUTSIDE the instance to actually advance across jobs. Two clients Self-contained so the provider layer doesn't depend on the ingester
# pointing at the same instance set share one rotator; clients with package. ``is_open`` auto-probes after the cooldown elapses: a single
# different sets get independent rotators. request is allowed through, and its success closes the breaker while another
failure re-opens it.
"""
def __init__(
self,
failure_threshold: int,
cooldown_s: float,
now_fn: Callable[[], float],
):
self._threshold = failure_threshold
self._cooldown_s = cooldown_s
self._now = now_fn
self._consecutive_failures = 0
self._opened_at: float | None = None
@property
def is_open(self) -> bool:
if self._opened_at is None:
return False
# Cooldown elapsed → let the next request probe the instance.
return self._now() - self._opened_at < self._cooldown_s
def record_success(self) -> None:
self._consecutive_failures = 0
self._opened_at = None
def record_failure(self) -> None:
self._consecutive_failures += 1
if self._consecutive_failures >= self._threshold:
self._opened_at = self._now()
# Process-global registries over docling-serve instances. Clients are
# constructed per-job by `get_converter` / `get_chunker`, so this shared state
# has to live OUTSIDE the instance to persist across jobs: the round-robin
# cursor (keyed by the sorted tuple of base URLs) advances across jobs, and the
# per-URL breakers keep a crashed instance skipped across subsequent jobs until
# its cooldown elapses. The first client to touch a URL fixes that breaker's
# threshold/cooldown/clock; in practice all clients share one config.
_instance_rotators: dict[tuple[str, ...], "itertools.cycle[str]"] = {} _instance_rotators: dict[tuple[str, ...], "itertools.cycle[str]"] = {}
_instance_breakers: dict[str, "_InstanceBreaker"] = {}
class DoclingServeClient: class DoclingServeClient:
@ -58,6 +103,9 @@ class DoclingServeClient:
max_attempts: int = 3, max_attempts: int = 3,
retry_base_delay: float = 0.5, retry_base_delay: float = 0.5,
retry_max_delay: float = 8.0, retry_max_delay: float = 8.0,
breaker_failure_threshold: int = 3,
breaker_cooldown_s: float = 30.0,
now_fn: Callable[[], float] = time.monotonic,
): ):
urls = [base_urls] if isinstance(base_urls, str) else list(base_urls) urls = [base_urls] if isinstance(base_urls, str) else list(base_urls)
if not urls: if not urls:
@ -67,12 +115,17 @@ class DoclingServeClient:
self.timeout = timeout self.timeout = timeout
# transport is for testing — production callers leave it None. # transport is for testing — production callers leave it None.
self._transport = transport self._transport = transport
# Bounded retry with failover: docling-serve instances crash (memory # Bounded retry with failover + per-instance circuit breaking:
# leaks), so a request that hits a dying instance is retried, preferring # docling-serve instances crash (memory leaks), so a request that hits a
# an instance that hasn't already failed this request. # dying instance is retried on another instance, and repeated failures
# trip that instance's breaker so later requests skip it while it
# recovers.
self._max_attempts = max(1, max_attempts) self._max_attempts = max(1, max_attempts)
self._retry_base_delay = retry_base_delay self._retry_base_delay = retry_base_delay
self._retry_max_delay = retry_max_delay self._retry_max_delay = retry_max_delay
self._breaker_failure_threshold = breaker_failure_threshold
self._breaker_cooldown_s = breaker_cooldown_s
self._now = now_fn
# setdefault is atomic under the GIL — concurrent constructors with # setdefault is atomic under the GIL — concurrent constructors with
# the same instance set will end up sharing one rotator. # the same instance set will end up sharing one rotator.
key = tuple(sorted(self.base_urls)) key = tuple(sorted(self.base_urls))
@ -89,25 +142,59 @@ class DoclingServeClient:
callers should let `_pick_url` round-robin per request.""" callers should let `_pick_url` round-robin per request."""
return self.base_urls[0] return self.base_urls[0]
def _breaker_for(self, url: str) -> _InstanceBreaker:
breaker = _instance_breakers.get(url)
if breaker is None:
breaker = _InstanceBreaker(
self._breaker_failure_threshold,
self._breaker_cooldown_s,
self._now,
)
_instance_breakers[url] = breaker
return breaker
def _pick_url(self, exclude: frozenset[str] = frozenset()) -> str: def _pick_url(self, exclude: frozenset[str] = frozenset()) -> str:
"""Next instance in the round-robin, preferring one not in ``exclude`` """Next instance in the round-robin, skipping instances that already
(instances that already failed this request). Falls back to a failed this request (``exclude``) or whose circuit breaker is open
possibly-excluded instance when every instance is excluded a (recently crashed / overloaded).
single-instance fleet, or one where all instances failed, still gets
retried after backoff in case the instance has since restarted.""" Prefers a not-excluded, breaker-closed instance. Falls back to a
not-excluded breaker-open one (worth a probe over re-hitting one that
already failed this request), and finally to any instance so a
single-instance fleet, or one where everything is excluded/open, still
gets an attempt rather than nothing to pick."""
not_excluded: str | None = None
for _ in range(len(self.base_urls)):
url = next(self._instance_rotator) url = next(self._instance_rotator)
if url not in exclude: if url in exclude:
return url continue
for _ in range(len(self.base_urls) - 1): if not_excluded is None:
url = next(self._instance_rotator) not_excluded = url
if url not in exclude: if not self._breaker_for(url).is_open:
return url
return url return url
if not_excluded is not None:
return not_excluded
return next(self._instance_rotator)
def _retry_delay(self, attempt_no: int) -> float: def _retry_delay(self, attempt_no: int) -> float:
"""Capped exponential backoff between retry attempts.""" """Capped exponential backoff between retry attempts."""
return min(self._retry_base_delay * (2**attempt_no), self._retry_max_delay) return min(self._retry_base_delay * (2**attempt_no), self._retry_max_delay)
def _record_failure(self, base_url: str) -> None:
"""Count an instance-health failure against ``base_url``'s breaker,
logging when the breaker transitions to open."""
breaker = self._breaker_for(base_url)
was_open = breaker.is_open
breaker.record_failure()
if not was_open and breaker.is_open:
logger.warning(
"docling-serve instance %s breaker opened after %d consecutive "
"failure(s); skipping it for %.0fs",
base_url,
self._breaker_failure_threshold,
self._breaker_cooldown_s,
)
def _get_headers(self) -> dict[str, str]: def _get_headers(self) -> dict[str, str]:
"""Get headers for API requests.""" """Get headers for API requests."""
headers: dict[str, str] = {} headers: dict[str, str] = {}
@ -126,8 +213,11 @@ class DoclingServeClient:
instance (task IDs are instance-local, so a trio can't be split across instance (task IDs are instance-local, so a trio can't be split across
instances). On a retryable failure a transport error (crashed / instances). On a retryable failure a transport error (crashed /
unresponsive instance) or 5xx/overload the next attempt prefers an unresponsive instance) or 5xx/overload the next attempt prefers an
instance that hasn't already failed this request. Non-retryable errors instance that hasn't already failed this request, and the failure is
(4xx, task-level ``ValueError``) propagate immediately. counted against that instance's circuit breaker so later requests skip
it while it recovers; a successful attempt closes the breaker.
Non-retryable errors (4xx, task-level ``ValueError``) propagate
immediately and leave the breaker untouched.
""" """
tried: set[str] = set() tried: set[str] = set()
last_exc: BaseException | None = None last_exc: BaseException | None = None
@ -135,10 +225,13 @@ class DoclingServeClient:
base_url = self._pick_url(exclude=frozenset(tried)) base_url = self._pick_url(exclude=frozenset(tried))
try: try:
async with self._httpx_client() as client: async with self._httpx_client() as client:
return await attempt(client, base_url) result = await attempt(client, base_url)
except Exception as exc: except Exception as exc:
if not _is_retryable(exc): if not _is_retryable(exc):
raise raise
# Retryable == an unhealthy instance: count it against the
# breaker so later requests route around it.
self._record_failure(base_url)
last_exc = exc last_exc = exc
tried.add(base_url) tried.add(base_url)
remaining = self._max_attempts - attempt_no - 1 remaining = self._max_attempts - attempt_no - 1
@ -153,6 +246,9 @@ class DoclingServeClient:
remaining, remaining,
) )
await asyncio.sleep(self._retry_delay(attempt_no)) await asyncio.sleep(self._retry_delay(attempt_no))
else:
self._breaker_for(base_url).record_success()
return result
# Unreachable: max_attempts >= 1, and the final iteration either returns # Unreachable: max_attempts >= 1, and the final iteration either returns
# on success or re-raises on failure (remaining <= 0). # on success or re-raises on failure (remaining <= 0).
raise last_exc or RuntimeError( # pragma: no cover raise last_exc or RuntimeError( # pragma: no cover

View file

@ -321,6 +321,145 @@ def test_pick_url_all_excluded_falls_back():
assert picked == "http://solo:5001" assert picked == "http://solo:5001"
async def _convert(client: DoclingServeClient) -> dict:
return await client.submit_and_poll(
endpoint="/v1/convert/file/async",
files={"file": ("x.md", b"x", "text/markdown")},
data={},
)
@pytest.mark.asyncio
async def test_breaker_opens_after_repeated_failures_then_skips_instance():
"""Repeated failures on an instance (each failing over to a healthy one)
trip its breaker; once open, _pick_url skips it entirely."""
down = {"crash-h"}
transport, seen = _failover_transport(down, "t", {"ok": True})
crash = "http://crash-h:5001"
client = DoclingServeClient(
base_urls=[crash, "http://live-h:5001"],
transport=transport,
breaker_failure_threshold=2,
retry_base_delay=0.0,
)
# Each call fails on crash-h (recorded) then fails over to live-h.
assert await _convert(client) == {"ok": True}
assert await _convert(client) == {"ok": True} # 2nd failure opens crash-h
assert client._breaker_for(crash).is_open
seen.clear()
assert await _convert(client) == {"ok": True}
assert "crash-h" not in seen # breaker open → never attempted
assert set(seen) == {"live-h"}
@pytest.mark.asyncio
async def test_breaker_recovers_after_cooldown():
"""An open breaker auto-probes after its cooldown; once the instance is
healthy again a successful request closes it."""
clock = [1000.0]
down = {"flip-j"}
transport, seen = _failover_transport(down, "t", {"ok": True})
flip = "http://flip-j:5001"
client = DoclingServeClient(
base_urls=[flip, "http://spare-j:5001"],
transport=transport,
breaker_failure_threshold=2,
breaker_cooldown_s=30.0,
retry_base_delay=0.0,
now_fn=lambda: clock[0],
)
await _convert(client)
await _convert(client) # opens flip-j's breaker
assert client._breaker_for(flip).is_open
down.clear() # instance recovers
assert client._breaker_for(flip).is_open # still skipped within cooldown
clock[0] += 31.0
assert not client._breaker_for(flip).is_open # cooldown elapsed → probe ok
seen.clear()
await _convert(client)
await _convert(client)
assert "flip-j" in seen # traffic returned
assert not client._breaker_for(flip).is_open # success closed it
@pytest.mark.asyncio
async def test_4xx_does_not_trip_breaker():
"""A 4xx is the caller's fault — it must not count against instance health,
even at a 1-failure threshold."""
seen: list[str] = []
def handler(request: httpx.Request) -> httpx.Response:
seen.append(request.url.host)
return httpx.Response(400, json={"detail": "bad request"})
bad = "http://bad-k:5001"
client = DoclingServeClient(
base_urls=[bad],
transport=httpx.MockTransport(handler),
breaker_failure_threshold=1,
retry_base_delay=0.0,
)
for _ in range(3):
with pytest.raises(httpx.HTTPStatusError):
await _convert(client)
assert not client._breaker_for(bad).is_open
@pytest.mark.asyncio
async def test_pick_url_falls_back_when_all_breakers_open():
"""When every instance's breaker is open (here a single-instance fleet that
just failed), _pick_url returns one anyway so the request can still probe
it rather than having nothing to pick."""
lone = "http://lone-m:5001"
transport, _ = _failover_transport({"lone-m"}, "t", {"ok": True})
client = DoclingServeClient(
base_urls=[lone],
transport=transport,
breaker_failure_threshold=1,
max_attempts=1,
retry_base_delay=0.0,
)
with pytest.raises(httpx.ConnectError):
await _convert(client)
assert client._breaker_for(lone).is_open
# All (one) instances open → fall back to it.
assert client._pick_url() == lone
def test_config_knobs_reach_the_client():
"""Retry/breaker knobs set in DoclingServeConfig must reach the client via
get_converter / get_chunker (previously hardcoded constructor defaults)."""
from haiku.rag.chunkers.docling_serve import DoclingServeChunker
from haiku.rag.config import AppConfig
from haiku.rag.converters.docling_serve import DoclingServeConverter
config = AppConfig()
ds = config.providers.docling_serve
ds.base_url = "http://cfg-n:5001"
ds.max_attempts = 7
ds.retry_base_delay = 1.25
ds.retry_max_delay = 20.0
ds.breaker_failure_threshold = 9
ds.breaker_cooldown_s = 90.0
for component in (DoclingServeConverter(config), DoclingServeChunker(config)):
client = component.client
assert client._max_attempts == 7
assert client._retry_base_delay == 1.25
assert client._retry_max_delay == 20.0
assert client._breaker_failure_threshold == 9
assert client._breaker_cooldown_s == 90.0
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_zip_endpoint_uses_round_robin_too(): async def test_zip_endpoint_uses_round_robin_too():
transport, seen = _scripted_transport( transport, seen = _scripted_transport(