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

310 lines
12 KiB
Python
Raw Normal View History

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