fix: reranker crashes on provider error (#403)

* fix: reranker crashes on provider error

* fix: reranker crashes on provider error
This commit is contained in:
Nicolò Boschi 2026-02-19 11:38:37 +01:00 committed by GitHub
parent c3ef1555bf
commit 58c4d65778
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 73 additions and 1 deletions

View file

@ -2628,6 +2628,8 @@ class MemoryEngine(MemoryEngineInterface):
rerank_span.set_attribute("hindsight.bank_id", bank_id)
rerank_span.set_attribute("hindsight.candidates_count", len(merged_candidates))
scored_results: list = []
pre_filtered_count = 0
try:
# Ensure reranker is initialized (for lazy initialization mode)
await reranker_instance.ensure_initialized()
@ -2635,7 +2637,6 @@ class MemoryEngine(MemoryEngineInterface):
# Pre-filter candidates to reduce reranking cost (RRF already provides good ranking)
# This is especially important for remote rerankers with network latency
reranker_max_candidates = get_config().reranker_max_candidates
pre_filtered_count = 0
if len(merged_candidates) > reranker_max_candidates:
# Sort by RRF score and take top candidates
merged_candidates.sort(key=lambda mc: mc.rrf_score, reverse=True)

View file

@ -0,0 +1,71 @@
"""
Regression test for UnboundLocalError in recall when the reranker raises.
Before the fix, `scored_results` and `pre_filtered_count` were only assigned
inside the `try` block, but referenced in the `finally` block. If
`reranker_instance.rerank()` (or `ensure_initialized()`) raised, the `finally`
block crashed with `UnboundLocalError` instead of propagating the original
exception.
Fix: initialise both variables to safe defaults before the try/finally block.
"""
from datetime import datetime, timezone
from unittest.mock import AsyncMock, patch
import pytest
@pytest.mark.asyncio
async def test_recall_reranker_error_does_not_raise_unbound_local(memory, request_context):
"""Recall must propagate the reranker's exception, not an UnboundLocalError."""
bank_id = f"test_reranker_err_{datetime.now(timezone.utc).timestamp()}"
try:
await memory.retain_async(
bank_id=bank_id,
content="Paris is the capital of France",
request_context=request_context,
)
# Simulate a reranker failure (e.g. Cohere API error on empty/small candidate set)
rerank_mock = AsyncMock(side_effect=RuntimeError("reranker API error"))
memory._cross_encoder_reranker._initialized = True # skip ensure_initialized
with patch.object(memory._cross_encoder_reranker, "rerank", rerank_mock):
with pytest.raises(Exception, match="reranker API error"):
await memory.recall_async(
bank_id=bank_id,
query="capital of France",
request_context=request_context,
)
finally:
await memory.delete_bank(bank_id, request_context=request_context)
@pytest.mark.asyncio
async def test_recall_reranker_init_error_does_not_raise_unbound_local(memory, request_context):
"""Same regression when ensure_initialized() raises (before pre_filtered_count is set)."""
bank_id = f"test_reranker_init_err_{datetime.now(timezone.utc).timestamp()}"
try:
await memory.retain_async(
bank_id=bank_id,
content="Paris is the capital of France",
request_context=request_context,
)
init_mock = AsyncMock(side_effect=RuntimeError("reranker init failed"))
memory._cross_encoder_reranker._initialized = False
with patch.object(memory._cross_encoder_reranker, "ensure_initialized", init_mock):
with pytest.raises(Exception, match="reranker init failed"):
await memory.recall_async(
bank_id=bank_id,
query="capital of France",
request_context=request_context,
)
finally:
await memory.delete_bank(bank_id, request_context=request_context)