216 lines
7.4 KiB
Python
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"
|