diff options
| -rw-r--r-- | server/szurubooru/api/metric_api.py | 18 | ||||
| -rw-r--r-- | server/szurubooru/func/metrics.py | 4 | ||||
| -rw-r--r-- | server/szurubooru/tests/func/test_metrics.py | 13 | ||||
| -rw-r--r-- | server/szurubooru/tests/search/configs/test_post_metric_search_config.py | 28 |
4 files changed, 58 insertions, 5 deletions
diff --git a/server/szurubooru/api/metric_api.py b/server/szurubooru/api/metric_api.py index a23ac45..12babbd 100644 --- a/server/szurubooru/api/metric_api.py +++ b/server/szurubooru/api/metric_api.py @@ -22,12 +22,13 @@ def _serialize_post_metric( @rest.routes.get('/metrics/?') -def get_post_metrics( +def get_metrics( ctx: rest.Context, params: Dict[str, str] = {}) -> rest.Response: auth.verify_privilege(ctx.user, 'metrics:list') - _search_executor_config.user = ctx.user - return _search_executor.execute_and_serialize( - ctx, lambda post_metric: _serialize_post_metric(ctx, post_metric)) + all_metrics = metrics.get_all_metrics() + return { + 'results': [_serialize_metric(ctx, metric) for metric in all_metrics] + } @rest.routes.post('/metrics/?') @@ -44,3 +45,12 @@ def create_metric( # snapshots.create(metric, ctx.user) ctx.session.commit() return _serialize_metric(ctx, metric) + + +@rest.routes.get('/post-metrics/?') +def get_post_metrics( + ctx: rest.Context, params: Dict[str, str] = {}) -> rest.Response: + auth.verify_privilege(ctx.user, 'metrics:list') + _search_executor_config.user = ctx.user + return _search_executor.execute_and_serialize( + ctx, lambda post_metric: _serialize_post_metric(ctx, post_metric)) diff --git a/server/szurubooru/func/metrics.py b/server/szurubooru/func/metrics.py index 210409c..d2ef8e6 100644 --- a/server/szurubooru/func/metrics.py +++ b/server/szurubooru/func/metrics.py @@ -73,6 +73,10 @@ def serialize_post_metric( return PostMetricSerializer(post_metric).serialize(options) +def get_all_metrics() -> List[model.Metric]: + return db.session.query(model.Metric).all() + + def try_get_post_metric( post: model.Post, metric: model.Metric) -> Optional[model.PostMetric]: diff --git a/server/szurubooru/tests/func/test_metrics.py b/server/szurubooru/tests/func/test_metrics.py index 1a1bec4..997d726 100644 --- a/server/szurubooru/tests/func/test_metrics.py +++ b/server/szurubooru/tests/func/test_metrics.py @@ -50,6 +50,19 @@ def test_try_get_post_metric_range( assert metrics.try_get_post_metric_range(post, metric) is post_metric_range +def test_get_all_metrics(metric_factory): + metric1 = metric_factory() + metric2 = metric_factory() + metric3 = metric_factory() + db.session.add_all([metric1, metric2, metric3]) + db.session.flush() + all_metrics = metrics.get_all_metrics() + assert len(all_metrics) == 3 + assert metric1 in all_metrics + assert metric2 in all_metrics + assert metric3 in all_metrics + + def test_create_metric(tag_factory): tag = tag_factory() db.session.add(tag) diff --git a/server/szurubooru/tests/search/configs/test_post_metric_search_config.py b/server/szurubooru/tests/search/configs/test_post_metric_search_config.py index f2b098e..0b1a199 100644 --- a/server/szurubooru/tests/search/configs/test_post_metric_search_config.py +++ b/server/szurubooru/tests/search/configs/test_post_metric_search_config.py @@ -37,8 +37,10 @@ def test_refresh_metrics(tag_factory, metric_factory): @pytest.mark.parametrize('input,expected_tag_names', [ ('', ['t1:10', 't2:20.5', 't1:30', 't2:40']), + ('*', ['t1:10', 't2:20.5', 't1:30', 't2:40']), ('t1', ['t1:10', 't1:30']), ('t2', ['t2:20.5', 't2:40']), + ('t*', ['t1:10', 't2:20.5', 't1:30', 't2:40']), ('t1,t2', ['t1:10', 't2:20.5', 't1:30', 't2:40']), ('T1,T2', ['t1:10', 't2:20.5', 't1:30', 't2:40']), ]) @@ -80,4 +82,28 @@ def test_filter_by_value( t4 = post_metric_factory(metric=metric, value=40) db.session.add_all([tag, metric, t1, t2, t3, t4]) db.session.flush() - verify_unpaged(input, expected_tag_names)
\ No newline at end of file + verify_unpaged(input, expected_tag_names) + + +@pytest.mark.parametrize('input,expected_tag_names', [ + ('t:13', []), + ('t:10', ['t:10']), + ('t:20.5', ['t:20.5']), + ('t:18.6..', ['t:20.5', 't:30', 't:40']), + ('t-min:18.6', ['t:20.5', 't:30', 't:40']), + ('t:..21.4', ['t:10', 't:20.5']), + ('t-max:21.4', ['t:10', 't:20.5']), + ('t:17..33', ['t:20.5', 't:30']), +]) +def test_filter_by_value( + verify_unpaged, tag_factory, metric_factory, post_metric_factory, + input, expected_tag_names): + tag = tag_factory(names=['t']) + metric = metric_factory(tag) + t1 = post_metric_factory(metric=metric, value=10) + t2 = post_metric_factory(metric=metric, value=30) + t3 = post_metric_factory(metric=metric, value=20.5) + t4 = post_metric_factory(metric=metric, value=40) + db.session.add_all([tag, metric, t1, t2, t3, t4]) + db.session.flush() + verify_unpaged(input, expected_tag_names) |