254 lines
9.1 KiB
Python
254 lines
9.1 KiB
Python
"""
|
|
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"])
|