Fix stuck retries in extraction and enable configurable retry limit. Resolves #207

This commit is contained in:
KaifAhmad1
2026-01-25 21:27:03 +05:30
parent 246119f48a
commit bc55dcc57a
4 changed files with 290 additions and 1 deletions
+10
View File
@@ -7,6 +7,13 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
### Added
- **Configurable LLM Retry Logic**:
- Exposed `max_retries` parameter in `NERExtractor`, `RelationExtractor`, `TripletExtractor` and low-level extraction methods (`extract_entities_llm`, `extract_relations_llm`, `extract_triplets_llm`).
- Defaults to 3 retries to prevent infinite loops during JSON validation failures or API timeouts.
- Propagated retry configuration through chunked processing helpers to ensure consistent behavior for long documents.
- Updated `03_Earnings_Call_Analysis.ipynb` to use `max_retries=3` by default.
### Added
- **Bring Your Own Model (BYOM) Support**:
- Enabled full support for custom Hugging Face models in `NERExtractor`, `RelationExtractor`, and `TripletExtractor`.
@@ -24,6 +31,9 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
- Implemented post-processing logic to clean and validate generated triplets.
### Fixed
- **LLM Extraction Stability**:
- Fixed infinite retry loops in `BaseProvider` by strictly enforcing `max_retries` limit during structured output generation.
- Resolved stuck execution in earnings call analysis notebooks when using smaller models (e.g., Llama 3 8B) that frequently produce invalid JSON.
- **Model Parameter Precedence**:
- Fixed issue where configuration defaults took precedence over runtime arguments in Hugging Face extractors. Runtime options now correctly override config values.
- **Import Handling**:
@@ -288,6 +288,7 @@
" llm_model=\"llama-3.1-8b-instant\",\n",
" temperature=0.0,\n",
" api_key=GROQ_API_KEY,\n",
" max_retries=3,\n",
")\n",
"\n",
"ENTITY_TYPES = [\"ORGANIZATION\", \"PERSON\", \"MONEY\", \"PERCENT\", \"DATE\", \"EVENT\"]\n",
+15 -1
View File
@@ -1650,6 +1650,7 @@ def extract_relations_llm(
silent_fail: bool = False,
max_text_length: Optional[int] = None,
structured_output_mode: str = "typed",
max_retries: int = 3,
**kwargs,
) -> List[Relation]:
"""
@@ -1662,6 +1663,7 @@ def extract_relations_llm(
model: LLM model
silent_fail: If True, return empty list on error. If False (default), raise exception.
max_text_length: Maximum text length before auto-chunking. None = provider default.
max_retries: Maximum number of retries for LLM calls (default: 3)
**kwargs: Additional options
"""
# Support llm_model parameter to disambiguate from ML model
@@ -1674,6 +1676,7 @@ def extract_relations_llm(
"model": model,
"max_text_length": max_text_length,
"structured_output_mode": structured_output_mode,
"max_retries": max_retries,
"relation_types": kwargs.get("relation_types"),
# Include entities hash/str in cache key implicitly via **cache_params
"entities_hash": hash(tuple(sorted([e.text for e in entities]))) if entities else 0
@@ -1745,6 +1748,7 @@ def extract_relations_llm(
return _extract_relations_chunked(
text, entities, provider=provider, model=model,
silent_fail=silent_fail, max_text_length=max_text_length,
max_retries=max_retries,
**kwargs
)
@@ -1816,6 +1820,8 @@ Entities found in text: {entities_str}"""
call_kwargs["temperature"] = kwargs["temperature"]
if "verbose" in kwargs:
call_kwargs["verbose"] = kwargs["verbose"]
call_kwargs["max_retries"] = max_retries
result_obj = llm.generate_typed(prompt, schema=RelationsResponse, **call_kwargs)
if verbose_mode:
@@ -1962,6 +1968,7 @@ def _extract_relations_chunked(
silent_fail: bool,
max_text_length: int,
structured_output_mode: str = "typed",
max_retries: int = 3,
**kwargs
) -> List[Relation]:
"""Internal helper to extract relations from long text by chunking."""
@@ -2004,6 +2011,7 @@ def _extract_relations_chunked(
silent_fail=False,
max_text_length=len(chunk.text) + 1,
structured_output_mode=structured_output_mode,
max_retries=max_retries,
**limited_kwargs
)
future_to_chunk[future] = i
@@ -2176,6 +2184,7 @@ def extract_triplets_llm(
silent_fail: bool = False,
max_text_length: Optional[int] = None,
structured_output_mode: str = "typed",
max_retries: int = 3,
**kwargs,
) -> List[Triplet]:
"""
@@ -2189,6 +2198,7 @@ def extract_triplets_llm(
model: LLM model
silent_fail: If True, return empty list on error. If False (default), raise exception.
max_text_length: Maximum text length before auto-chunking. None = provider default.
max_retries: Maximum number of retries for LLM calls (default: 3)
**kwargs: Additional options
"""
# Support llm_model parameter to disambiguate from ML model
@@ -2201,6 +2211,7 @@ def extract_triplets_llm(
"model": model,
"max_text_length": max_text_length,
"structured_output_mode": structured_output_mode,
"max_retries": max_retries,
"triplet_types": kwargs.get("triplet_types"),
# Include entities/relations hash in cache key implicitly via **cache_params
"entities_hash": hash(tuple(sorted([e.text for e in entities]))) if entities else 0,
@@ -2266,6 +2277,7 @@ def extract_triplets_llm(
return _extract_triplets_chunked(
text, provider=provider, model=model,
silent_fail=silent_fail, max_text_length=max_text_length,
max_retries=max_retries,
**kwargs
)
@@ -2308,7 +2320,9 @@ Text to extract from:
try:
# Use typed generation with Pydantic schema
result_obj = llm.generate_typed(prompt, schema=TripletsResponse, **kwargs)
call_kwargs = kwargs.copy()
call_kwargs["max_retries"] = max_retries
result_obj = llm.generate_typed(prompt, schema=TripletsResponse, **call_kwargs)
# Convert back to internal Triplet format
triplets = []
+264
View File
@@ -0,0 +1,264 @@
import unittest
from unittest.mock import MagicMock, patch, ANY
import sys
import os
from typing import List, Optional
from pydantic import BaseModel
import importlib.util
# Add project root to path
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../..")))
# Mock external dependencies with __spec__ for importlib checks
mock_spacy = MagicMock()
mock_spacy.__spec__ = MagicMock()
sys.modules["spacy"] = mock_spacy
sys.modules["instructor"] = MagicMock()
sys.modules["groq"] = MagicMock()
# Better mock for openai
mock_openai = MagicMock()
mock_openai.__spec__ = MagicMock()
sys.modules["openai"] = mock_openai
# Mock sentence_transformers and transformers to avoid heavy imports and dependency checks
sys.modules["sentence_transformers"] = MagicMock()
mock_transformers = MagicMock()
mock_transformers.__spec__ = MagicMock()
sys.modules["transformers"] = mock_transformers
from semantica.semantic_extract import NERExtractor
from semantica.semantic_extract.methods import extract_entities_llm, _extract_entities_chunked, extract_relations_llm, extract_triplets_llm
from semantica.semantic_extract.providers import BaseProvider
class EntitiesResponse(BaseModel):
entities: List[dict]
class TestRetryLogic(unittest.TestCase):
def setUp(self):
self.mock_provider = MagicMock()
self.mock_provider.is_available.return_value = True
self.mock_provider.generate_typed.return_value = MagicMock(entities=[])
def test_ner_extractor_init_default(self):
"""Test default max_retries in NERExtractor"""
ner = NERExtractor(method="llm", provider="test")
# Check internal config, max_retries not in config means default behavior downstream
self.assertIsNone(ner.config.get("max_retries"))
def test_ner_extractor_init_custom(self):
"""Test custom max_retries in NERExtractor init"""
ner = NERExtractor(method="llm", provider="test", max_retries=5)
self.assertEqual(ner.config.get("max_retries"), 5)
@patch('semantica.semantic_extract.methods.create_provider')
def test_extract_entities_uses_init_value(self, mock_create_provider):
"""Test extract_entities uses initialized max_retries"""
mock_create_provider.return_value = self.mock_provider
ner = NERExtractor(method="llm", provider="test", max_retries=5)
ner.extract_entities("test text")
# Verify generate_typed called with max_retries=5
args, kwargs = self.mock_provider.generate_typed.call_args
self.assertEqual(kwargs.get("max_retries"), 5)
@patch('semantica.semantic_extract.methods.create_provider')
def test_extract_entities_override(self, mock_create_provider):
"""Test extract_entities override max_retries"""
mock_create_provider.return_value = self.mock_provider
ner = NERExtractor(method="llm", provider="test", max_retries=5)
# Override with 1
ner.extract_entities("test text", max_retries=1)
args, kwargs = self.mock_provider.generate_typed.call_args
self.assertEqual(kwargs.get("max_retries"), 1)
@patch('semantica.semantic_extract.methods.create_provider')
def test_chunked_extraction_propagation(self, mock_create_provider):
"""Test max_retries propagation in chunked extraction"""
mock_create_provider.return_value = self.mock_provider
# Patch TextSplitter where it lives
with patch('semantica.split.TextSplitter') as MockSplitter:
mock_splitter_instance = MockSplitter.return_value
# Mock split to return 2 chunks
mock_chunk1 = MagicMock()
mock_chunk1.text = "chunk1"
mock_chunk2 = MagicMock()
mock_chunk2.text = "chunk2"
mock_splitter_instance.split.return_value = [mock_chunk1, mock_chunk2]
# Force chunking by setting max_text_length small
extract_entities_llm(
"very long text",
provider="test",
model="test-model",
max_text_length=10, # Force chunking
max_retries=7,
structured_output_mode="typed"
)
# Check if generate_typed was called with max_retries=7 for chunks
# It should be called twice (once for each chunk)
self.assertEqual(self.mock_provider.generate_typed.call_count, 2)
# Check arguments of the calls
call_args_list = self.mock_provider.generate_typed.call_args_list
for args, kwargs in call_args_list:
self.assertEqual(kwargs.get("max_retries"), 7)
def test_provider_base_logic(self):
"""Test BaseProvider logic for max_retries with manual loop"""
provider = BaseProvider()
provider.client = MagicMock()
provider.logger = MagicMock()
provider.generate_structured = MagicMock(side_effect=Exception("Fail"))
# Mock instructor failing
with patch('semantica.semantic_extract.providers.instructor') as mock_instructor:
# Make instructor client fail
mock_client = MagicMock()
mock_client.chat.completions.create.side_effect = Exception("Instructor Fail")
mock_instructor.from_provider.return_value = mock_client
mock_instructor.from_openai.return_value = mock_client
# Run with max_retries=2
try:
provider.generate_typed("prompt", EntitiesResponse, max_retries=2)
except Exception:
pass
# Should try manual generation exactly 2 times
self.assertEqual(provider.generate_structured.call_count, 2)
def test_provider_zero_retries(self):
"""Test BaseProvider with max_retries=0"""
provider = BaseProvider()
provider.client = MagicMock()
provider.logger = MagicMock()
provider.generate_structured = MagicMock(side_effect=Exception("Fail"))
# Mock instructor failing
with patch('semantica.semantic_extract.providers.instructor') as mock_instructor:
mock_client = MagicMock()
mock_client.chat.completions.create.side_effect = Exception("Instructor Fail")
mock_instructor.from_provider.return_value = mock_client
mock_instructor.from_openai.return_value = mock_client
try:
provider.generate_typed("prompt", EntitiesResponse, max_retries=0)
except Exception:
pass
# Should NOT try manual generation loop (range(0) is empty)
self.assertEqual(provider.generate_structured.call_count, 0)
@patch('semantica.semantic_extract.methods.create_provider')
def test_relations_retry_propagation(self, mock_create_provider):
"""Test max_retries propagation in relation extraction"""
mock_create_provider.return_value = self.mock_provider
# Create a mock entity
mock_entity = MagicMock()
mock_entity.text = "entity"
mock_entity.start_char = 0
mock_entity.end_char = 5
extract_relations_llm(
"test text",
entities=[mock_entity],
provider="test",
max_retries=4
)
args, kwargs = self.mock_provider.generate_typed.call_args
self.assertEqual(kwargs.get("max_retries"), 4)
@patch('semantica.semantic_extract.methods.create_provider')
def test_relations_chunked_propagation(self, mock_create_provider):
"""Test max_retries propagation in chunked relation extraction"""
mock_create_provider.return_value = self.mock_provider
with patch('semantica.split.TextSplitter') as MockSplitter:
mock_splitter_instance = MockSplitter.return_value
mock_chunk1 = MagicMock()
mock_chunk1.text = "chunk1"
mock_chunk1.start_index = 0
mock_chunk1.end_index = 6
mock_splitter_instance.split.return_value = [mock_chunk1]
# Create a mock entity
mock_entity = MagicMock()
mock_entity.text = "entity"
mock_entity.start_char = 0
mock_entity.end_char = 5
extract_relations_llm(
"very long text",
entities=[mock_entity],
provider="test",
max_text_length=10,
max_retries=6
)
# Check call count - should be called for the chunk
# Note: _extract_relations_chunked creates a new future for each chunk
# which calls extract_relations_llm, which calls generate_typed
self.assertEqual(self.mock_provider.generate_typed.call_count, 1)
args, kwargs = self.mock_provider.generate_typed.call_args
self.assertEqual(kwargs.get("max_retries"), 6)
@patch('semantica.semantic_extract.methods.create_provider')
def test_triplets_retry_propagation(self, mock_create_provider):
"""Test max_retries propagation in triplet extraction"""
mock_create_provider.return_value = self.mock_provider
extract_triplets_llm(
"test text",
entities=[],
relations=[],
provider="test",
max_retries=7
)
args, kwargs = self.mock_provider.generate_typed.call_args
self.assertEqual(kwargs.get("max_retries"), 7)
@patch('semantica.semantic_extract.methods.create_provider')
def test_triplets_chunked_propagation(self, mock_create_provider):
"""Test max_retries propagation in chunked triplet extraction"""
mock_create_provider.return_value = self.mock_provider
with patch('semantica.split.TextSplitter') as MockSplitter:
mock_splitter_instance = MockSplitter.return_value
mock_chunk1 = MagicMock()
mock_chunk1.text = "chunk1"
mock_chunk1.start_index = 0
mock_chunk1.end_index = 6
mock_splitter_instance.split.return_value = [mock_chunk1]
# Use max_text_length > 100 to pass the minimum viable chunk size check
# and make text longer than that
extract_triplets_llm(
"very long text " * 20, # length > 101
entities=[],
relations=[],
provider="test",
max_text_length=101,
max_retries=8
)
# Check call count
self.assertEqual(self.mock_provider.generate_typed.call_count, 1)
args, kwargs = self.mock_provider.generate_typed.call_args
self.assertEqual(kwargs.get("max_retries"), 8)
if __name__ == "__main__":
unittest.main()