feat/stubgen: fix incorrect stub generation
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user