reddit-media-collector/tests/test_database.py

216 lines
7.4 KiB
Python

"""Tests for database operations."""
from datetime import datetime
from src.database import PostRecord
def _make_post(post_id="test123", **kwargs):
"""Helper to create a PostRecord with defaults."""
defaults = {
"id": post_id,
"subreddit": "pics",
"author": "testuser",
"title": "Test Post",
"url": "https://reddit.com/test",
"media_url": "https://i.redd.it/test.jpg",
"media_type": "image",
"score": 100,
"created_utc": 1700000000.0,
"downloaded_at": None,
"local_path": None,
"file_hash": None,
"permalink": "/r/pics/comments/test123/test_post/",
"source_type": "subreddit",
"flair": "OC",
}
defaults.update(kwargs)
return PostRecord(**defaults)
class TestDatabase:
def test_init_creates_tables(self, db):
with db._get_connection() as conn:
cursor = conn.execute("SELECT name FROM sqlite_master WHERE type='table'")
tables = {row[0] for row in cursor.fetchall()}
assert "posts" in tables
assert "favorites" in tables
assert "scheduler_history" in tables
def test_add_and_get_post(self, db):
post = _make_post()
db.add_post(post)
retrieved = db.get_post("test123")
assert retrieved is not None
assert retrieved.id == "test123"
assert retrieved.subreddit == "pics"
assert retrieved.author == "testuser"
def test_post_exists(self, db):
assert db.post_exists("test123") is False
db.add_post(_make_post())
assert db.post_exists("test123") is True
def test_mark_downloaded(self, db):
db.add_post(_make_post())
db.mark_downloaded("test123", "/path/to/file.jpg", "abc123hash")
post = db.get_post("test123")
assert post.local_path == "/path/to/file.jpg"
assert post.file_hash == "abc123hash"
assert post.downloaded_at is not None
def test_hash_exists(self, db):
assert db.hash_exists("abc123hash") is None
post = _make_post()
db.add_post(post)
db.mark_downloaded("test123", "/path/to/file.jpg", "abc123hash")
assert db.hash_exists("abc123hash") == "/path/to/file.jpg"
def test_get_stats(self, db):
db.add_post(_make_post("p1"))
db.mark_downloaded("p1", "/path/p1.jpg", "hash1")
db.add_post(_make_post("p2", media_type="video"))
db.mark_downloaded("p2", "/path/p2.mp4", "hash2")
stats = db.get_stats()
assert stats["total_posts"] == 2
assert stats["downloaded"] == 2
def test_get_media_files_sorting(self, db):
db.add_post(_make_post("p1", score=10, created_utc=1000.0))
db.mark_downloaded("p1", "/p1.jpg", "h1")
db.add_post(_make_post("p2", score=500, created_utc=2000.0))
db.mark_downloaded("p2", "/p2.jpg", "h2")
newest = db.get_media_files(sort="newest")
assert newest[0]["id"] == "p2"
oldest = db.get_media_files(sort="oldest")
assert oldest[0]["id"] == "p1"
score_high = db.get_media_files(sort="score_high")
assert score_high[0]["id"] == "p2"
def test_get_media_files_filtering(self, db):
db.add_post(_make_post("p1", subreddit="pics"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
db.add_post(_make_post("p2", subreddit="videos", media_type="video"))
db.mark_downloaded("p2", "/p2.mp4", "h2")
pics_only = db.get_media_files(subreddit="pics")
assert len(pics_only) == 1
assert pics_only[0]["subreddit"] == "pics"
videos_only = db.get_media_files(media_type="video")
assert len(videos_only) == 1
def test_total_media_count(self, db):
db.add_post(_make_post("p1"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
db.add_post(_make_post("p2"))
db.mark_downloaded("p2", "/p2.jpg", "h2")
assert db.get_total_media_count() == 2
assert db.get_total_media_count(subreddit="pics") == 2
assert db.get_total_media_count(subreddit="nonexistent") == 0
class TestFavorites:
def test_add_favorite(self, db):
db.add_post(_make_post())
assert db.add_favorite("test123") is True
assert db.add_favorite("test123") is False # duplicate
def test_remove_favorite(self, db):
db.add_post(_make_post())
db.add_favorite("test123")
assert db.remove_favorite("test123") is True
assert db.remove_favorite("test123") is False
def test_is_favorite(self, db):
db.add_post(_make_post())
assert db.is_favorite("test123") is False
db.add_favorite("test123")
assert db.is_favorite("test123") is True
def test_get_favorites(self, db):
db.add_post(_make_post("p1"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
db.add_post(_make_post("p2"))
db.mark_downloaded("p2", "/p2.jpg", "h2")
db.add_favorite("p1")
favs = db.get_favorites()
assert len(favs) == 1
assert favs[0]["id"] == "p1"
def test_count_favorites(self, db):
db.add_post(_make_post("p1"))
db.add_post(_make_post("p2"))
db.add_favorite("p1")
db.add_favorite("p2")
assert db.count_favorites() == 2
def test_get_favorite_authors(self, db):
db.add_post(_make_post("p1", author="alice"))
db.add_post(_make_post("p2", author="bob"))
db.add_favorite("p1")
authors = db.get_favorite_authors()
assert "alice" in authors
assert "bob" not in authors
class TestSchedulerHistory:
def test_add_and_finish_run(self, db):
run_id = db.add_scheduler_run(datetime.now())
assert run_id > 0
db.finish_scheduler_run(run_id, "success", 10, 5)
history = db.get_scheduler_history()
assert len(history) == 1
assert history[0]["status"] == "success"
assert history[0]["posts_processed"] == 10
def test_get_last_scheduler_run(self, db):
assert db.get_last_scheduler_run() is None
db.add_scheduler_run(datetime.now())
assert db.get_last_scheduler_run() is not None
class TestPostsByAuthorsAndSubreddits:
def test_get_posts_by_authors(self, db):
db.add_post(_make_post("p1", author="alice"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
db.add_post(_make_post("p2", author="bob"))
db.mark_downloaded("p2", "/p2.jpg", "h2")
posts = db.get_posts_by_authors(["alice"])
assert len(posts) == 1
assert posts[0].author == "alice"
def test_case_insensitive_author_search(self, db):
db.add_post(_make_post("p1", author="Alice"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
posts = db.get_posts_by_authors(["alice"])
assert len(posts) == 1
def test_count_posts_by_authors(self, db):
db.add_post(_make_post("p1", author="alice"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
db.add_post(_make_post("p2", author="alice"))
db.mark_downloaded("p2", "/p2.jpg", "h2")
assert db.count_posts_by_authors(["alice"]) == 2
assert db.count_posts_by_authors(["bob"]) == 0
assert db.count_posts_by_authors([]) == 0
def test_get_posts_by_subreddits(self, db):
db.add_post(_make_post("p1", subreddit="pics"))
db.mark_downloaded("p1", "/p1.jpg", "h1")
posts = db.get_posts_by_subreddits(["pics"])
assert len(posts) == 1
assert posts[0].subreddit == "pics"