Skip to content

Commit f2fbf79

Browse files
committed
feat: add support for resolving string annotations in typed method discovery
1 parent 5f81dc6 commit f2fbf79

2 files changed

Lines changed: 72 additions & 1 deletion

File tree

src/iop/messages/dispatch.py

Lines changed: 38 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -306,8 +306,9 @@ def get_handler_info(host: Any, method_name: str) -> tuple[str, str] | None:
306306
if len(params) != 1:
307307
return None
308308

309+
method = getattr(host, method_name)
309310
param: Parameter = next(iter(params.values()))
310-
annotation = param.annotation
311+
annotation = _resolve_annotation(host, method, param.annotation)
311312

312313
if annotation == Parameter.empty:
313314
return None
@@ -322,6 +323,42 @@ def get_handler_info(host: Any, method_name: str) -> tuple[str, str] | None:
322323
return None
323324

324325

326+
def _resolve_annotation(host: Any, method: Callable, annotation: Any) -> Any:
327+
if not isinstance(annotation, str):
328+
return annotation
329+
330+
globalns = _annotation_globalns(method)
331+
localns = _annotation_localns(host)
332+
333+
try:
334+
resolved = eval(annotation, globalns, localns) # noqa: B307
335+
except Exception:
336+
resolved = annotation
337+
338+
if isinstance(resolved, str):
339+
# Quoted postponed annotations evaluate to a string first.
340+
try:
341+
resolved = eval(resolved, globalns, localns) # noqa: B307
342+
except Exception:
343+
if "." in resolved:
344+
return resolved
345+
return Parameter.empty
346+
347+
return resolved
348+
349+
350+
def _annotation_globalns(method: Callable) -> dict[str, Any]:
351+
function = getattr(method, "__func__", method)
352+
return getattr(function, "__globals__", {})
353+
354+
355+
def _annotation_localns(host: Any) -> dict[str, Any]:
356+
namespace: dict[str, Any] = {}
357+
for klass in reversed(type(host).__mro__):
358+
namespace.update(vars(klass))
359+
return namespace
360+
361+
325362
def _message_class_name(message_type: Any) -> str | None:
326363
if isinstance(message_type, str):
327364
return message_type

src/tests/unit/test_dispatch.py

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -287,6 +287,40 @@ def log_warning(self, message):
287287
assert logs == []
288288

289289

290+
def test_typed_method_discovery_resolves_string_annotations():
291+
class Host:
292+
def on_message(self, request):
293+
return "fallback"
294+
295+
def handle_message(self, request: "MessageTest"):
296+
return "handled"
297+
298+
host = Host()
299+
create_dispatch(host)
300+
301+
assert host.DISPATCH == [
302+
(f"{MessageTest.__module__}.{MessageTest.__name__}", "handle_message")
303+
]
304+
assert dispatch_message(host, MessageTest(text="test", number=1)) == "handled"
305+
306+
307+
def test_typed_method_discovery_ignores_unresolved_bare_string_annotations():
308+
class Host:
309+
def on_message(self, request):
310+
return "fallback"
311+
312+
def handle_message(self, request):
313+
return "handled"
314+
315+
Host.handle_message.__annotations__["request"] = "UnknownMessage"
316+
317+
host = Host()
318+
create_dispatch(host)
319+
320+
assert host.DISPATCH == []
321+
assert dispatch_message(host, MessageTest(text="test", number=1)) == "fallback"
322+
323+
290324
def test_duplicate_legacy_mappings_log_discarded_handler():
291325
logs = []
292326

0 commit comments

Comments
 (0)