310 lines
12 KiB
Python
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"])
|