From 52872f5d7d412dff617bee69738b4976469be0b6 Mon Sep 17 00:00:00 2001 From: Alan Date: Thu, 1 Oct 2026 09:15:27 +0800 Subject: [PATCH] Fix set operations on frame flags --- CHANGELOG.rst | 3 ++- src/hyperframe/flags.py | 13 +++++++-- tests/test_flags.py | 58 +++++++++++++++++++++++++++++++++++++++++ 3 files changed, 71 insertions(+), 3 deletions(-) diff --git a/CHANGELOG.rst b/CHANGELOG.rst index 750102f..b62d493 100644 --- a/CHANGELOG.rst +++ b/CHANGELOG.rst @@ -19,7 +19,8 @@ dev **Bugfixes** -- +- Fix set operations on frame flags raising ``AttributeError``, including + in-place intersection. Non-mutating operations now return ordinary sets. 6.1.0 (2025-01-22) ------------------ diff --git a/src/hyperframe/flags.py b/src/hyperframe/flags.py index e5f4a22..2e0888d 100644 --- a/src/hyperframe/flags.py +++ b/src/hyperframe/flags.py @@ -4,7 +4,9 @@ from __future__ import annotations from collections.abc import Iterable, Iterator, MutableSet -from typing import NamedTuple +from typing import NamedTuple, TypeVar + +_T = TypeVar("_T") class Flag(NamedTuple): @@ -18,13 +20,20 @@ class Flags(MutableSet): # type: ignore elements. Will behave like a regular set(), except that a ValueError will be thrown - when .add()ing unexpected flags. + when .add()ing unexpected flags. Non-mutating set operations return + regular sets. """ def __init__(self, defined_flags: Iterable[Flag]) -> None: self._valid_flags = {flag.name for flag in defined_flags} self._flags: set[str] = set() + @classmethod + def _from_iterable(cls, it: Iterable[_T]) -> set[_T]: + # The set mixins pass flag names, not the definitions our constructor + # needs. Results are ordinary sets; in-place operations still validate. + return set(it) + def __repr__(self) -> str: return repr(sorted(self._flags)) diff --git a/tests/test_flags.py b/tests/test_flags.py index 17b6a98..1b75aa2 100644 --- a/tests/test_flags.py +++ b/tests/test_flags.py @@ -1,3 +1,5 @@ +import operator + from hyperframe.frame import ( Flags, Flag, ) @@ -40,3 +42,59 @@ def test_repr(self): assert repr(flags) == "['VALID_FLAG']" flags.add("OTHER_FLAG") assert repr(flags) == "['OTHER_FLAG', 'VALID_FLAG']" + + +@pytest.mark.parametrize("operation, expected", [ + (operator.and_, {"FIRST"}), + (operator.or_, {"FIRST", "SECOND", "OTHER"}), + (operator.sub, {"SECOND"}), + (operator.xor, {"SECOND", "OTHER"}), +]) +@pytest.mark.parametrize("reflected", [False, True]) +def test_set_operations(operation, expected, reflected): + flags = Flags([Flag("FIRST", 0x01), Flag("SECOND", 0x02)]) + flags |= {"FIRST", "SECOND"} + other = {"FIRST", "OTHER"} + + if reflected: + result = operation(other, flags) + if operation is operator.sub: + expected = {"OTHER"} + else: + result = operation(flags, other) + + assert result == expected + assert isinstance(result, set) + assert flags == {"FIRST", "SECOND"} + + +@pytest.mark.parametrize("operation, expected", [ + (operator.iand, {"FIRST"}), + (operator.isub, {"SECOND"}), + (operator.ixor, {"SECOND"}), + (operator.ior, {"FIRST", "SECOND"}), +]) +def test_in_place_set_operations(operation, expected): + flags = Flags([Flag("FIRST", 0x01), Flag("SECOND", 0x02)]) + flags |= {"FIRST", "SECOND"} + + result = operation(flags, {"FIRST"}) + + assert result is flags + assert flags == expected + with pytest.raises(ValueError): + flags.add("OTHER") + + +@pytest.mark.parametrize("operation", [operator.ior, operator.ixor]) +def test_in_place_operations_reject_unknown_flags(operation): + flags = Flags([Flag("FIRST", 0x01)]) + with pytest.raises(ValueError): + operation(flags, {"OTHER"}) + + +def test_empty_set_operation(): + flags = Flags([Flag("FIRST", 0x01)]) + result = flags & {"FIRST"} + assert result == set() + assert isinstance(result, set)