""" 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"])