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