Leave no database reading when a federated fan-out fails
`asyncio.gather` propagates the first failure while its siblings run on, and the caller unwinding from that closes the set through `async with` — so a sibling still reading reads through a closed session. Seven fan-outs were affected: lookup, search, image enrichment, multimodal picture loading, context expansion, document listing and counting, and the sandbox's document load. `gather_all` cancels and drains the rest, then re-raises the original exception. A `TaskGroup` would drain them too but raise an `ExceptionGroup`, which every caller and both CLIs' exception handlers would have to unwrap. The two `return_exceptions=True` gathers in session opening and teardown already drain their children and are left alone.
This commit is contained in:
parent
225a37ca73
commit
bfd09f6880
6 changed files with 120 additions and 12 deletions
|
|
@ -43,7 +43,7 @@ from haiku.rag.store.repositories.chunk import ChunkRepository
|
||||||
from haiku.rag.store.repositories.document import DocumentRepository
|
from haiku.rag.store.repositories.document import DocumentRepository
|
||||||
from haiku.rag.store.repositories.document_item import DocumentItemRepository
|
from haiku.rag.store.repositories.document_item import DocumentItemRepository
|
||||||
from haiku.rag.store.repositories.settings import SettingsRepository
|
from haiku.rag.store.repositories.settings import SettingsRepository
|
||||||
from haiku.rag.utils import escape_sql_string
|
from haiku.rag.utils import escape_sql_string, gather_all
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from docling_core.types.doc.document import DoclingDocument
|
from docling_core.types.doc.document import DoclingDocument
|
||||||
|
|
@ -69,7 +69,7 @@ async def all_found(
|
||||||
asked at once. Asking in turn would cost a round trip per database for an
|
asked at once. Asking in turn would cost a round trip per database for an
|
||||||
identifier that is missing or held by the last of them.
|
identifier that is missing or held by the last of them.
|
||||||
"""
|
"""
|
||||||
found_by_client = await asyncio.gather(*(lookup(client) for client in clients))
|
found_by_client = await gather_all(*(lookup(client) for client in clients))
|
||||||
return [
|
return [
|
||||||
(client, found)
|
(client, found)
|
||||||
for client, found in zip(clients, found_by_client, strict=True)
|
for client, found in zip(clients, found_by_client, strict=True)
|
||||||
|
|
@ -784,7 +784,7 @@ class HaikuRAG:
|
||||||
# the window is applied to the merged listing: a limit means that
|
# the window is applied to the merged listing: a limit means that
|
||||||
# many documents in total, not that many per database.
|
# many documents in total, not that many per database.
|
||||||
wanted = None if limit is None else limit + (offset or 0)
|
wanted = None if limit is None else limit + (offset or 0)
|
||||||
groups = await asyncio.gather(
|
groups = await gather_all(
|
||||||
*(
|
*(
|
||||||
owner.list_documents(
|
owner.list_documents(
|
||||||
limit=wanted, filter=filter, include_content=include_content
|
limit=wanted, filter=filter, include_content=include_content
|
||||||
|
|
@ -813,7 +813,7 @@ class HaikuRAG:
|
||||||
Number of documents matching the criteria.
|
Number of documents matching the criteria.
|
||||||
"""
|
"""
|
||||||
if self.covers_multiple:
|
if self.covers_multiple:
|
||||||
counts = await asyncio.gather(
|
counts = await gather_all(
|
||||||
*(
|
*(
|
||||||
owner.count_documents(filter=filter)
|
owner.count_documents(filter=filter)
|
||||||
for owner in await self.clients_covering()
|
for owner in await self.clients_covering()
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,3 @@
|
||||||
import asyncio
|
|
||||||
import base64
|
import base64
|
||||||
from collections.abc import Sequence
|
from collections.abc import Sequence
|
||||||
from typing import TYPE_CHECKING
|
from typing import TYPE_CHECKING
|
||||||
|
|
@ -10,6 +9,7 @@ from haiku.rag.store.models.chunk import (
|
||||||
qualified_id,
|
qualified_id,
|
||||||
)
|
)
|
||||||
from haiku.rag.store.models.document_item import PICTURE_REF_PREFIX
|
from haiku.rag.store.models.document_item import PICTURE_REF_PREFIX
|
||||||
|
from haiku.rag.utils import gather_all
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from PIL import Image as PILImage
|
from PIL import Image as PILImage
|
||||||
|
|
@ -109,7 +109,7 @@ async def search_sources(
|
||||||
fetch_limit = _fetch_limit(client, query, limit)
|
fetch_limit = _fetch_limit(client, query, limit)
|
||||||
query_vector = await _embed_query(selected[0], query, resolved)
|
query_vector = await _embed_query(selected[0], query, resolved)
|
||||||
text = query if isinstance(query, str) else ""
|
text = query if isinstance(query, str) else ""
|
||||||
per_source = await asyncio.gather(
|
per_source = await gather_all(
|
||||||
*(
|
*(
|
||||||
c.chunk_repository.search(
|
c.chunk_repository.search(
|
||||||
query=text,
|
query=text,
|
||||||
|
|
@ -137,7 +137,7 @@ async def search_sources(
|
||||||
if result.source:
|
if result.source:
|
||||||
by_owner.setdefault(result.source, []).append(result)
|
by_owner.setdefault(result.source, []).append(result)
|
||||||
owners = await client.clients_for(list(by_owner))
|
owners = await client.clients_for(list(by_owner))
|
||||||
await asyncio.gather(
|
await gather_all(
|
||||||
*(
|
*(
|
||||||
_populate_image_data(owner, by_owner[name])
|
_populate_image_data(owner, by_owner[name])
|
||||||
for name, owner in zip(by_owner, owners, strict=True)
|
for name, owner in zip(by_owner, owners, strict=True)
|
||||||
|
|
@ -176,7 +176,7 @@ async def _fuse(
|
||||||
if reranker is not None:
|
if reranker is not None:
|
||||||
chunks = [chunk for _, chunk, _ in owned]
|
chunks = [chunk for _, chunk, _ in owned]
|
||||||
if federator._config.reranking.multimodal:
|
if federator._config.reranking.multimodal:
|
||||||
await asyncio.gather(
|
await gather_all(
|
||||||
*(
|
*(
|
||||||
_attach_picture_data(
|
_attach_picture_data(
|
||||||
c, [chunk for owner, chunk, _ in owned if owner is c]
|
c, [chunk for owner, chunk, _ in owned if owner is c]
|
||||||
|
|
@ -438,7 +438,7 @@ async def expand_sources(
|
||||||
unsourced.append(result)
|
unsourced.append(result)
|
||||||
names = list(by_source)
|
names = list(by_source)
|
||||||
sessions = await federated.sessions_for(names)
|
sessions = await federated.sessions_for(names)
|
||||||
expanded_groups = await asyncio.gather(
|
expanded_groups = await gather_all(
|
||||||
*(
|
*(
|
||||||
expand_context(session, by_source[name])
|
expand_context(session, by_source[name])
|
||||||
for name, session in zip(names, sessions, strict=True)
|
for name, session in zip(names, sessions, strict=True)
|
||||||
|
|
|
||||||
|
|
@ -20,6 +20,7 @@ from haiku.rag.config.models import AppConfig
|
||||||
from haiku.rag.sandbox.dependencies import AnalysisContext
|
from haiku.rag.sandbox.dependencies import AnalysisContext
|
||||||
from haiku.rag.store.models.chunk import SearchResult
|
from haiku.rag.store.models.chunk import SearchResult
|
||||||
from haiku.rag.store.models.document_item import PICTURE_REF_PREFIX, DocumentItem
|
from haiku.rag.store.models.document_item import PICTURE_REF_PREFIX, DocumentItem
|
||||||
|
from haiku.rag.utils import gather_all
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from pathlib import Path, PurePosixPath
|
from pathlib import Path, PurePosixPath
|
||||||
|
|
@ -280,7 +281,7 @@ class Sandbox:
|
||||||
owners = await rag.clients_covering(self._context.sources)
|
owners = await rag.clients_covering(self._context.sources)
|
||||||
# Resolve owners under the shared-client lock; owner sessions perform
|
# Resolve owners under the shared-client lock; owner sessions perform
|
||||||
# reads independently.
|
# reads independently.
|
||||||
groups = await asyncio.gather(
|
groups = await gather_all(
|
||||||
*(owner.list_documents(filter=self._context.filter) for owner in owners)
|
*(owner.list_documents(filter=self._context.filter) for owner in owners)
|
||||||
)
|
)
|
||||||
# Interleaved, not concatenated: code that prints the listing is read
|
# Interleaved, not concatenated: code that prints the listing is read
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,7 @@
|
||||||
|
import asyncio
|
||||||
import math
|
import math
|
||||||
import sys
|
import sys
|
||||||
|
from collections.abc import Awaitable
|
||||||
from importlib import metadata
|
from importlib import metadata
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import TYPE_CHECKING, Any, NoReturn, cast
|
from typing import TYPE_CHECKING, Any, NoReturn, cast
|
||||||
|
|
@ -16,6 +18,22 @@ if TYPE_CHECKING:
|
||||||
from haiku.rag.store.models.citation import Citation
|
from haiku.rag.store.models.citation import Citation
|
||||||
|
|
||||||
|
|
||||||
|
async def gather_all[T](*awaitables: Awaitable[T]) -> list[T]:
|
||||||
|
"""Run `awaitables` concurrently, leaving none of them running when one fails.
|
||||||
|
|
||||||
|
`asyncio.gather` leaves its siblings running; `TaskGroup` wraps the failure in
|
||||||
|
an `ExceptionGroup`.
|
||||||
|
"""
|
||||||
|
tasks = [asyncio.ensure_future(awaitable) for awaitable in awaitables]
|
||||||
|
try:
|
||||||
|
return await asyncio.gather(*tasks)
|
||||||
|
except BaseException:
|
||||||
|
for task in tasks:
|
||||||
|
task.cancel()
|
||||||
|
await asyncio.gather(*tasks, return_exceptions=True)
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
def parse_model_option(value: str) -> "ModelConfig":
|
def parse_model_option(value: str) -> "ModelConfig":
|
||||||
"""Parse a 'provider:name' string into a ModelConfig."""
|
"""Parse a 'provider:name' string into a ModelConfig."""
|
||||||
from haiku.rag.config.models import ModelConfig
|
from haiku.rag.config.models import ModelConfig
|
||||||
|
|
|
||||||
|
|
@ -7,10 +7,11 @@ from docling_core.types.doc.document import DoclingDocument
|
||||||
from docling_core.types.doc.labels import DocItemLabel
|
from docling_core.types.doc.labels import DocItemLabel
|
||||||
|
|
||||||
from haiku.rag.client import HaikuRAG
|
from haiku.rag.client import HaikuRAG
|
||||||
from haiku.rag.client.session import FederatedSession
|
from haiku.rag.client.session import FederatedSession, SingleDatabaseSession
|
||||||
from haiku.rag.config import get_config
|
from haiku.rag.config import get_config
|
||||||
from haiku.rag.store.exceptions import (
|
from haiku.rag.store.exceptions import (
|
||||||
AmbiguousDatabaseError,
|
AmbiguousDatabaseError,
|
||||||
|
ReadOnlyError,
|
||||||
SourceUnavailableError,
|
SourceUnavailableError,
|
||||||
)
|
)
|
||||||
from haiku.rag.store.models import Chunk
|
from haiku.rag.store.models import Chunk
|
||||||
|
|
@ -93,6 +94,38 @@ class TestOpeningDatabases:
|
||||||
|
|
||||||
assert not alpha.store.db.is_open()
|
assert not alpha.store.db.is_open()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_failing_read_leaves_no_sibling_reading(
|
||||||
|
self, tmp_path, monkeypatch
|
||||||
|
):
|
||||||
|
"""Unwinding through `async with` closes every session, so a sibling still
|
||||||
|
reading reads through a closed one."""
|
||||||
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
|
await _seed(config, "beta", ["beta document about cats"])
|
||||||
|
|
||||||
|
reading = asyncio.Event()
|
||||||
|
unwound = asyncio.Event()
|
||||||
|
|
||||||
|
async def listing(self, *args, **kwargs):
|
||||||
|
if self.source == "beta":
|
||||||
|
await reading.wait()
|
||||||
|
raise ReadOnlyError("beta is read-only")
|
||||||
|
reading.set()
|
||||||
|
try:
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
finally:
|
||||||
|
unwound.set()
|
||||||
|
return []
|
||||||
|
|
||||||
|
monkeypatch.setattr(SingleDatabaseSession, "list_documents", listing)
|
||||||
|
|
||||||
|
async with HaikuRAG(config=config) as rag:
|
||||||
|
with pytest.raises(ReadOnlyError, match="beta"):
|
||||||
|
await rag.list_documents()
|
||||||
|
|
||||||
|
assert unwound.is_set()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_a_database_named_twice_is_opened_once(self, tmp_path):
|
async def test_a_database_named_twice_is_opened_once(self, tmp_path):
|
||||||
"""Fusion would count a repeated database as two rank lists."""
|
"""Fusion would count a repeated database as two rank lists."""
|
||||||
|
|
|
||||||
|
|
@ -1,3 +1,4 @@
|
||||||
|
import asyncio
|
||||||
import importlib.util
|
import importlib.util
|
||||||
from unittest.mock import AsyncMock
|
from unittest.mock import AsyncMock
|
||||||
|
|
||||||
|
|
@ -7,7 +8,8 @@ from pydantic_ai.models.openai import OpenAIChatModel
|
||||||
from haiku.rag.config import get_config
|
from haiku.rag.config import get_config
|
||||||
from haiku.rag.config.models import ModelConfig
|
from haiku.rag.config.models import ModelConfig
|
||||||
from haiku.rag.converters import get_converter
|
from haiku.rag.converters import get_converter
|
||||||
from haiku.rag.utils import get_model
|
from haiku.rag.store.exceptions import ReadOnlyError
|
||||||
|
from haiku.rag.utils import gather_all, get_model
|
||||||
|
|
||||||
# Check for optional dependencies
|
# Check for optional dependencies
|
||||||
HAS_ANTHROPIC = importlib.util.find_spec("anthropic") is not None
|
HAS_ANTHROPIC = importlib.util.find_spec("anthropic") is not None
|
||||||
|
|
@ -1104,3 +1106,57 @@ def test_get_model_api_key_rejected_on_unplumbed_provider():
|
||||||
get_model(
|
get_model(
|
||||||
ModelConfig(provider="anthropic", name="claude-sonnet-4-5", api_key="sk-x")
|
ModelConfig(provider="anthropic", name="claude-sonnet-4-5", api_key="sk-x")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestGatherAll:
|
||||||
|
"""A fan-out leaves nothing running: a caller unwinding from a failure closes
|
||||||
|
the sessions its siblings are still reading through."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_results_arrive_in_the_order_asked_for(self):
|
||||||
|
async def slow(value):
|
||||||
|
await asyncio.sleep(0.01)
|
||||||
|
return value
|
||||||
|
|
||||||
|
async def fast(value):
|
||||||
|
return value
|
||||||
|
|
||||||
|
assert await gather_all(slow("a"), fast("b"), slow("c")) == ["a", "b", "c"]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_failure_leaves_no_sibling_running(self):
|
||||||
|
started = asyncio.Event()
|
||||||
|
unwound = asyncio.Event()
|
||||||
|
|
||||||
|
async def sibling():
|
||||||
|
started.set()
|
||||||
|
try:
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
finally:
|
||||||
|
unwound.set()
|
||||||
|
|
||||||
|
async def failing():
|
||||||
|
await started.wait()
|
||||||
|
raise ReadOnlyError("beta is read-only")
|
||||||
|
|
||||||
|
before = asyncio.all_tasks()
|
||||||
|
with pytest.raises(ReadOnlyError, match="beta"):
|
||||||
|
await gather_all(sibling(), failing())
|
||||||
|
|
||||||
|
assert unwound.is_set()
|
||||||
|
assert asyncio.all_tasks() - before == set()
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_the_failure_arrives_as_itself(self):
|
||||||
|
"""A `TaskGroup` drains the siblings too, but raises an `ExceptionGroup`."""
|
||||||
|
|
||||||
|
async def failing():
|
||||||
|
raise ReadOnlyError("beta is read-only")
|
||||||
|
|
||||||
|
async def sibling():
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
|
||||||
|
with pytest.raises(ReadOnlyError) as raised:
|
||||||
|
await gather_all(sibling(), failing())
|
||||||
|
|
||||||
|
assert not isinstance(raised.value, BaseExceptionGroup)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue