feat/stubgen: fix incorrect stub generation

This commit is contained in:
2026-07-09 18:02:23 +08:00
parent 7156567302
commit 556969a054
3 changed files with 75 additions and 89 deletions
+5 -2
View File
@@ -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)
+7 -7
View File
@@ -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,
)
+63 -80
View File
@@ -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")