feat/stubgen: fix incorrect stub generation
This commit is contained in:
@@ -33,5 +33,8 @@ def generate_enum(
|
|||||||
literals.append(f'"{e["name"]}"')
|
literals.append(f'"{e["name"]}"')
|
||||||
|
|
||||||
emitter.add_typing_import("Literal")
|
emitter.add_typing_import("Literal")
|
||||||
emitter.add_type_alias(type_info.name, f"Literal[{', '.join(literals)}]")
|
line = f"Literal[{', '.join(literals)}]"
|
||||||
emitter.add_blank_line()
|
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)
|
sinfo = registry.get(sid)
|
||||||
if sinfo is not None:
|
if sinfo is not None:
|
||||||
sinfo.parent_type_id = type_info.type_id
|
sinfo.parent_type_id = type_info.type_id
|
||||||
|
dname = f"{mname}{suffix}"
|
||||||
sinfo.scoped_name = (
|
sinfo.scoped_name = (
|
||||||
f"{type_info.scoped_name}.{mname}.{suffix}"
|
f"{type_info.scoped_name}.{dname}"
|
||||||
)
|
|
||||||
# Override name for class nesting
|
|
||||||
sinfo.display_name = (
|
|
||||||
f"{mname}{suffix}"
|
|
||||||
)
|
)
|
||||||
|
sinfo.display_name = dname
|
||||||
|
|
||||||
|
|
||||||
def generate_interface(
|
def generate_interface(
|
||||||
@@ -76,6 +74,8 @@ def generate_interface(
|
|||||||
"result_info": result_info,
|
"result_info": result_info,
|
||||||
})
|
})
|
||||||
|
|
||||||
|
emitter.add_typing_import("Any")
|
||||||
|
|
||||||
# ── Main interface class ───────────────────────────────────────────
|
# ── Main interface class ───────────────────────────────────────────
|
||||||
emitter.begin_class(name)
|
emitter.begin_class(name)
|
||||||
|
|
||||||
@@ -106,7 +106,7 @@ def generate_interface(
|
|||||||
emitter.begin_class(client_name, bases=[name])
|
emitter.begin_class(client_name, bases=[name])
|
||||||
emitter.add_method(
|
emitter.add_method(
|
||||||
"_new_client",
|
"_new_client",
|
||||||
params=["self", "server"],
|
params=["self", "server: Any"],
|
||||||
return_type=client_name,
|
return_type=client_name,
|
||||||
)
|
)
|
||||||
emitter.end_class()
|
emitter.end_class()
|
||||||
@@ -121,7 +121,7 @@ def generate_interface(
|
|||||||
|
|
||||||
emitter.add_method(
|
emitter.add_method(
|
||||||
mname,
|
mname,
|
||||||
params=["self", f"_params: {param_type}", "_context"],
|
params=["self", f"_params: {param_type}", "_context: Any"],
|
||||||
return_type=None,
|
return_type=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -76,61 +76,71 @@ def generate_struct(
|
|||||||
tv_names = generate_type_vars(emitter, type_info)
|
tv_names = generate_type_vars(emitter, type_info)
|
||||||
emitter.add_typing_import("Generic")
|
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:
|
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:
|
else:
|
||||||
emitter.begin_class(name)
|
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:
|
for field_name, type_str in non_union_fields:
|
||||||
emitter.add_field(field_name, type_str)
|
emitter.add_field(field_name, type_str)
|
||||||
|
|
||||||
if discriminant_count > 0 and union_fields:
|
if discriminant_count > 0 and union_fields:
|
||||||
_generate_union_methods(emitter, 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
|
# Factory methods + to_dict
|
||||||
_generate_factory_methods(emitter, name)
|
_generate_factory_methods(emitter, qbuilder, qreader)
|
||||||
emitter.add_method("to_dict", params=["self"], return_type="dict")
|
emitter.add_method("to_dict", params=["self"], return_type="dict[str, Any]")
|
||||||
|
|
||||||
# Nested structs
|
|
||||||
for nt in nested_types:
|
|
||||||
if nt.kind == "struct":
|
|
||||||
emitter.add_blank_line()
|
|
||||||
generate_struct(emitter, nt, registry)
|
|
||||||
|
|
||||||
emitter.end_class()
|
emitter.end_class()
|
||||||
|
|
||||||
# ── class: Reader ──────────────────────────────────────────────────
|
# ── class: Reader ──────────────────────────────────────────────────
|
||||||
reader_name = name + "Reader"
|
emitter.begin_class(reader_name, bases=[base_name])
|
||||||
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.add_method(
|
emitter.add_method(
|
||||||
"as_builder", params=["self"], return_type=name + "Builder",
|
"as_builder", params=["self"], return_type=qbuilder,
|
||||||
)
|
)
|
||||||
emitter.end_class()
|
emitter.end_class()
|
||||||
|
|
||||||
# ── class: Builder ─────────────────────────────────────────────────
|
# ── class: Builder ─────────────────────────────────────────────────
|
||||||
builder_name = name + "Builder"
|
emitter.begin_class(builder_name, bases=[base_name])
|
||||||
emitter.begin_class(builder_name, bases=[name])
|
|
||||||
|
|
||||||
for field_name, type_str in non_union_fields:
|
_generate_builder_methods(emitter, qbuilder, qreader)
|
||||||
emitter.add_field(field_name, _to_builder_type(type_str))
|
|
||||||
|
|
||||||
_generate_builder_methods(emitter, name)
|
|
||||||
emitter.end_class()
|
emitter.end_class()
|
||||||
|
|
||||||
|
|
||||||
@@ -148,7 +158,10 @@ def _resolve_field_type(
|
|||||||
py_type = resolve_type(type_dict, registry)
|
py_type = resolve_type(type_dict, registry)
|
||||||
if py_type is not None:
|
if py_type is not None:
|
||||||
q = py_type.render()
|
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"
|
return "Any"
|
||||||
|
|
||||||
if which == "list":
|
if which == "list":
|
||||||
@@ -173,11 +186,9 @@ def _resolve_list_field_type(
|
|||||||
inner_which = _type_which(inner)
|
inner_which = _type_which(inner)
|
||||||
|
|
||||||
if inner_which == "struct":
|
if inner_which == "struct":
|
||||||
type_id = type_id_to_int(inner["struct"]["typeId"])
|
pt = resolve_type(inner, registry)
|
||||||
info = registry.get(type_id)
|
if pt is not None:
|
||||||
if info is not None:
|
el_type = pt.render()
|
||||||
q = info.scoped_name
|
|
||||||
el_type = f"{q} | {q}Builder | {q}Reader"
|
|
||||||
else:
|
else:
|
||||||
el_type = "Any"
|
el_type = "Any"
|
||||||
elif inner_which == "enum":
|
elif inner_which == "enum":
|
||||||
@@ -206,33 +217,6 @@ def _type_which(type_dict: NodeDict) -> str | None:
|
|||||||
return 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 ──────────────────────────────────────────────────────
|
# ── method generation ──────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
@@ -257,16 +241,16 @@ def _generate_union_methods(
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _generate_factory_methods(emitter: Emitter, name: str) -> None:
|
def _generate_factory_methods(
|
||||||
emitter.add_typing_import("Iterator", "collections.abc")
|
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")
|
emitter.add_typing_import("contextmanager", "contextlib")
|
||||||
|
|
||||||
reader_name = name + "Reader"
|
|
||||||
builder_name = name + "Builder"
|
|
||||||
|
|
||||||
emitter.add_blank_line()
|
emitter.add_blank_line()
|
||||||
emitter.add_static_method(
|
emitter.add_static_method(
|
||||||
"new_message", params=["**kwargs"], return_type=builder_name,
|
"new_message", params=["**kwargs: Any"], return_type=qbuilder,
|
||||||
)
|
)
|
||||||
emitter.add_blank_line()
|
emitter.add_blank_line()
|
||||||
|
|
||||||
@@ -279,7 +263,7 @@ def _generate_factory_methods(emitter: Emitter, name: str) -> None:
|
|||||||
"traversal_limit_in_words: int | None = ...",
|
"traversal_limit_in_words: int | None = ...",
|
||||||
"nesting_limit: int | None = ...",
|
"nesting_limit: int | None = ...",
|
||||||
],
|
],
|
||||||
return_type=f"Iterator[{reader_name}]",
|
return_type=f"Generator[{qreader}, None, None]",
|
||||||
)
|
)
|
||||||
emitter.add_blank_line()
|
emitter.add_blank_line()
|
||||||
|
|
||||||
@@ -290,20 +274,19 @@ def _generate_factory_methods(emitter: Emitter, name: str) -> None:
|
|||||||
"traversal_limit_in_words: int | None = ...",
|
"traversal_limit_in_words: int | None = ...",
|
||||||
"nesting_limit: int | None = ...",
|
"nesting_limit: int | None = ...",
|
||||||
],
|
],
|
||||||
return_type=reader_name,
|
return_type=qreader,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
def _generate_builder_methods(emitter: Emitter, name: str) -> None:
|
def _generate_builder_methods(
|
||||||
reader_name = name + "Reader"
|
emitter: Emitter, qbuilder: str, qreader: str
|
||||||
builder_name = name + "Builder"
|
) -> None:
|
||||||
|
|
||||||
emitter.add_blank_line()
|
emitter.add_blank_line()
|
||||||
emitter.add_static_method(
|
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_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_blank_line()
|
||||||
emitter.add_method("to_bytes", params=["self"], return_type="bytes")
|
emitter.add_method("to_bytes", params=["self"], return_type="bytes")
|
||||||
emitter.add_blank_line()
|
emitter.add_blank_line()
|
||||||
@@ -311,7 +294,7 @@ def _generate_builder_methods(emitter: Emitter, name: str) -> None:
|
|||||||
emitter.add_blank_line()
|
emitter.add_blank_line()
|
||||||
emitter.add_method("to_segments", params=["self"], return_type="list[bytes]")
|
emitter.add_method("to_segments", params=["self"], return_type="list[bytes]")
|
||||||
emitter.add_blank_line()
|
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_blank_line()
|
||||||
|
|
||||||
emitter.add_typing_import("IO", "typing")
|
emitter.add_typing_import("IO", "typing")
|
||||||
|
|||||||
Reference in New Issue
Block a user