diff --git a/beets/dbcore/queryparse.py b/beets/dbcore/queryparse.py index 288b71c89..b610ec369 100644 --- a/beets/dbcore/queryparse.py +++ b/beets/dbcore/queryparse.py @@ -7,6 +7,7 @@ import re import shlex from dataclasses import dataclass from functools import partial, reduce +from itertools import groupby from typing import TYPE_CHECKING, NamedTuple from typing_extensions import Self @@ -17,7 +18,7 @@ from beets.util import cached_classproperty from . import query, sort if TYPE_CHECKING: - from collections.abc import Collection, Sequence + from collections.abc import Iterable, Sequence from beets.library.models import LibModel @@ -139,7 +140,7 @@ class QueryTerm: # TYPING ERROR def build_and_query( - model_cls: type[LibModel], query_parts: Collection[str] + model_cls: type[LibModel], query_parts: Iterable[str] ) -> query.Query: """Build a query by combining search terms with a logical AND.""" if not query_parts: @@ -221,31 +222,21 @@ def parse_sorted_query( """Given a list of strings, create the `Query` and `Sort` that they represent. """ - # Separate query token and sort token. - query_parts = [] - sort_parts = [] + sort_parts = [p for p in parts if SortTerm.check_valid(p)] + query_parts = [p for p in parts if not SortTerm.check_valid(p)] - # Split up query in to comma-separated subqueries, each representing - # an AndQuery, which need to be joined together in one OrQuery - subquery_parts = [] - for part in [*parts, ","]: - if part.endswith(","): - # Ensure we can catch "foo, bar" as well as "foo , bar" - last_subquery_part = part[:-1] - if last_subquery_part: - subquery_parts.append(last_subquery_part) - # Parse the subquery in to a single Query - query_parts.append(build_and_query(model_cls, subquery_parts)) - del subquery_parts[:] - elif SortTerm.check_valid(part): - sort_parts.append(part) - else: - subquery_parts.append(part) + queries = [ + build_and_query(model_cls, g) + for k, g in groupby(query_parts, lambda p: p == ",") + if not k + ] + if not queries or "," in {query_parts[0], query_parts[-1]}: + queries.append(query.TrueQuery()) - # Avoid needlessly wrapping single statements in an OR - q = query.OrQuery(query_parts) if len(query_parts) > 1 else query_parts[0] - s = sort_from_strings(model_cls, sort_parts) - return q, s + return ( + reduce(operator.or_, queries), + sort_from_strings(model_cls, sort_parts), + ) class ModelQuery(NamedTuple): diff --git a/test/dbcore/test_queryparse.py b/test/dbcore/test_queryparse.py index 0765b872d..213b53e52 100644 --- a/test/dbcore/test_queryparse.py +++ b/test/dbcore/test_queryparse.py @@ -145,13 +145,16 @@ class ParseSortedQueryTest(unittest.TestCase): q, s = self.psq("foo , bar") assert isinstance(q, query.OrQuery) assert isinstance(s, sort.NullSort) - assert len(q.subqueries) == 2 + assert len(q.subqueries) == 4 def test_no_space_before_comma_or_query(self): q, s = self.psq("foo, bar") - assert isinstance(q, query.OrQuery) + assert isinstance(q, query.AndQuery) + assert isinstance(q.subqueries[0], query.OrQuery) + assert isinstance(q.subqueries[1], query.OrQuery) + assert len(q.subqueries[0].subqueries) == 2 + assert len(q.subqueries[1].subqueries) == 2 assert isinstance(s, sort.NullSort) - assert len(q.subqueries) == 2 def test_no_spaces_or_query(self): q, s = self.psq("foo,bar") @@ -163,13 +166,13 @@ class ParseSortedQueryTest(unittest.TestCase): q, s = self.psq("foo , bar ,") assert isinstance(q, query.OrQuery) assert isinstance(s, sort.NullSort) - assert len(q.subqueries) == 3 + assert len(q.subqueries) == 5 def test_leading_comma_or_query(self): q, s = self.psq(", foo , bar") assert isinstance(q, query.OrQuery) assert isinstance(s, sort.NullSort) - assert len(q.subqueries) == 3 + assert len(q.subqueries) == 5 def test_only_direction(self): q, s = self.psq("-")