Track all in-flight background vacuum tasks, not just the last one
This commit is contained in:
parent
57385506c1
commit
b6ea07d2de
3 changed files with 152 additions and 15 deletions
|
|
@ -91,7 +91,7 @@ class HaikuRAG:
|
||||||
self._create = create
|
self._create = create
|
||||||
self._read_only = read_only
|
self._read_only = read_only
|
||||||
self._before = before
|
self._before = before
|
||||||
self._vacuum_task: asyncio.Task | None = None
|
self._vacuum_tasks: set[asyncio.Task] = set()
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def is_read_only(self) -> bool:
|
def is_read_only(self) -> bool:
|
||||||
|
|
@ -116,12 +116,26 @@ class HaikuRAG:
|
||||||
|
|
||||||
async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002
|
async def __aexit__(self, exc_type, exc_val, exc_tb): # noqa: ARG002
|
||||||
"""Async context manager exit."""
|
"""Async context manager exit."""
|
||||||
# Wait for any pending background vacuum to complete before closing
|
await self._await_vacuum_tasks()
|
||||||
if self._vacuum_task is not None and not self._vacuum_task.done():
|
|
||||||
await self._vacuum_task
|
|
||||||
self.close()
|
self.close()
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
async def _await_vacuum_tasks(self) -> None:
|
||||||
|
"""Wait for all in-flight background vacuum tasks to complete.
|
||||||
|
|
||||||
|
Each create_document / update_document can schedule its own vacuum task;
|
||||||
|
all must be awaited before tearing down the connection, not just the
|
||||||
|
most recently scheduled one.
|
||||||
|
"""
|
||||||
|
if self._vacuum_tasks:
|
||||||
|
await asyncio.gather(*self._vacuum_tasks, return_exceptions=True)
|
||||||
|
|
||||||
|
def _schedule_vacuum(self) -> None:
|
||||||
|
"""Schedule a background vacuum and track the task for later awaiting."""
|
||||||
|
task = asyncio.create_task(self.store.vacuum())
|
||||||
|
self._vacuum_tasks.add(task)
|
||||||
|
task.add_done_callback(self._vacuum_tasks.discard)
|
||||||
|
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
# Processing Primitives
|
# Processing Primitives
|
||||||
# =========================================================================
|
# =========================================================================
|
||||||
|
|
@ -376,8 +390,6 @@ class HaikuRAG:
|
||||||
Returns:
|
Returns:
|
||||||
The created Document instance with ID set.
|
The created Document instance with ID set.
|
||||||
"""
|
"""
|
||||||
import asyncio
|
|
||||||
|
|
||||||
# Ensure all chunks have embeddings before storing
|
# Ensure all chunks have embeddings before storing
|
||||||
chunks = await self._ensure_chunks_embedded(chunks)
|
chunks = await self._ensure_chunks_embedded(chunks)
|
||||||
|
|
||||||
|
|
@ -405,7 +417,7 @@ class HaikuRAG:
|
||||||
|
|
||||||
# Vacuum old versions in background (non-blocking) if auto_vacuum enabled
|
# Vacuum old versions in background (non-blocking) if auto_vacuum enabled
|
||||||
if self._config.storage.auto_vacuum:
|
if self._config.storage.auto_vacuum:
|
||||||
self._vacuum_task = asyncio.create_task(self.store.vacuum())
|
self._schedule_vacuum()
|
||||||
|
|
||||||
return created_doc
|
return created_doc
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|
@ -432,8 +444,6 @@ class HaikuRAG:
|
||||||
Returns:
|
Returns:
|
||||||
The updated Document instance.
|
The updated Document instance.
|
||||||
"""
|
"""
|
||||||
import asyncio
|
|
||||||
|
|
||||||
assert document.id is not None, "Document ID is required for update"
|
assert document.id is not None, "Document ID is required for update"
|
||||||
|
|
||||||
# Ensure all chunks have embeddings before storing
|
# Ensure all chunks have embeddings before storing
|
||||||
|
|
@ -468,7 +478,7 @@ class HaikuRAG:
|
||||||
|
|
||||||
# Vacuum old versions in background (non-blocking) if auto_vacuum enabled
|
# Vacuum old versions in background (non-blocking) if auto_vacuum enabled
|
||||||
if self._config.storage.auto_vacuum:
|
if self._config.storage.auto_vacuum:
|
||||||
self._vacuum_task = asyncio.create_task(self.store.vacuum())
|
self._schedule_vacuum()
|
||||||
|
|
||||||
return updated_doc
|
return updated_doc
|
||||||
except Exception:
|
except Exception:
|
||||||
|
|
@ -1403,8 +1413,7 @@ class HaikuRAG:
|
||||||
The ID of the document currently being processed.
|
The ID of the document currently being processed.
|
||||||
"""
|
"""
|
||||||
# Wait for any background vacuum before destructive table operations
|
# Wait for any background vacuum before destructive table operations
|
||||||
if self._vacuum_task is not None and not self._vacuum_task.done():
|
await self._await_vacuum_tasks()
|
||||||
await self._vacuum_task
|
|
||||||
|
|
||||||
# Update settings to current config
|
# Update settings to current config
|
||||||
settings_repo = SettingsRepository(self.store)
|
settings_repo = SettingsRepository(self.store)
|
||||||
|
|
|
||||||
File diff suppressed because one or more lines are too long
|
|
@ -101,9 +101,8 @@ async def test_existing_database_checks_migrations(monkeypatch, temp_db_path):
|
||||||
|
|
||||||
|
|
||||||
async def _wait_for_background_vacuum(client):
|
async def _wait_for_background_vacuum(client):
|
||||||
"""Wait for any background vacuum task to complete."""
|
"""Wait for any in-flight background vacuum tasks to complete."""
|
||||||
if client._vacuum_task is not None and not client._vacuum_task.done():
|
await client._await_vacuum_tasks()
|
||||||
await client._vacuum_task
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.vcr()
|
@pytest.mark.vcr()
|
||||||
|
|
@ -226,6 +225,53 @@ async def test_aexit_awaits_background_vacuum(temp_db_path, monkeypatch):
|
||||||
assert vacuum_completed.is_set(), "__aexit__ exited before vacuum finished"
|
assert vacuum_completed.is_set(), "__aexit__ exited before vacuum finished"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.vcr()
|
||||||
|
async def test_aexit_awaits_all_background_vacuums(temp_db_path, monkeypatch):
|
||||||
|
"""Multiple create_document calls schedule multiple vacuum tasks; __aexit__
|
||||||
|
must await all of them, not just the last-scheduled one.
|
||||||
|
|
||||||
|
Scenario: Task A acquires the vacuum lock and is slow. Task B is scheduled
|
||||||
|
while Task A still holds the lock — Task B sees the lock held and returns
|
||||||
|
immediately. If the client only tracks the most recently scheduled task,
|
||||||
|
__aexit__ awaits the fast no-op B and closes the connection while Task A
|
||||||
|
is still running.
|
||||||
|
"""
|
||||||
|
from haiku.rag.config import Config
|
||||||
|
|
||||||
|
monkeypatch.setattr(Config.storage, "auto_vacuum", True)
|
||||||
|
|
||||||
|
first_vacuum_completed = asyncio.Event()
|
||||||
|
|
||||||
|
async with HaikuRAG(db_path=temp_db_path, create=True) as client:
|
||||||
|
call_count = 0
|
||||||
|
|
||||||
|
async def slow_vacuum(*_args, **_kwargs):
|
||||||
|
nonlocal call_count
|
||||||
|
call_count += 1
|
||||||
|
my_num = call_count
|
||||||
|
# Mimic the real vacuum's skip-if-running behavior.
|
||||||
|
if client.store._vacuum_lock.locked():
|
||||||
|
return
|
||||||
|
async with client.store._vacuum_lock:
|
||||||
|
if my_num == 1:
|
||||||
|
# Hold the lock longer than any other operation in the
|
||||||
|
# test so Task A cannot finish incidentally. __aexit__
|
||||||
|
# must explicitly wait for this task.
|
||||||
|
await asyncio.sleep(2.0)
|
||||||
|
first_vacuum_completed.set()
|
||||||
|
|
||||||
|
client.store.vacuum = slow_vacuum
|
||||||
|
|
||||||
|
await client.create_document(content="triggers first vacuum")
|
||||||
|
# Let Task A start and acquire the vacuum lock before scheduling B.
|
||||||
|
await asyncio.sleep(0.02)
|
||||||
|
await client.create_document(content="triggers second vacuum")
|
||||||
|
|
||||||
|
assert first_vacuum_completed.is_set(), (
|
||||||
|
"__aexit__ returned before the first vacuum task finished"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.vcr()
|
@pytest.mark.vcr()
|
||||||
async def test_auto_vacuum_disabled_skips_vacuum(temp_db_path, monkeypatch):
|
async def test_auto_vacuum_disabled_skips_vacuum(temp_db_path, monkeypatch):
|
||||||
"""Test that auto_vacuum=False prevents automatic vacuum after operations."""
|
"""Test that auto_vacuum=False prevents automatic vacuum after operations."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue