Files
sentiment-engine/sentiment_engine/tests/unit/test_entity_extraction_comprehensive.py

254 lines
9.1 KiB
Python
Raw Normal View History

"""
Comprehensive tests for EntityExtractor and AssetMapper.
"""
import pytest
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch
from sentiment_engine.nlp.entity_extraction import (
EntityExtractor, AssetMapper
)
from sentiment_engine.schemas.processed import EntityExtraction
from sentiment_engine.schemas.payload import AssetMention
class TestAssetMapper:
"""Tests for AssetMapper"""
@pytest.fixture
def mapper(self):
return AssetMapper()
def test_map_known_crypto_tickers(self, mapper):
"""Should map known crypto tickers with high confidence"""
for ticker in ["BTC", "ETH", "SOL", "AVAX", "MATIC", "DOT", "LINK", "UNI", "AAVE", "ARB", "OP", "SUI"]:
asset_id, confidence = mapper.map_ticker(ticker)
assert asset_id == ticker
assert confidence >= 0.9
def test_map_ticker_case_insensitive(self, mapper):
"""Should handle case insensitive tickers"""
asset_id, confidence = mapper.map_ticker("btc")
assert asset_id == "BTC"
assert confidence >= 0.9
def test_map_ticker_with_dollar_prefix(self, mapper):
"""Should handle $ prefix"""
asset_id, confidence = mapper.map_ticker("$BTC")
assert asset_id == "BTC"
assert confidence >= 0.9
def test_map_unknown_ticker(self, mapper):
"""Should return unknown ticker with low confidence"""
asset_id, confidence = mapper.map_ticker("UNKNOWNTICKER")
assert asset_id == "UNKNOWNTICKER"
assert confidence == 0.5
def test_map_known_contracts(self, mapper):
"""Should map known contract addresses"""
# ETH contract
asset_id, confidence, chain = mapper.map_contract("0xC02aaA39b223FE8D0A0e5C4F27eAD9083C756Cc2")
assert asset_id == "ETH"
assert confidence >= 0.99
assert chain == "ethereum"
def test_map_unknown_contract(self, mapper):
"""Should return unknown contract with low confidence"""
asset_id, confidence, chain = mapper.map_contract("0x1234567890123456789012345678901234567890")
assert asset_id == "0x1234567890123456789012345678901234567890"
assert confidence == 0.3
assert chain is None
def test_resolve_aliases(self, mapper):
"""Should resolve known aliases"""
results = mapper.resolve_alias("Vitalik Buterin says ETH will moon")
assert any(r[1] == "ETH" for r in results)
results = mapper.resolve_alias("CZ buys BNB")
assert any(r[1] == "BNB" for r in results)
def test_resolve_aliases_case_insensitive(self, mapper):
"""Should resolve aliases case insensitive"""
results = mapper.resolve_alias("VITALIK BUTERIN")
assert any(r[1] == "ETH" for r in results)
class TestEntityExtractor:
"""Tests for EntityExtractor"""
@pytest.fixture
def extractor(self):
return EntityExtractor(AssetMapper())
@pytest.mark.asyncio
async def test_initialize_loads_spacy(self, extractor):
"""Should initialize spaCy if available"""
await extractor.initialize()
# May or may not load spaCy depending on availability
assert extractor.asset_mapper is not None
def test_extract_tickers_basic(self, extractor):
"""Should extract basic tickers"""
text = "BTC and ETH are pumping"
mentions = extractor.extract_tickers(text)
assert len(mentions) == 2
asset_ids = [m.asset_id for m in mentions]
assert "BTC" in asset_ids
assert "ETH" in asset_ids
def test_extract_tickers_with_dollar(self, extractor):
"""Should extract tickers with $ prefix"""
text = "$BTC $ETH are pumping"
mentions = extractor.extract_tickers(text)
assert len(mentions) == 2
def test_extract_tickers_filters_false_positives(self, extractor):
"""Should filter common false positives"""
text = "THE CEO OF API COMPANY SAYS BTC"
mentions = extractor.extract_tickers(text)
asset_ids = [m.asset_id for m in mentions]
assert "THE" not in asset_ids
assert "CEO" not in asset_ids
assert "API" not in asset_ids
assert "BTC" in asset_ids
def test_extract_tickers_deduplicates(self, extractor):
"""Should deduplicate repeated tickers"""
text = "BTC BTC BTC"
mentions = extractor.extract_tickers(text)
assert len(mentions) == 1
assert mentions[0].asset_id == "BTC"
def test_extract_contracts(self, extractor):
"""Should extract contract addresses"""
text = "Send to 0xC02aaA39b223FE8D0A0e5C4F27eAD9083C756Cc2"
mentions = extractor.extract_contracts(text)
assert len(mentions) == 1
assert mentions[0].asset_id == "ETH"
def test_extract_aliases(self, extractor):
"""Should extract known aliases"""
text = "Vitalik says ETH to the moon"
mentions = extractor.extract_aliases(text)
assert len(mentions) >= 1
assert mentions[0].asset_id == "ETH"
@pytest.mark.asyncio
async def test_extract_all_combines_sources(self, extractor):
"""extract_all should combine all extraction methods"""
await extractor.initialize()
text = "BTC surges. Vitalik buys ETH. Send to 0xC02aaA39b223FE8D0A0e5C4F27eAD9083C756Cc2"
entities = await extractor.extract_all(text)
asset_ids = [e.asset_id for e in entities]
assert "BTC" in asset_ids
assert "ETH" in asset_ids
@pytest.mark.asyncio
async def test_extract_all_deduplicates(self, extractor):
"""extract_all should deduplicate overlapping mentions"""
await extractor.initialize()
text = "BTC BTC BTC"
entities = await extractor.extract_all(text)
btc_entities = [e for e in entities if e.asset_id == "BTC"]
assert len(btc_entities) == 1
@pytest.mark.asyncio
async def test_extract_all_no_overlaps(self, extractor):
"""extract_all should not return overlapping spans"""
await extractor.initialize()
text = "BTC and ETH are different assets"
entities = await extractor.extract_all(text)
spans = [(e.mention_span[0], e.mention_span[1]) for e in entities]
for i, (s1, e1) in enumerate(spans):
for j, (s2, e2) in enumerate(spans):
if i != j:
assert not (s1 < e2 and s2 < e1)
def test_deduplicate_tickers_keeps_highest_confidence(self, extractor):
"""Deduplication should keep highest confidence"""
mentions = [
AssetMention(asset_id="BTC", mention_span=(0, 3), confidence=0.5, source_text="BTC", mention_type="ticker"),
AssetMention(asset_id="BTC", mention_span=(5, 8), confidence=0.9, source_text="BTC", mention_type="ticker"),
]
deduped = extractor._deduplicate_tickers(mentions)
assert len(deduped) == 1
assert deduped[0].confidence == 0.9
def test_mention_span_accuracy(self, extractor):
"""Mention spans should accurately reflect position in text"""
text = "BTC surges to $100k"
mentions = extractor.extract_tickers(text)
for mention in mentions:
start, end = mention.mention_span
assert text[start:end] == mention.source_text
class TestEntityExtractionEdgeCases:
"""Tests for edge cases in entity extraction"""
@pytest.fixture
def extractor(self):
return EntityExtractor(AssetMapper())
def test_empty_text(self, extractor):
"""Should handle empty text"""
mentions = extractor.extract_tickers("")
assert mentions == []
def test_text_without_tickers(self, extractor):
"""Should handle text without tickers"""
text = "The market is moving today"
mentions = extractor.extract_tickers(text)
assert mentions == []
def test_mixed_case_tickers(self, extractor):
"""Should handle mixed case"""
text = "btc eth Btc Eth"
mentions = extractor.extract_tickers(text)
assert len(mentions) == 2
def test_ticker_adjacent_to_punctuation(self, extractor):
"""Should handle tickers adjacent to punctuation"""
text = "BTC, ETH; SOL."
mentions = extractor.extract_tickers(text)
assert len(mentions) == 3
def test_ticker_with_numbers(self, extractor):
"""Should handle tickers with numbers"""
text = "SHIB1000 DOGE2"
mentions = extractor.extract_tickers(text)
# These might not be standard tickers but should be extracted
assert len(mentions) >= 0
def test_very_long_text(self, extractor):
"""Should handle very long text"""
text = "BTC " * 1000
mentions = extractor.extract_tickers(text)
# Should deduplicate
assert len(mentions) == 1
def test_unicode_text(self, extractor):
"""Should handle unicode"""
text = "BTC 🚀 ETH 💎"
mentions = extractor.extract_tickers(text)
assert len(mentions) == 2
if __name__ == "__main__":
pytest.main([__file__, "-v"])