495 lines
16 KiB
Python
495 lines
16 KiB
Python
"""
|
|
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"])
|