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

495 lines
16 KiB
Python
Raw Normal View History

"""
Comprehensive tests for ingestion connectors.
"""
import pytest
import asyncio
from unittest.mock import AsyncMock, MagicMock, patch, mock_open
from datetime import datetime
from sentiment_engine.ingestion.base import BaseConnector
from sentiment_engine.ingestion.rss import RSSConnector
from sentiment_engine.ingestion.api import APIConnector
from sentiment_engine.ingestion.reddit import RedditConnector
from sentiment_engine.ingestion.telegram import TelegramConnector
from sentiment_engine.ingestion.web_crawl import WebCrawlConnector
from sentiment_engine.ingestion.router import IngestionRouter
from sentiment_engine.schemas.config import (
RSSConnectorConfig, APIConnectorConfig, RedditConnectorConfig,
TelegramConnectorConfig, WebCrawlConnectorConfig
)
from sentiment_engine.schemas.payload import NormalizedPayload, SourceType, AssetMention, EngagementMetrics
from sentiment_engine.nlp.credibility import CredibilityScorer
class TestBaseConnector:
"""Tests for BaseConnector"""
@pytest.fixture
def config(self):
from sentiment_engine.schemas.config import ConnectorConfig
return ConnectorConfig(
name="test_connector",
source_type="news",
poll_interval_seconds=60,
timeout_seconds=30,
rate_limit_rps=1.0,
rate_limit_burst=5,
max_concurrent_requests=2
)
@pytest.fixture
def credibility_registry(self):
return {"test_connector": 0.8}
@pytest.fixture
def connector(self, config, credibility_registry):
return BaseConnector(config, credibility_registry)
@pytest.mark.asyncio
async def test_start_sets_running(self, connector):
"""Start should set _running to True"""
await connector.start()
assert connector._running is True
@pytest.mark.asyncio
async def test_stop_sets_not_running(self, connector):
"""Stop should set _running to False"""
await connector.start()
await connector.stop()
assert connector._running is False
@pytest.mark.asyncio
async def test_fetch_not_implemented(self, connector):
"""Base fetch should raise NotImplementedError"""
with pytest.raises(NotImplementedError):
async for _ in connector.fetch():
pass
@pytest.mark.asyncio
async def test_health_check_not_implemented(self, connector):
"""Base health_check should raise NotImplementedError"""
with pytest.raises(NotImplementedError):
await connector.health_check()
def test_stats_initialization(self, connector):
"""Stats should initialize to zero"""
assert connector.stats["total_fetched"] == 0
assert connector.stats["total_errors"] == 0
assert connector.stats["last_fetch"] is None
def test_rate_limiter_tokens(self, connector):
"""Rate limiter should initialize with burst tokens"""
assert connector._rate_limiter._tokens == connector.config.rate_limit_burst
@pytest.mark.asyncio
async def test_backoff_increases_on_error(self, connector):
"""Backoff should increase on consecutive errors"""
initial_backoff = connector._current_backoff
# Simulate error
connector._handle_error(Exception("test"))
assert connector._current_backoff > initial_backoff
@pytest.mark.asyncio
async def test_backoff_resets_on_success(self, connector):
"""Backoff should reset on success"""
connector._current_backoff = 10.0
connector._handle_success()
assert connector._current_backoff == connector.config.backoff_base_seconds
class TestRSSConnector:
"""Tests for RSSConnector"""
@pytest.fixture
def config(self):
return RSSConnectorConfig(
name="rss_test",
source_type="news",
feed_urls=["https://example.com/rss"],
max_items_per_feed=10,
metadata={"user_agent": "test-agent"}
)
@pytest.fixture
def credibility_registry(self):
return {"rss_test": 0.8}
@pytest.fixture
def connector(self, config, credibility_registry):
return RSSConnector(config, credibility_registry)
@pytest.mark.asyncio
async def test_fetch_parses_rss(self, connector):
"""Should parse RSS feed"""
mock_rss = """<?xml version="1.0"?>
<rss version="2.0">
<channel>
<item>
<title>Test Title</title>
<link>https://example.com/item1</link>
<description>Test description</description>
<pubDate>Mon, 01 Jan 2024 12:00:00 GMT</pubDate>
</item>
</channel>
</rss>"""
with patch('aiohttp.ClientSession.get') as mock_get:
mock_response = AsyncMock()
mock_response.status = 200
mock_response.text = AsyncMock(return_value=mock_rss)
mock_get.return_value.__aenter__.return_value = mock_response
connector._session = AsyncMock()
connector._session.get = mock_get
items = []
async for item in connector.fetch():
items.append(item)
assert len(items) == 1
assert items[0].title == "Test Title"
def test_extract_source_id(self, connector):
"""Should extract source ID from feed URL"""
source_id = connector._extract_source_id("https://www.coindesk.com/rss")
assert source_id == "rss:coindesk.com"
def test_extract_source_id_no_www(self, connector):
"""Should handle URLs without www"""
source_id = connector._extract_source_id("https://coindesk.com/rss")
assert source_id == "rss:coindesk.com"
@pytest.mark.asyncio
async def test_parse_entry_creates_payload(self, connector):
"""_parse_entry should create NormalizedPayload"""
import feedparser
entry = feedparser.parse("""<item>
<title>Test</title>
<link>https://example.com</link>
<description>Desc</description>
<pubDate>Mon, 01 Jan 2024 12:00:00 GMT</pubDate>
</item>""").entries[0]
payload = await connector._parse_entry("https://example.com/rss", entry)
assert payload is not None
assert isinstance(payload, NormalizedPayload)
assert payload.source_id == "rss:example.com"
@pytest.mark.asyncio
async def test_deduplication(self, connector):
"""Should deduplicate entries by content hash"""
connector._seen_ids.add("abc123")
# Mock entry with same ID
import feedparser
entry = feedparser.parse("""<item>
<title>Test</title>
<link>https://example.com</link>
<description>Desc</description>
</item>""").entries[0]
entry.id = "abc123"
payload = await connector._parse_entry("https://example.com/rss", entry)
assert payload is None
class TestAPIConnector:
"""Tests for APIConnector"""
@pytest.fixture
def config(self):
return APIConnectorConfig(
name="api_test",
source_type="news",
base_url="https://api.example.com",
endpoints=["/v1/news"],
auth_type="bearer",
headers={"Authorization": "Bearer test"}
)
@pytest.fixture
def credibility_registry(self):
return {"api_test": 0.9}
@pytest.fixture
def connector(self, config, credibility_registry):
return APIConnector(config, credibility_registry)
@pytest.mark.asyncio
async def test_fetch_calls_endpoints(self, connector):
"""Should call each endpoint"""
mock_response = {"data": [{"title": "Test", "url": "https://example.com", "content": "Test"}]}
with patch('aiohttp.ClientSession.get') as mock_get:
mock_response_obj = AsyncMock()
mock_response_obj.status = 200
mock_response_obj.json = AsyncMock(return_value=mock_response)
mock_get.return_value.__aenter__.return_value = mock_response_obj
connector._session = AsyncMock()
connector._session.get = mock_get
items = []
async for item in connector.fetch():
items.append(item)
assert len(items) == 1
class TestRedditConnector:
"""Tests for RedditConnector"""
@pytest.fixture
def config(self):
return RedditConnectorConfig(
name="reddit_test",
source_type="social",
client_id="test_id",
client_secret="test_secret",
subreddits=["CryptoCurrency", "Bitcoin"],
use_pushshift=True
)
@pytest.fixture
def credibility_registry(self):
return {"reddit_test": 0.7}
@pytest.fixture
def connector(self, config, credibility_registry):
return RedditConnector(config, credibility_registry)
@pytest.mark.asyncio
async def test_initialize_creates_reddit_client(self, connector):
"""Should initialize Reddit client"""
with patch('asyncpraw.Reddit') as mock_reddit:
mock_reddit.return_value = AsyncMock()
await connector.initialize()
assert connector._reddit is not None
def test_extract_tickers_from_title(self, connector):
"""Should extract tickers from post title"""
text = "BTC and ETH are mooning"
tickers = connector._extract_tickers(text)
assert "BTC" in tickers
assert "ETH" in tickers
class TestTelegramConnector:
"""Tests for TelegramConnector"""
@pytest.fixture
def config(self):
return TelegramConnectorConfig(
name="telegram_test",
source_type="social",
bot_token="test_token",
channel_usernames=["@channel1", "@channel2"]
)
@pytest.fixture
def credibility_registry(self):
return {"telegram_test": 0.7}
@pytest.fixture
def connector(self, config, credibility_registry):
return TelegramConnector(config, credibility_registry)
@pytest.mark.asyncio
async def test_initialize_creates_bot(self, connector):
"""Should initialize bot"""
with patch('aiogram.Bot') as mock_bot:
mock_bot.return_value = AsyncMock()
await connector.initialize()
assert connector._bot is not None
class TestWebCrawlConnector:
"""Tests for WebCrawlConnector"""
@pytest.fixture
def config(self):
return WebCrawlConnectorConfig(
name="web_crawl_test",
source_type="news",
seed_urls=["https://example.com"],
allowed_domains=["example.com"],
max_depth=2
)
@pytest.fixture
def credibility_registry(self):
return {"web_crawl_test": 0.6}
@pytest.fixture
def connector(self, config, credibility_registry):
return WebCrawlConnector(config, credibility_registry)
def test_normalize_url(self, connector):
"""Should normalize URLs"""
url = "https://example.com/path?query=1#fragment"
normalized = connector._normalize_url(url)
assert "fragment" not in normalized
def test_is_allowed_domain(self, connector):
"""Should check allowed domains"""
assert connector._is_allowed_domain("https://example.com/page") is True
assert connector._is_allowed_domain("https://other.com/page") is False
@pytest.mark.asyncio
async def test_fetch_crawls_pages(self, connector):
"""Should crawl pages up to max depth"""
mock_html = """<html><body><a href="/page2">Link</a></body></html>"""
with patch('aiohttp.ClientSession.get') as mock_get:
mock_response = AsyncMock()
mock_response.status = 200
mock_response.text = AsyncMock(return_value=mock_html)
mock_get.return_value.__aenter__.return_value = mock_response
connector._session = AsyncMock()
connector._session.get = mock_get
items = []
async for item in connector.fetch():
items.append(item)
assert len(items) >= 1
class TestIngestionRouter:
"""Tests for IngestionRouter"""
@pytest.fixture
def router(self):
return IngestionRouter()
@pytest.fixture
def mock_connector(self):
connector = AsyncMock()
connector.name = "test_connector"
connector.fetch = AsyncMock()
return connector
@pytest.mark.asyncio
async def test_register_connector(self, router, mock_connector):
"""Should register connector"""
router.register_connector(mock_connector)
assert "test_connector" in router._connectors
@pytest.mark.asyncio
async def test_route_publishes_to_nats(self, router, mock_connector):
"""Should publish payloads to NATS"""
mock_connector.fetch.return_value = AsyncMock()
# Create async generator
async def mock_fetch():
from sentiment_engine.schemas.payload import NormalizedPayload, SourceType
yield NormalizedPayload(
source_id="test",
source_type=SourceType.NEWS,
source_credibility_base=0.8,
ingest_ts=1700000000.0,
content_length=10,
raw_text="Test",
metadata={}
)
mock_connector.fetch.return_value = mock_fetch()
router.register_connector(mock_connector)
with patch('sentiment_engine.ingestion.router.NATSJetStreamPublisher') as mock_publisher:
mock_publisher.return_value.publish = AsyncMock()
await router.route_all()
# Should have attempted to publish
assert True # Basic test
class TestConnectorEdgeCases:
"""Edge case tests for connectors"""
@pytest.fixture
def config(self):
from sentiment_engine.schemas.config import ConnectorConfig
return ConnectorConfig(
name="edge_test",
source_type="news",
poll_interval_seconds=60,
timeout_seconds=30,
rate_limit_rps=1.0,
rate_limit_burst=5
)
@pytest.fixture
def connector(self, config):
from sentiment_engine.nlp.credibility import CredibilityScorer
cred = CredibilityScorer()
cred.load_registry({"edge_test": 0.8})
return BaseConnector(config, cred._source_registry)
@pytest.mark.asyncio
async def test_concurrent_fetches(self, connector):
"""Should handle concurrent fetches with semaphore"""
connector.config.max_concurrent_requests = 2
async def slow_fetch():
await asyncio.sleep(0.1)
return []
connector.fetch = slow_fetch
# Run multiple fetches concurrently
tasks = [connector.fetch() for _ in range(4)]
results = await asyncio.gather(*tasks)
assert len(results) == 4
def test_query_windows(self, connector):
"""Should respect preferred and avoid windows"""
# Set preferred window to current hour
now = datetime.now()
connector.config.preferred_query_windows = [{"start_hour": now.hour, "end_hour": now.hour + 1}]
assert connector._in_preferred_window() is True
# Set avoid window to current hour
connector.config.avoid_query_windows = [{"start_hour": now.hour, "end_hour": now.hour + 1}]
assert connector._in_avoid_window() is True
def test_query_windows_wrap_midnight(self, connector):
"""Should handle windows wrapping midnight"""
connector.config.preferred_query_windows = [{"start_hour": 22, "end_hour": 2}]
# 23:00 should be in window
with patch('datetime.datetime') as mock_datetime:
mock_datetime.utcnow.return_value.hour = 23
assert connector._in_preferred_window() is True
mock_datetime.utcnow.return_value.hour = 1
assert connector._in_preferred_window() is True
mock_datetime.utcnow.return_value.hour = 10
assert connector._in_preferred_window() is False
if __name__ == "__main__":
pytest.main([__file__, "-v"])