Keep every database a cancelled fan-out opened
`sessions_for` recorded what opened from `gather`'s results, which arrive only when it runs to completion. Cancelling it discarded them, so a database that opened while a sibling was still pending was never recorded and `aclose` never closed it. `return_exceptions=True` covers a failing child, not a cancelled parent. `_open` now registers its own session, so what opened is reachable however the fan-out ends.
This commit is contained in:
parent
b2c617f9f2
commit
cfb880e8a0
2 changed files with 46 additions and 17 deletions
|
|
@ -288,25 +288,22 @@ class FederatedSession:
|
||||||
missing = [name for name in names if name not in self._sessions]
|
missing = [name for name in names if name not in self._sessions]
|
||||||
if missing:
|
if missing:
|
||||||
opened = await asyncio.gather(
|
opened = await asyncio.gather(
|
||||||
*(self._open(self._refs[name]) for name in missing),
|
*(self._open(name) for name in missing),
|
||||||
return_exceptions=True,
|
return_exceptions=True,
|
||||||
)
|
)
|
||||||
# Whatever opened is tracked before the failure is reported, so
|
for result in opened:
|
||||||
# teardown closes it: `gather` does not cancel the siblings of the
|
|
||||||
# one that raised, and an untracked connection leaks.
|
|
||||||
failure: BaseException | None = None
|
|
||||||
for name, result in zip(missing, opened, strict=True):
|
|
||||||
if isinstance(result, BaseException):
|
if isinstance(result, BaseException):
|
||||||
failure = failure or result
|
raise result
|
||||||
else:
|
|
||||||
self._sessions[name] = result
|
|
||||||
if failure is not None:
|
|
||||||
raise failure
|
|
||||||
return [self._sessions[name] for name in names]
|
return [self._sessions[name] for name in names]
|
||||||
|
|
||||||
async def _open(self, ref: DatabaseRef) -> SingleDatabaseSession:
|
async def _open(self, name: str) -> None:
|
||||||
|
"""Open and register one database before returning to the fan-out.
|
||||||
|
|
||||||
|
Registered here because a cancelled `gather` discards its results.
|
||||||
|
"""
|
||||||
|
ref = self._refs[name]
|
||||||
one, db_path = ref.connection(self._config)
|
one, db_path = ref.connection(self._config)
|
||||||
return await SingleDatabaseSession(
|
self._sessions[name] = await SingleDatabaseSession(
|
||||||
db_path if db_path is not None else default_db_path(one),
|
db_path if db_path is not None else default_db_path(one),
|
||||||
one,
|
one,
|
||||||
skip_validation=self._skip_validation,
|
skip_validation=self._skip_validation,
|
||||||
|
|
|
||||||
|
|
@ -35,11 +35,11 @@ class TestOpeningDatabases:
|
||||||
barrier = asyncio.Barrier(len(names))
|
barrier = asyncio.Barrier(len(names))
|
||||||
open_one = rag._session._open
|
open_one = rag._session._open
|
||||||
|
|
||||||
async def gated(ref):
|
async def gated(name):
|
||||||
# Every open has to be in flight before any of them finishes, so
|
# Every open has to be in flight before any of them finishes, so
|
||||||
# a serial loop cannot get past this and the wait times out.
|
# a serial loop cannot get past this and the wait times out.
|
||||||
await barrier.wait()
|
await barrier.wait()
|
||||||
return await open_one(ref)
|
await open_one(name)
|
||||||
|
|
||||||
rag._session._open = gated
|
rag._session._open = gated
|
||||||
clients = await asyncio.wait_for(rag.clients_for(names), timeout=15)
|
clients = await asyncio.wait_for(rag.clients_for(names), timeout=15)
|
||||||
|
|
@ -61,6 +61,38 @@ class TestOpeningDatabases:
|
||||||
assert isinstance(rag._session, FederatedSession)
|
assert isinstance(rag._session, FederatedSession)
|
||||||
assert set(rag._session._sessions) == {"alpha"}
|
assert set(rag._session._sessions) == {"alpha"}
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_a_cancelled_open_does_not_leak_the_ones_that_worked(self, tmp_path):
|
||||||
|
"""Cancellation discards the fan-out's results rather than returning them,
|
||||||
|
so a database that opened while a sibling was still pending is reachable
|
||||||
|
only because the opener recorded it."""
|
||||||
|
config = _config(tmp_path, ["alpha", "beta"])
|
||||||
|
await _seed(config, "alpha", ["alpha document about cats"])
|
||||||
|
await _seed(config, "beta", ["beta document about cats"])
|
||||||
|
|
||||||
|
async with HaikuRAG(config=config) as rag:
|
||||||
|
assert isinstance(rag._session, FederatedSession)
|
||||||
|
open_one = rag._session._open
|
||||||
|
alpha_open = asyncio.Event()
|
||||||
|
|
||||||
|
async def staged(name):
|
||||||
|
if name == "beta":
|
||||||
|
await asyncio.sleep(60)
|
||||||
|
await open_one(name)
|
||||||
|
alpha_open.set()
|
||||||
|
|
||||||
|
rag._session._open = staged
|
||||||
|
fanout = asyncio.create_task(rag.clients_for(["alpha", "beta"]))
|
||||||
|
await asyncio.wait_for(alpha_open.wait(), timeout=15)
|
||||||
|
fanout.cancel()
|
||||||
|
with pytest.raises(asyncio.CancelledError):
|
||||||
|
await fanout
|
||||||
|
|
||||||
|
assert set(rag._session._sessions) == {"alpha"}
|
||||||
|
alpha = rag._session._sessions["alpha"]
|
||||||
|
|
||||||
|
assert not alpha.store.db.is_open()
|
||||||
|
|
||||||
@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."""
|
||||||
|
|
@ -392,8 +424,8 @@ class TestDatabaseIndependentWork:
|
||||||
|
|
||||||
opened: list[str] = []
|
opened: list[str] = []
|
||||||
|
|
||||||
async def refuse(self, ref):
|
async def refuse(self, name):
|
||||||
opened.append(ref.name)
|
opened.append(name)
|
||||||
raise AssertionError("opened a database to chunk a document")
|
raise AssertionError("opened a database to chunk a document")
|
||||||
|
|
||||||
monkeypatch.setattr(FederatedSession, "_open", refuse)
|
monkeypatch.setattr(FederatedSession, "_open", refuse)
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue