From 556969a0546e200075c257edc2c6e40db37f28a4 Mon Sep 17 00:00:00 2001 From: worldmozara Date: Thu, 9 Jul 2026 18:02:23 +0800 Subject: [PATCH] feat/stubgen: fix incorrect stub generation --- src/capnp_stubgen/gen_enum.py | 7 +- src/capnp_stubgen/gen_interface.py | 14 +-- src/capnp_stubgen/gen_struct.py | 143 +++++++++++++---------------- 3 files changed, 75 insertions(+), 89 deletions(-) diff --git a/src/capnp_stubgen/gen_enum.py b/src/capnp_stubgen/gen_enum.py index 940f0a9..f561972 100644 --- a/src/capnp_stubgen/gen_enum.py +++ b/src/capnp_stubgen/gen_enum.py @@ -33,5 +33,8 @@ def generate_enum( literals.append(f'"{e["name"]}"') emitter.add_typing_import("Literal") - emitter.add_type_alias(type_info.name, f"Literal[{', '.join(literals)}]") - emitter.add_blank_line() + line = f"Literal[{', '.join(literals)}]" + if type_info.parent_type_id is not None and type_info.scoped_name.count(".") > 0: + # Inside a struct/interface class — type alias needs annotation + line += " # pyright: ignore[reportUnannotatedClassAttribute]" + emitter.add_type_alias(type_info.name, line) diff --git a/src/capnp_stubgen/gen_interface.py b/src/capnp_stubgen/gen_interface.py index 0a1e6bf..80d73f0 100644 --- a/src/capnp_stubgen/gen_interface.py +++ b/src/capnp_stubgen/gen_interface.py @@ -43,13 +43,11 @@ def fixup_interface_methods(registry: TypeRegistry) -> None: sinfo = registry.get(sid) if sinfo is not None: sinfo.parent_type_id = type_info.type_id + dname = f"{mname}{suffix}" sinfo.scoped_name = ( - f"{type_info.scoped_name}.{mname}.{suffix}" - ) - # Override name for class nesting - sinfo.display_name = ( - f"{mname}{suffix}" + f"{type_info.scoped_name}.{dname}" ) + sinfo.display_name = dname def generate_interface( @@ -76,6 +74,8 @@ def generate_interface( "result_info": result_info, }) + emitter.add_typing_import("Any") + # ── Main interface class ─────────────────────────────────────────── emitter.begin_class(name) @@ -106,7 +106,7 @@ def generate_interface( emitter.begin_class(client_name, bases=[name]) emitter.add_method( "_new_client", - params=["self", "server"], + params=["self", "server: Any"], return_type=client_name, ) emitter.end_class() @@ -121,7 +121,7 @@ def generate_interface( emitter.add_method( mname, - params=["self", f"_params: {param_type}", "_context"], + params=["self", f"_params: {param_type}", "_context: Any"], return_type=None, ) diff --git a/src/capnp_stubgen/gen_struct.py b/src/capnp_stubgen/gen_struct.py index 000ba20..04dbbdc 100644 --- a/src/capnp_stubgen/gen_struct.py +++ b/src/capnp_stubgen/gen_struct.py @@ -76,61 +76,71 @@ def generate_struct( tv_names = generate_type_vars(emitter, type_info) emitter.add_typing_import("Generic") + # Simple names for class definitions (no prefix) + reader_name = name + "Reader" + builder_name = name + "Builder" + # Fully-qualified names for type references (with parent scope) + parent_prefix = "" + if "." in type_info.scoped_name: + parent_prefix = type_info.scoped_name.rsplit(".", 1)[0] + "." + qreader = parent_prefix + reader_name + qbuilder = parent_prefix + builder_name + qualified_name = parent_prefix + name + tv_params = ", ".join(tv_names) if tv_names: - emitter.begin_class(name, bases=[f"Generic[{', '.join(tv_names)}]"]) + qreader += f"[{tv_params}]" + qbuilder += f"[{tv_params}]" + qualified_name += f"[{tv_params}]" + base_name = qualified_name + + if tv_names: + emitter.begin_class(name, bases=[f"Generic[{tv_params}]"]) else: emitter.begin_class(name) + # Nested structs first (so forward refs within class body resolve) + for nt in nested_types: + if nt.kind == "struct": + emitter.add_blank_line() + generate_struct(emitter, nt, registry) + + # Nested enums + for nt in nested_types: + if nt.kind == "enum": + emitter.add_blank_line() + generate_enum(emitter, nt) + + # Nested constants + for nt in nested_types: + if nt.kind == "const": + emitter.add_blank_line() + generate_const(emitter, nt, registry) + for field_name, type_str in non_union_fields: emitter.add_field(field_name, type_str) if discriminant_count > 0 and union_fields: _generate_union_methods(emitter, union_fields) - # Nested enums first - for nt in nested_types: - if nt.kind == "enum": - emitter.add_blank_line() - generate_enum(emitter, nt) - - # Nested constants (class-level) - for nt in nested_types: - if nt.kind == "const": - emitter.add_blank_line() - generate_const(emitter, nt, registry) - # Factory methods + to_dict - _generate_factory_methods(emitter, name) - emitter.add_method("to_dict", params=["self"], return_type="dict") - - # Nested structs - for nt in nested_types: - if nt.kind == "struct": - emitter.add_blank_line() - generate_struct(emitter, nt, registry) + _generate_factory_methods(emitter, qbuilder, qreader) + emitter.add_method("to_dict", params=["self"], return_type="dict[str, Any]") emitter.end_class() # ── class: Reader ────────────────────────────────────────────────── - reader_name = name + "Reader" - emitter.begin_class(reader_name, bases=[name]) - - for field_name, type_str in non_union_fields: - emitter.add_field(field_name, _to_reader_type(type_str)) + emitter.begin_class(reader_name, bases=[base_name]) emitter.add_method( - "as_builder", params=["self"], return_type=name + "Builder", + "as_builder", params=["self"], return_type=qbuilder, ) emitter.end_class() # ── class: Builder ───────────────────────────────────────────────── - builder_name = name + "Builder" - emitter.begin_class(builder_name, bases=[name]) + emitter.begin_class(builder_name, bases=[base_name]) - for field_name, type_str in non_union_fields: - emitter.add_field(field_name, _to_builder_type(type_str)) + _generate_builder_methods(emitter, qbuilder, qreader) - _generate_builder_methods(emitter, name) emitter.end_class() @@ -148,7 +158,10 @@ def _resolve_field_type( py_type = resolve_type(type_dict, registry) if py_type is not None: q = py_type.render() - return f"{q} | {q}Builder | {q}Reader" + if py_type.params: + # Branded types (e.g. Holder[str]) — no separate Builder/Reader + return q + return q return "Any" if which == "list": @@ -173,11 +186,9 @@ def _resolve_list_field_type( inner_which = _type_which(inner) if inner_which == "struct": - type_id = type_id_to_int(inner["struct"]["typeId"]) - info = registry.get(type_id) - if info is not None: - q = info.scoped_name - el_type = f"{q} | {q}Builder | {q}Reader" + pt = resolve_type(inner, registry) + if pt is not None: + el_type = pt.render() else: el_type = "Any" elif inner_which == "enum": @@ -206,33 +217,6 @@ def _type_which(type_dict: NodeDict) -> str | None: return None -# ── Reader/Builder type transformation ───────────────────────────────────── - - -def _to_reader_type(type_str: str) -> str: - return _transform_variant(type_str, "Reader") - - -def _to_builder_type(type_str: str) -> str: - return _transform_variant(type_str, "Builder") - - -def _transform_variant(type_str: str, variant: str) -> str: - if " | " not in type_str: - return type_str - - if type_str.startswith("Sequence["): - prefix = "Sequence[" - inner = type_str[len(prefix):-1] - return f"Sequence[{_transform_variant(inner, variant)}]" - - parts = [p.strip() for p in type_str.split("|")] - for p in parts: - if p.endswith(variant): - return p - return parts[0] - - # ── method generation ────────────────────────────────────────────────────── @@ -257,16 +241,16 @@ def _generate_union_methods( ) -def _generate_factory_methods(emitter: Emitter, name: str) -> None: - emitter.add_typing_import("Iterator", "collections.abc") +def _generate_factory_methods( + emitter: Emitter, qbuilder: str, qreader: str +) -> None: + emitter.add_typing_import("Any") + emitter.add_typing_import("Generator", "collections.abc") emitter.add_typing_import("contextmanager", "contextlib") - reader_name = name + "Reader" - builder_name = name + "Builder" - emitter.add_blank_line() emitter.add_static_method( - "new_message", params=["**kwargs"], return_type=builder_name, + "new_message", params=["**kwargs: Any"], return_type=qbuilder, ) emitter.add_blank_line() @@ -279,7 +263,7 @@ def _generate_factory_methods(emitter: Emitter, name: str) -> None: "traversal_limit_in_words: int | None = ...", "nesting_limit: int | None = ...", ], - return_type=f"Iterator[{reader_name}]", + return_type=f"Generator[{qreader}, None, None]", ) emitter.add_blank_line() @@ -290,20 +274,19 @@ def _generate_factory_methods(emitter: Emitter, name: str) -> None: "traversal_limit_in_words: int | None = ...", "nesting_limit: int | None = ...", ], - return_type=reader_name, + return_type=qreader, ) -def _generate_builder_methods(emitter: Emitter, name: str) -> None: - reader_name = name + "Reader" - builder_name = name + "Builder" - +def _generate_builder_methods( + emitter: Emitter, qbuilder: str, qreader: str +) -> None: emitter.add_blank_line() emitter.add_static_method( - "from_dict", params=["dictionary: dict"], return_type=builder_name, + "from_dict", params=["dictionary: dict[str, Any]"], return_type=qbuilder, ) emitter.add_blank_line() - emitter.add_method("copy", params=["self"], return_type=builder_name) + emitter.add_method("copy", params=["self"], return_type=qbuilder) emitter.add_blank_line() emitter.add_method("to_bytes", params=["self"], return_type="bytes") emitter.add_blank_line() @@ -311,7 +294,7 @@ def _generate_builder_methods(emitter: Emitter, name: str) -> None: emitter.add_blank_line() emitter.add_method("to_segments", params=["self"], return_type="list[bytes]") emitter.add_blank_line() - emitter.add_method("as_reader", params=["self"], return_type=reader_name) + emitter.add_method("as_reader", params=["self"], return_type=qreader) emitter.add_blank_line() emitter.add_typing_import("IO", "typing")