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

310 lines
12 KiB
Python

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