aboutsummaryrefslogtreecommitdiff
path: root/server/szurubooru/tests/conftest.py
diff options
context:
space:
mode:
authorRuin0x112020-05-04 19:12:54 -0700
committerRuin0x112020-05-04 19:12:54 -0700
commit5ca21f9e7f5447a2c7841c1317dbe8df01b04e97 (patch)
treeb807a323d84653c2a1af707436bd002973cc132f /server/szurubooru/tests/conftest.py
parent6b8e3f251f69c9d7d560148100cdb8df941a81c2 (diff)
Add pool tests
Diffstat (limited to 'server/szurubooru/tests/conftest.py')
-rw-r--r--server/szurubooru/tests/conftest.py46
1 files changed, 46 insertions, 0 deletions
diff --git a/server/szurubooru/tests/conftest.py b/server/szurubooru/tests/conftest.py
index a131fec..249df1a 100644
--- a/server/szurubooru/tests/conftest.py
+++ b/server/szurubooru/tests/conftest.py
@@ -202,6 +202,52 @@ def post_favorite_factory(user_factory, post_factory):
@pytest.fixture
+def pool_category_factory():
+ def factory(name=None, color='dummy', default=False):
+ category = model.PoolCategory()
+ category.name = name or get_unique_name()
+ category.color = color
+ category.default = default
+ return category
+ return factory
+
+
+@pytest.fixture
+def pool_factory():
+ def factory(id=None, names=None, description=None, category=None, time=None):
+ if not category:
+ category = model.PoolCategory(get_unique_name())
+ db.session.add(category)
+ pool = model.Pool()
+ pool.pool_id = id
+ pool.names = []
+ for i, name in enumerate(names or [get_unique_name()]):
+ pool.names.append(model.PoolName(name, i))
+ pool.description = description
+ pool.category = category
+ pool.creation_time = time or datetime(1996, 1, 1)
+ return pool
+ return factory
+
+
+@pytest.fixture
+def pool_post_factory(pool_factory, post_factory):
+ def factory(pool=None, post=None, order=None):
+ if not pool:
+ pool = pool_factory()
+ db.session.add(pool)
+ if not post:
+ post = post_factory()
+ db.session.add(post)
+ pool_post = model.PoolPost(post)
+ pool_post.pool = pool
+ pool_post.post = post
+ pool_post.order = order or 0
+ return pool_post
+ return factory
+
+
+@pytest.fixture
def read_asset():
def get(path):
path = os.path.join(os.path.dirname(__file__), 'assets', path)

© 2015 - 2026 Jakob L. Kreuze