Skip to content

Commit bf481f5

Browse files
committed
Add Registry.extensions_for
1 parent c7f2d04 commit bf481f5

3 files changed

Lines changed: 34 additions & 13 deletions

File tree

src/protobuf/_registry.py

Lines changed: 26 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -27,6 +27,7 @@
2727
from ._enum import Enum
2828

2929
_Types = DescMessage | DescEnum | DescExtension | DescService
30+
_MessageTypeInfo = DescMessage | str | Message
3031
_T = TypeVar("_T", bound=_Types)
3132

3233

@@ -176,8 +177,21 @@ def extension(self, type_name: str) -> DescExtension | None:
176177
msg = self._types.get(type_name)
177178
return msg if isinstance(msg, DescExtension) else None
178179

180+
def extensions_for(self, type_info: _MessageTypeInfo) -> list[DescExtension]:
181+
"""Look up all extensions for a given message type.
182+
183+
Args:
184+
type_info: The extended message, either as a
185+
DescMessage, its fully qualified name, or a Message instance.
186+
187+
Returns:
188+
A list of extension descriptors for the given message type.
189+
"""
190+
msg_type_name = _resolve_type_name(type_info)
191+
return list(self._extendees.get(msg_type_name, {}).values())
192+
179193
def extension_for(
180-
self, type_info: DescMessage | str | Message, number: int
194+
self, type_info: _MessageTypeInfo, number: int
181195
) -> DescExtension | None:
182196
"""Look up an extension by the message it extends and field number.
183197
@@ -189,14 +203,7 @@ def extension_for(
189203
Returns:
190204
The descriptor for the extension, or `None` if not found.
191205
"""
192-
match type_info:
193-
case DescMessage():
194-
msg_type_name = type_info.type_name
195-
case Message():
196-
msg_type_name = type_info.desc().type_name
197-
case str():
198-
msg_type_name = type_info
199-
206+
msg_type_name = _resolve_type_name(type_info)
200207
return self._extendees[msg_type_name].get(number)
201208

202209
def _get_type(self, type_name: str, typ: type[_T]) -> None | _T:
@@ -220,3 +227,13 @@ def __iter__(
220227
yield from self._types.values()
221228

222229
__slots__ = "_extendees", "_files", "_types"
230+
231+
232+
def _resolve_type_name(type_info: _MessageTypeInfo) -> str:
233+
match type_info:
234+
case DescMessage():
235+
return type_info.type_name
236+
case Message():
237+
return type_info.desc().type_name
238+
case str():
239+
return type_info

src/protobuf/_to_json.py

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -183,10 +183,9 @@ def _message_to_json_value(message: Message, opts: ToJsonOptions) -> JsonValue:
183183
result[json_key] = json_value
184184

185185
# Extension fields
186-
if opts.registry and (uf := message._unknown_fields):
187-
for field_number in uf:
188-
ext_desc = opts.registry.extension_for(message._desc, field_number)
189-
if not ext_desc:
186+
if opts.registry:
187+
for ext_desc in opts.registry.extensions_for(message):
188+
if ext_desc.type not in message:
190189
continue
191190
value = message[ext_desc.type]
192191
match field_value := ext_desc.value:

tests/test_registry.py

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ def test_empty_registry() -> None:
3232
assert reg.extension("foo") is None
3333
assert reg.service("foo") is None
3434
assert reg.extension_for("foo", 1) is None
35+
assert reg.extensions_for("foo") == []
3536

3637

3738
@pytest.fixture
@@ -100,6 +101,8 @@ def test_extension_without_message(file: DescFile) -> None:
100101
assert reg.message("P.M") is None
101102
assert reg.extension_for("P.M", 100) is file.extensions[0]
102103
assert reg.extension_for(file.messages[0], 100) is file.extensions[0]
104+
assert reg.extensions_for("P.M") == [file.extensions[0]]
105+
assert reg.extensions_for(file.messages[0]) == [file.extensions[0]]
103106

104107

105108
def test_multiple(protoc: Protoc, file: DescFile) -> None:
@@ -130,6 +133,7 @@ def test_multiple(protoc: Protoc, file: DescFile) -> None:
130133
assert reg.enum("O.E") is other.enums[0]
131134
assert reg.extension("O.ext") is other.extensions[0]
132135
assert reg.extension_for("O.M", 100) is other.extensions[0]
136+
assert reg.extensions_for("O.M") == [other.extensions[0]]
133137

134138

135139
def test_last_win(protoc: Protoc, file: DescFile) -> None:
@@ -183,3 +187,4 @@ def assert_all(file: DescFile, reg: Registry) -> None:
183187
assert reg.enum("P.E") is file.enums[0]
184188
assert reg.extension("P.ext") is file.extensions[0]
185189
assert reg.extension_for("P.M", 100) is file.extensions[0]
190+
assert file.extensions[0] in reg.extensions_for("P.M")

0 commit comments

Comments
 (0)