""" Comprehensive tests for EventClassifier. """ import pytest import asyncio import numpy as np from unittest.mock import AsyncMock, MagicMock, patch from sentiment_engine.nlp.event_classification import ( EventClassifier, ONNXEventModel, EventType ) from sentiment_engine.schemas.processed import EventClassification class TestEventClassifier: """Tests for EventClassifier""" @pytest.fixture def classifier(self): return EventClassifier() @pytest.mark.asyncio async def test_initialize_loads_model(self, classifier): """Should initialize and load model""" await classifier.initialize() # May use ONNX, PyTorch, or keyword fallback def test_classify_listing_keywords(self, classifier): """Should classify listing events""" text = "Binance will list new token ABC tomorrow" events = classifier._classify_sync(text, ["ABC"]) assert len(events) >= 1 assert events[0].event_type == EventType.LISTING def test_classify_hack_keywords(self, classifier): """Should classify hack events""" text = "Exchange hacked, millions stolen in exploit" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.HACK def test_classify_regulatory_keywords(self, classifier): """Should classify regulatory events""" text = "SEC investigation into crypto exchange" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.REGULATORY def test_classify_upgrade_keywords(self, classifier): """Should classify upgrade events""" text = "Ethereum Dencun upgrade activates Proto-Danksharding" events = classifier._classify_sync(text, ["ETH"]) assert len(events) >= 1 assert events[0].event_type == EventType.UPGRADE def test_classify_partnership_keywords(self, classifier): """Should classify partnership events""" text = "JPMorgan and Coinbase announce strategic partnership" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.PARTNERSHIP def test_classify_earnings_keywords(self, classifier): """Should classify earnings events""" text = "Coinbase Q2 earnings beat estimates. Revenue up 50%" events = classifier._classify_sync(text, ["COIN"]) assert len(events) >= 1 assert events[0].event_type == EventType.EARNINGS def test_classify_macro_keywords(self, classifier): """Should classify macro events""" text = "Fed cuts rates 50bps. Bitcoin rallies on macro pivot" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.MACRO def test_classify_liquidation_keywords(self, classifier): """Should classify liquidation events""" text = "Massive liquidation cascade wipes out $500M in longs" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.LIQUIDATION def test_classify_whale_keywords(self, classifier): """Should classify whale events""" text = "Whale moves 10,000 BTC after 5 years dormancy" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.WHALE def test_classify_manipulation_keywords(self, classifier): """Should classify manipulation events""" text = "Pump and dump scheme detected on low cap token" events = classifier._classify_sync(text, ["SHITCOIN"]) assert len(events) >= 1 assert events[0].event_type == EventType.MANIPULATION def test_classify_delisting_keywords(self, classifier): """Should classify delisting events""" text = "Binance delists privacy coins XMR and ZEC" events = classifier._classify_sync(text, ["XMR", "ZEC"]) assert len(events) >= 1 assert events[0].event_type == EventType.DELISTING def test_classify_governance_keywords(self, classifier): """Should classify governance events""" text = "Arbitrum DAO proposal passes with 95% approval" events = classifier._classify_sync(text, ["ARB"]) assert len(events) >= 1 assert events[0].event_type == EventType.GOVERNANCE def test_multiple_events_detected(self, classifier): """Should detect multiple events in one text""" text = "SEC approves ETF and Binance lists new token" events = classifier._classify_sync(text, ["BTC"]) event_types = [e.event_type for e in events] assert EventType.REGULATORY in event_types assert EventType.LISTING in event_types def test_confidence_calculation(self, classifier): """Confidence should increase with more keyword matches""" text1 = "listing" text2 = "listing listed debut launch trading starts" events1 = classifier._classify_sync(text1, ["BTC"]) events2 = classifier._classify_sync(text2, ["BTC"]) assert events2[0].confidence >= events1[0].confidence def test_severity_estimation(self, classifier): """Severity should be higher for stronger language""" text1 = "hack" text2 = "major hack massive exploit emergency" events1 = classifier._classify_sync(text1, ["BTC"]) events2 = classifier._classify_sync(text2, ["BTC"]) assert events2[0].severity >= events1[0].severity def test_find_involved_assets(self, classifier): """Should find mentioned assets in text""" text = "BTC and ETH both surge on news" assets = ["BTC", "ETH", "SOL"] involved = classifier._find_involved_assets(text, assets, EventType.LISTING) assert "BTC" in involved assert "ETH" in involved assert "SOL" not in involved def test_market_wide_events(self, classifier): """Should assign MARKET for macro/regulatory without specific assets""" text = "Fed cuts rates 50bps" involved = classifier._find_involved_assets(text, [], EventType.MACRO) assert involved == ["MARKET"] class TestONNXEventModel: """Tests for ONNXEventModel wrapper""" def test_init_loads_session(self): """Should load ONNX session""" with patch('onnxruntime.InferenceSession') as mock_session: mock_session.return_value.get_inputs.return_value = [ MagicMock(name="input_ids"), MagicMock(name="attention_mask"), MagicMock(name="token_type_ids") ] mock_session.return_value.get_outputs.return_value = [MagicMock(name="logits")] with patch('transformers.AutoTokenizer.from_pretrained'): model = ONNXEventModel("path", "tokenizer_path") assert model.session is not None def test_predict_returns_probabilities(self): """predict should return probabilities summing to 1""" with patch('onnxruntime.InferenceSession') as mock_session: mock_session.return_value.get_inputs.return_value = [ MagicMock(name="input_ids"), MagicMock(name="attention_mask"), MagicMock(name="token_type_ids") ] mock_session.return_value.get_outputs.return_value = [MagicMock(name="logits")] mock_session.return_value.run.return_value = [np.array([[1.0, 2.0, 0.5] + [0.0]*9])] with patch('transformers.AutoTokenizer.from_pretrained'): model = ONNXEventModel("path", "tokenizer_path") probs = model.predict( np.ones((1, 10), dtype=np.int64), np.ones((1, 10), dtype=np.int64) ) assert len(probs) == 12 assert abs(probs.sum() - 1.0) < 0.001 class TestEventClassifierONNX: """Tests for EventClassifier with ONNX model""" @pytest.fixture def classifier(self): return EventClassifier() @pytest.mark.asyncio async def test_classify_with_onnx(self, classifier): """Should use ONNX model when available""" with patch('onnxruntime.InferenceSession') as mock_session: mock_session.return_value.get_inputs.return_value = [ MagicMock(name="input_ids"), MagicMock(name="attention_mask"), MagicMock(name="token_type_ids") ] mock_session.return_value.get_outputs.return_value = [MagicMock(name="logits")] mock_session.return_value.run.return_value = [np.array([[0.8] + [0.02]*11])] with patch('transformers.AutoTokenizer.from_pretrained') as mock_tokenizer: mock_tokenizer.return_value.return_value = { "input_ids": np.ones((1, 10), dtype=np.int64), "attention_mask": np.ones((1, 10), dtype=np.int64), "token_type_ids": np.zeros((1, 10), dtype=np.int64) } classifier._onnx_model = ONNXEventModel("path", "tokenizer_path") classifier._use_onnx = True events = await classifier.classify("SEC approves ETF", ["BTC"]) assert len(events) >= 1 class TestEventClassifierEdgeCases: """Tests for edge cases in event classification""" @pytest.fixture def classifier(self): return EventClassifier() def test_empty_text(self, classifier): """Should handle empty text""" events = classifier._classify_sync("", ["BTC"]) assert events == [] def test_no_keywords(self, classifier): """Should return empty when no keywords match""" text = "The weather is nice today" events = classifier._classify_sync(text, ["BTC"]) assert events == [] def test_case_insensitive_matching(self, classifier): """Should match keywords case insensitively""" text = "HACK EXPLOIT STOLEN" events = classifier._classify_sync(text, ["BTC"]) assert len(events) >= 1 assert events[0].event_type == EventType.HACK def test_partial_keyword_matching(self, classifier): """Should match partial keywords""" text = "hacking attempt detected" events = classifier._classify_sync(text, ["BTC"]) # "hacking" contains "hack" assert len(events) >= 1 def test_asset_not_in_text(self, classifier): """Should not include assets not in text""" text = "BTC surges" assets = ["BTC", "ETH", "SOL"] involved = classifier._find_involved_assets(text, assets, EventType.LISTING) assert "BTC" in involved assert "ETH" not in involved assert "SOL" not in involved def test_overlapping_keywords(self, classifier): """Should handle overlapping keyword categories""" text = "SEC hack investigation" # both regulatory and hack events = classifier._classify_sync(text, ["BTC"]) event_types = [e.event_type for e in events] # Should detect both or the stronger one assert len(events) >= 1 def test_confidence_threshold(self, classifier): """Should filter low confidence events""" # Single weak keyword match text = "maybe listing soon" events = classifier._classify_sync(text, ["BTC"]) # Confidence might be below threshold for e in events: assert e.confidence >= 0.3 if __name__ == "__main__": pytest.main([__file__, "-v"])