fix(entity_resolver): prevent _pending_stats/_pending_cooccurrences memory leak (#662)

* fix(entity_resolver): prevent _pending_stats/_pending_cooccurrences memory leak

Add discard_pending_stats() to EntityResolver to clean up both pending dicts
for the current task key. Call it at the start of each _run_db_work attempt so
that exceptions between accumulation and flush_pending_stats() — including
deadlock retries — never leave stale entries keyed by recycled task IDs.

Fixes #660

* test(entity_resolver): add unit tests for discard_pending_stats()

Covers: clears both dicts for current task, is idempotent when empty,
and does not touch entries belonging to other task keys.
No database required — purely in-memory logic.
This commit is contained in:
Nicolò Boschi 2026-03-23 16:06:04 +01:00 committed by GitHub
parent d886d3acb9
commit e6333719ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 61 additions and 0 deletions

View file

@ -86,6 +86,19 @@ class EntityResolver:
task = asyncio.current_task()
return id(task) if task is not None else 0
def discard_pending_stats(self) -> None:
"""
Discard accumulated entity stats and co-occurrence counts for the current task.
Call this on any exception path between resolve_entities_batch /
link_units_to_entities_batch and flush_pending_stats() to prevent the
per-task dicts from growing unbounded when tasks fail before flushing.
Safe to call even if no entries exist for the current task.
"""
key = self._task_key()
self._pending_stats.pop(key, None)
self._pending_cooccurrences.pop(key, None)
async def flush_pending_stats(self) -> None:
"""
Flush accumulated entity stats and co-occurrence counts for the current task.

View file

@ -292,6 +292,10 @@ async def retain_batch(
pf.document_id = None
pf.chunk_id = None
# Discard any leftover pending stats from a previous failed attempt so
# retries don't double-count or accumulate unbounded state.
entity_resolver.discard_pending_stats()
async with acquire_with_retry(pool) as conn:
async with conn.transaction():
# Handle document tracking for all documents

View file

@ -11,6 +11,50 @@ import pytest
from hindsight_api.engine.entity_resolver import EntityResolver
from hindsight_api.pg0 import resolve_database_url
# ---------------------------------------------------------------------------
# Unit tests for discard_pending_stats() — no database required
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_discard_pending_stats_clears_both_dicts():
"""discard_pending_stats() must remove entries for the current task from
both _pending_stats and _pending_cooccurrences."""
resolver = EntityResolver(pool=None) # type: ignore[arg-type]
key = resolver._task_key()
resolver._pending_stats[key] = [object()] # type: ignore[list-item]
resolver._pending_cooccurrences[key] = [object()] # type: ignore[list-item]
resolver.discard_pending_stats()
assert key not in resolver._pending_stats
assert key not in resolver._pending_cooccurrences
@pytest.mark.asyncio
async def test_discard_pending_stats_is_idempotent():
"""Calling discard_pending_stats() when nothing is pending must not raise."""
resolver = EntityResolver(pool=None) # type: ignore[arg-type]
resolver.discard_pending_stats()
resolver.discard_pending_stats() # second call — still safe
@pytest.mark.asyncio
async def test_discard_pending_stats_does_not_affect_other_task_keys():
"""discard_pending_stats() must only remove the current task's entries,
leaving entries keyed under other task IDs untouched."""
resolver = EntityResolver(pool=None) # type: ignore[arg-type]
other_key = -1 # A fake key that can never be a real task id
resolver._pending_stats[other_key] = [object()] # type: ignore[list-item]
resolver._pending_cooccurrences[other_key] = [object()] # type: ignore[list-item]
resolver.discard_pending_stats() # discards current task's key only
assert other_key in resolver._pending_stats, "other task's stats must be preserved"
assert other_key in resolver._pending_cooccurrences, "other task's cooccurrences must be preserved"
@pytest.mark.asyncio
async def test_resolve_entities_batch_handles_unicode_lower_conflicts(pg0_db_url):