Compare commits

...

10 Commits

20 changed files with 439 additions and 239 deletions
+49 -41
View File
@@ -17,7 +17,7 @@ Generates `.pyi` type stub files for Cap'n Proto schemas, providing IDE autocomp
| Arch Linux | `sudo pacman -S capnproto` |
| Ubuntu/Debian | `sudo apt install capnproto` |
| macOS | `brew install capnp` |
| Windows | Download from [capnproto.org](https://capnproto.org) |
| Windows | Download from [capnproto.org](https://capnproto.org/install.html) |
## Installation
@@ -33,82 +33,90 @@ capnp compile -opy myschema.capnp
This generates `myschema_capnp.pyi` — a Python type stub file that type checkers can use. Runtime import (`import myschema_capnp`) is handled automatically by pycapnp's import hook.
### Generated output
### Multi-file projects
For a schema like:
```bash
capnp compile -opy --src-prefix=. src/schema.capnp
```
Cross-file type references automatically produce relative imports in the generated stubs.
## Example
Input (`addressbook.capnp`):
```capnp
struct Person {
name @0 :Text;
age @1 :Int32;
}
phones @2 :List(PhoneNumber);
enum Gender { male @0; female @1; }
struct PhoneNumber {
number @0 :Text;
type @1 :Type;
enum Type { mobile @0; home @1; work @2; }
}
}
```
`myschema_capnp.pyi`:
Output (`addressbook_capnp.pyi`):
```python
"""Auto-generated type stub for myschema.capnp"""
"""Auto-generated type stub for addressbook.capnp"""
from __future__ import annotations
from collections.abc import Iterator, Sequence
from contextlib import contextmanager
from typing import IO, Literal
Gender = Literal["male", "female"]
class Person:
name: str
age: int
phones: Sequence[Person.PhoneNumber | Person.PhoneNumberBuilder | Person.PhoneNumberReader]
class PhoneNumber:
Type = Literal["mobile", "home", "work"]
type: Type
number: str
...
@staticmethod
def new_message(**kwargs) -> PersonBuilder: ...
@staticmethod
@contextmanager
def from_bytes(data: bytes, ...) -> Iterator[PersonReader]: ...
@staticmethod
def from_bytes_packed(data: bytes, ...) -> PersonReader: ...
def to_dict(self) -> dict: ...
class PersonReader(Person):
def as_builder(self) -> PersonBuilder: ...
class PersonBuilder(Person):
@staticmethod
def from_dict(dictionary: dict) -> PersonBuilder: ...
def copy(self) -> PersonBuilder: ...
def to_bytes(self) -> bytes: ...
def to_bytes_packed(self) -> bytes: ...
def to_segments(self) -> list[bytes]: ...
def as_reader(self) -> PersonReader: ...
@staticmethod
def write(file: IO[bytes]): ...
@staticmethod
def write_packed(file: IO[bytes]): ...
```
## Development
Uses [PDM](https://pdm-project.org/) for project management.
```bash
git clone https://github.com/capnproto/pycapnp # or your fork
cd capnp-py
pdm install # installs dependencies + dev tools
pdm run pytest # run tests
...
```
## Supported Features
| Feature | Status |
|---------|--------|
| Structs (Reader/Builder classes) | ✅ |
| Enums (Literal aliases) | ✅ |
| Structs (Reader / Builder classes) | ✅ |
| Enums (`Literal[...]` aliases) | ✅ |
| Nested types | ✅ |
| Lists (Sequence[T]) | ✅ |
| Unions (which/init) | ✅ |
| Lists (`Sequence[T]`) | ✅ |
| Unions (`which()` / `@overload init()`) | ✅ |
| Groups | ✅ |
| Self-referencing structs | ✅ |
| Interfaces (RPC) | ⏳ Phase 3 |
| Constants | ⏳ Phase 2 |
| Generics | ⏳ Phase 2 |
| Constants (primitives, enums) | ✅ |
| Generics (`TypeVar`, `Generic[T]`, brand resolution) | ✅ |
| Interfaces (RPC: Client / Server classes) | ✅ |
| Cross-file imports (relative `.` imports) | ✅ |
| Annotations (`$Cxx.name`, etc.) | ⏳ |
## Development
Uses [PDM](https://pdm-project.org/) for project management.
```bash
git clone https://github.com/NCBM/capnpc-py
cd capnp-py
pdm install # installs dependencies + dev tools
pdm run pytest # run tests
```
## License
Generated
+37 -2
View File
@@ -5,10 +5,24 @@
groups = ["default", "dev"]
strategy = ["inherit_metadata"]
lock_version = "4.5.0"
content_hash = "sha256:b8c78f3679a91711144cac4d1d558cad80470ddb4c350d7409f419b3db462e22"
content_hash = "sha256:9760696434f3c5496929539bd0d968d60732150e34e59713be5e2e8d9b1be47c"
[[metadata.targets]]
requires_python = ">=3.10"
requires_python = "~=3.10"
[[package]]
name = "basedpyright"
version = "1.39.9"
requires_python = ">=3.8"
summary = "static type checking for Python (but based)"
groups = ["dev"]
dependencies = [
"nodejs-wheel-binaries>=20.13.1",
]
files = [
{file = "basedpyright-1.39.9-py3-none-any.whl", hash = "sha256:6b0837b9eba972c71895167ab9b127e6afdbc17abc92312e3f8d15ca82a5611c"},
{file = "basedpyright-1.39.9.tar.gz", hash = "sha256:32cbea5fc8273e89df3db20daea56cb7286e419ccdfdc479c64759d2dc071901"},
]
[[package]]
name = "colorama"
@@ -149,6 +163,27 @@ files = [
{file = "markupsafe-3.0.3.tar.gz", hash = "sha256:722695808f4b6457b320fdc131280796bdceb04ab50fe1795cd540799ebe1698"},
]
[[package]]
name = "nodejs-wheel-binaries"
version = "24.16.0"
requires_python = ">=3.7"
summary = "unoffical Node.js package"
groups = ["dev"]
dependencies = [
"typing-extensions; python_version < \"3.8\"",
]
files = [
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-macosx_13_0_arm64.whl", hash = "sha256:d9f8f677dcf30e37ac244f07869726abe043f01eb0f45722b1df31cc2af7093c"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-macosx_13_0_x86_64.whl", hash = "sha256:3d0370fe7120ce9697a4f60d40480d2bd8808d9f30131458d5afc0040d4e5a51"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:85dc92bbb79c851569c5925dcc2a4c915a034efab375f99e4e7e6bbe9cca8342"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:2f3036292811514ba847b3708492644764f88a833ac425c5f55007014308ddfd"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:db8a8a76ebd2b28ecbfc9ad464baa3707241b9e050a30e2efdf6f60c0f886502"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:f1a3d8f7b4491cbbd023ba3fc4e901fcca2d9fb80d57f24ba3890de8b1dbac03"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-win_amd64.whl", hash = "sha256:bb136be9944f0662dcf1120f45193a6b75b13fac378971a95cc42c9f879a81aa"},
{file = "nodejs_wheel_binaries-24.16.0-py2.py3-none-win_arm64.whl", hash = "sha256:8308940b5edd0a50dc5267ea36ba21c9f668e83fe0d9f293937174d3a7e31c36"},
{file = "nodejs_wheel_binaries-24.16.0.tar.gz", hash = "sha256:c973cb69dc5fd16e6f6dc6e579e2c3d5534e2a1f57619dddf5ba070efa7dde37"},
]
[[package]]
name = "packaging"
version = "26.2"
+14 -1
View File
@@ -19,7 +19,8 @@ capnpc-py = "capnp_stubgen.plugin:main"
[dependency-groups]
dev = [
"pytest>=8",
"pytest>=8",
"basedpyright>=1.39.9",
]
[build-system]
@@ -31,6 +32,18 @@ distribution = true
[tool.pytest.ini_options]
testpaths = ["tests"]
markers = ["slow: slow tests (dummy.capnp, generics, interface)"]
[tool.setuptools.packages.find]
where = ["src"]
[tool.basedpyright]
typeCheckingMode = "recommended"
allowedUntypedLibraries = ["pycapnp"]
reportImportCycles = "hint"
# Downgraded to hint for capnp to_dict() dynamic data patterns
reportUnknownVariableType = "hint"
reportUnknownMemberType = "hint"
reportAny = "hint"
reportExplicitAny = "hint"
+1 -1
View File
@@ -25,7 +25,7 @@ class Emitter:
"""
def __init__(self, source_filename: str = "") -> None:
self._source_filename = source_filename
self._source_filename: str = source_filename
self._header_lines: list[str] = []
self._body_lines: list[str] = []
self._indent_level: int = 0
+10 -10
View File
@@ -11,7 +11,7 @@ import math
from typing import TYPE_CHECKING
from .type_resolver import resolve_type
from .utils import type_id_to_int
from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING:
from .emitter import Emitter
@@ -31,10 +31,10 @@ def generate_const(
globalText: str = "foobar"
voidConst: None = None
"""
node: dict = type_info.node
const_body: dict = node.get("const", {})
type_dict: dict = const_body.get("type", {})
value_dict: dict = const_body.get("value", {})
node: NodeDict = type_info.node
const_body: NodeDict = node.get("const", {})
type_dict: NodeDict = const_body.get("type", {})
value_dict: NodeDict = const_body.get("value", {})
python_type_str = _resolve_const_type(type_dict, registry)
python_value_str = _render_const_value(value_dict, type_dict, registry)
@@ -49,7 +49,7 @@ def generate_const(
# ── type resolution ────────────────────────────────────────────────────────
def _resolve_const_type(type_dict: dict, registry: TypeRegistry) -> str:
def _resolve_const_type(type_dict: NodeDict, registry: TypeRegistry) -> str:
"""Resolve a constant's type to a Python type annotation string."""
# Handle enum references
if "enum" in type_dict:
@@ -83,8 +83,8 @@ def _resolve_const_type(type_dict: dict, registry: TypeRegistry) -> str:
def _render_const_value(
value_dict: dict,
type_dict: dict,
value_dict: NodeDict,
type_dict: NodeDict,
registry: TypeRegistry,
) -> str:
"""Render a constant's value as a Python literal."""
@@ -131,7 +131,7 @@ def _render_const_value(
def _resolve_enum_value(
type_dict: dict,
type_dict: NodeDict,
ordinal: int,
registry: TypeRegistry,
) -> str:
@@ -144,7 +144,7 @@ def _resolve_enum_value(
if info is None:
return repr(ordinal)
enum_node: dict = info.node
enum_node: NodeDict = info.node
for e in enum_node.get("enum", {}).get("enumerants", []):
if e.get("ordinal") == ordinal or e.get("codeOrder") == ordinal:
return f'"{e["name"]}"'
+10 -5
View File
@@ -7,6 +7,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING
from .utils import NodeDict
if TYPE_CHECKING:
from .emitter import Emitter
from .models import TypeInfo
@@ -22,14 +24,17 @@ def generate_enum(
Gender = Literal["male", "female", "other"]
"""
node: dict = type_info.node
enum_body: dict = node.get("enum", {})
enumerants: list[dict] = enum_body.get("enumerants", [])
node: NodeDict = type_info.node
enum_body: NodeDict = node.get("enum", {})
enumerants: list[NodeDict] = enum_body.get("enumerants", [])
literals: list[str] = []
for e in enumerants:
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)
+4 -2
View File
@@ -8,6 +8,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING
from .utils import NodeDict
if TYPE_CHECKING:
from .emitter import Emitter
from .models import TypeInfo, TypeRegistry
@@ -21,7 +23,7 @@ def generate_type_vars(
Returns the list of TypeVar names for use in the ``Generic[...]`` base.
"""
node: dict = type_info.node
node: NodeDict = type_info.node
params = node.get("parameters", [])
tv_names: list[str] = []
@@ -38,7 +40,7 @@ def generate_type_vars(
def resolve_brand(
brand_dict: dict,
brand_dict: NodeDict,
registry: TypeRegistry,
) -> list[str] | None:
"""Resolve brand bindings to Python type strings.
+17 -19
View File
@@ -14,7 +14,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING
from .utils import type_id_to_int
from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING:
from .emitter import Emitter
@@ -29,11 +29,11 @@ def fixup_interface_methods(registry: TypeRegistry) -> None:
the interface by setting their ``parent_type_id`` and adjusting their
``scoped_name``.
"""
for type_info in list(registry._types.values()):
for type_info in registry.all_types():
if type_info.kind != "interface":
continue
node: dict = type_info.node
node: NodeDict = type_info.node
for method in node.get("interface", {}).get("methods", []):
mname = method["name"]
param_id = type_id_to_int(method["paramStructType"])
@@ -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(
@@ -58,13 +56,13 @@ def generate_interface(
registry: TypeRegistry,
) -> None:
"""Generate interface classes for a Cap'n Proto interface."""
node: dict = type_info.node
iface_body: dict = node.get("interface", {})
node: NodeDict = type_info.node
iface_body: NodeDict = node.get("interface", {})
name = type_info.name
methods: list[dict] = iface_body.get("methods", [])
methods: list[NodeDict] = iface_body.get("methods", [])
# Build method info list
method_infos: list[dict] = []
method_infos: list[NodeDict] = []
for method in methods:
param_id = type_id_to_int(method["paramStructType"])
result_id = type_id_to_int(method["resultStructType"])
@@ -106,7 +104,7 @@ def generate_interface(
emitter.begin_class(client_name, bases=[name])
emitter.add_method(
"_new_client",
params=["self", "server"],
params=["self", f"server: {name}Server"],
return_type=client_name,
)
emitter.end_class()
@@ -121,30 +119,30 @@ def generate_interface(
emitter.add_method(
mname,
params=["self", f"_params: {param_type}", "_context"],
return_type=None,
params=["self", f"_params: {param_type}", "_context: object"],
return_type="None",
)
emitter.end_class()
def _method_param_type(
mi: dict,
mi: NodeDict,
iface_name: str,
) -> str:
"""Get the type string for a method's params."""
param_info = mi.get("param_info")
if param_info is not None:
return f"{iface_name}.{param_info._display_name or param_info.name}Reader"
return f"{iface_name}.{param_info.display_name or param_info.name}Reader"
return f"{iface_name}.{mi['name']}ParamsReader"
def _method_result_type(
mi: dict,
mi: NodeDict,
iface_name: str,
) -> str:
"""Get the type string for a method's results."""
result_info = mi.get("result_info")
if result_info is not None:
return f"{iface_name}.{result_info._display_name or result_info.name}Reader"
return f"{iface_name}.{result_info.display_name or result_info.name}Reader"
return f"{iface_name}.{mi['name']}ResultsReader"
+89 -99
View File
@@ -16,7 +16,7 @@ from .gen_const import generate_const
from .gen_enum import generate_enum
from .gen_generic import generate_type_vars
from .type_resolver import resolve_type
from .utils import type_id_to_int
from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING:
from .emitter import Emitter
@@ -32,9 +32,9 @@ def generate_struct(
registry: TypeRegistry,
) -> None:
"""Generate the class triple for a Cap'n Proto struct."""
node: dict = type_info.node
struct_body: dict = node.get("struct", {})
name = type_info._display_name or type_info.name
node: NodeDict = type_info.node
struct_body: NodeDict = node.get("struct", {})
name = type_info.display_name or type_info.name
discriminant_count = struct_body.get("discriminantCount", 0)
nested_types = registry.get_children(type_info.type_id)
@@ -60,14 +60,9 @@ def generate_struct(
non_union_fields.append((field_name, field_type_str))
elif "group" in field:
group_type_id = type_id_to_int(field["group"]["typeId"])
group_info = registry.get(group_type_id)
if group_info is not None:
type_str = group_info.scoped_name
if disc_value != _NO_DISCRIMINANT:
union_fields.append((field_name, type_str))
else:
non_union_fields.append((field_name, type_str))
# Groups become inner classes — skip field annotation
# (the group class serves as both type and accessor)
pass
# ── class: main type ───────────────────────────────────────────────
# Generic type: generate TypeVars and Generic[T] base
@@ -76,61 +71,72 @@ 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, non_union_fields)
emitter.add_typing_import("Any")
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()
@@ -138,7 +144,7 @@ def generate_struct(
def _resolve_field_type(
type_dict: dict, registry: TypeRegistry, emitter: Emitter,
type_dict: NodeDict, registry: TypeRegistry, emitter: Emitter,
) -> str | None:
"""Resolve a field's type to a Python type annotation string."""
which = _type_which(type_dict)
@@ -148,7 +154,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":
@@ -159,7 +168,7 @@ def _resolve_field_type(
def _resolve_list_field_type(
list_dict: dict, registry: TypeRegistry, emitter: Emitter,
list_dict: NodeDict, registry: TypeRegistry, emitter: Emitter,
) -> str:
"""Resolve a ``List(T)`` field type."""
emitter.add_typing_import("Sequence", "collections.abc")
@@ -173,11 +182,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":
@@ -195,7 +202,7 @@ def _resolve_list_field_type(
return result
def _type_which(type_dict: dict) -> str | None:
def _type_which(type_dict: NodeDict) -> str | None:
"""Return the union variant key for a type dict."""
for key in ("void", "bool", "int8", "int16", "int32", "int64",
"uint8", "uint16", "uint32", "uint64",
@@ -206,33 +213,6 @@ def _type_which(type_dict: dict) -> 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 ──────────────────────────────────────────────────────
@@ -240,7 +220,6 @@ def _generate_union_methods(
emitter: Emitter, union_fields: list[tuple[str, str]],
) -> None:
emitter.add_typing_import("Literal")
emitter.add_typing_import("overload")
literal_values = ", ".join(f'"{name}"' for name, _ in union_fields)
emitter.add_method(
@@ -248,8 +227,12 @@ def _generate_union_methods(
return_type=f"Literal[{literal_values}]",
)
if len(union_fields) > 1:
emitter.add_typing_import("overload")
for field_name, type_str in union_fields:
emitter.add_decorator("overload")
if len(union_fields) > 1:
emitter.add_decorator("overload")
emitter.add_method(
"init",
params=["self", f'name: Literal["{field_name}"]'],
@@ -257,16 +240,24 @@ 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,
non_union_fields: list[tuple[str, str]],
) -> None:
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()
# Typed kwargs: name: str = ..., age: int = ...
if non_union_fields:
kwargs = [f"{name}: {typ} = ..." for name, typ in non_union_fields]
else:
emitter.add_typing_import("Any")
kwargs = ["**kwargs: Any"]
emitter.add_static_method(
"new_message", params=["**kwargs"], return_type=builder_name,
"new_message", params=kwargs, return_type=qbuilder,
)
emitter.add_blank_line()
@@ -279,7 +270,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 +281,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,10 +301,10 @@ 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")
emitter.add_static_method("write", params=["file: IO[bytes]"])
emitter.add_static_method("write", params=["file: IO[bytes]"], return_type="None")
emitter.add_blank_line()
emitter.add_static_method("write_packed", params=["file: IO[bytes]"])
emitter.add_static_method("write_packed", params=["file: IO[bytes]"], return_type="None")
+5 -1
View File
@@ -109,7 +109,7 @@ class TypeInfo:
"""The ID of the file node that contains this type (for import resolution)."""
# May be overridden for method param/result structs
_display_name: str | None = None
display_name: str | None = None
"""Override for the display name (used for method param/result structs)."""
@@ -172,5 +172,9 @@ class TypeRegistry:
def __len__(self) -> int:
return len(self._types)
def all_types(self) -> list[TypeInfo]:
"""Return all registered types (for iteration/filtering)."""
return list(self._types.values())
def __contains__(self, type_id: int) -> bool:
return type_id in self._types
+14 -21
View File
@@ -13,19 +13,16 @@ from __future__ import annotations
import os
import sys
from typing import TYPE_CHECKING
from typing import Any
from .emitter import Emitter
from .gen_const import generate_const
from .gen_enum import generate_enum
from .gen_interface import fixup_interface_methods, generate_interface
from .gen_struct import generate_struct
from .models import TypeRegistry
from .models import TypeInfo, TypeRegistry
from .schema_walker import build_type_registry
from .utils import filename_to_module_name, type_id_to_int
if TYPE_CHECKING:
from typing import Any
from .utils import NodeDict, filename_to_module_name, type_id_to_int
def main() -> int:
@@ -38,15 +35,13 @@ def main() -> int:
def _main_impl() -> int:
import capnp
# ── 1. Read CodeGeneratorRequest from stdin ────────────────────────
request = _read_request()
# Convert to dict for processing (avoids C++ DynamicStruct quirks)
request_dict = request.to_dict()
nodes: list[dict] = request_dict["nodes"]
requested_files: list[dict] = request_dict["requestedFiles"]
nodes: list[NodeDict] = request_dict["nodes"]
requested_files: list[NodeDict] = request_dict["requestedFiles"]
# ── 2. Build type registry ─────────────────────────────────────────
registry = build_type_registry(nodes, requested_files)
@@ -79,7 +74,7 @@ def _main_impl() -> int:
output_path = f"{module_name}.pyi"
with open(output_path, "w") as f:
f.write(content)
_ = f.write(content)
except Exception as exc:
print(
@@ -100,13 +95,13 @@ def _read_request() -> Any:
Uses pycapnp's ``SchemaParser`` to dynamically load ``schema.capnp``
(shipped with pycapnp), then casts the stdin message.
"""
import capnp
import capnp # pyright: ignore[reportMissingTypeStubs]
capnp_dir = os.path.dirname(capnp.__file__)
schema_path = os.path.join(capnp_dir, "schema.capnp")
import_base = os.path.dirname(capnp_dir)
parser = capnp.SchemaParser()
parser = capnp.SchemaParser() # pyright: ignore[reportAttributeAccessIssue]
schema_module = parser.load(schema_path, imports=[import_base])
return schema_module.CodeGeneratorRequest.read(sys.stdin.buffer)
@@ -130,7 +125,7 @@ def _generate_file(
_ORDER = {"enum": 0, "struct": 1, "const": 2, "interface": 3}
def _sort_key(info):
def _sort_key(info: TypeInfo) -> tuple[int, str]:
return (_ORDER.get(info.kind, 99), info.name)
top_level.sort(key=_sort_key)
@@ -163,12 +158,12 @@ def _resolve_cross_imports(
# Collect all types belonging to this file
own_types: set[int] = set()
for info in registry._types.values():
for info in registry.all_types():
if info.file_id == file_id and info.kind in ("struct", "enum", "interface"):
own_types.add(info.type_id)
# Walk types and find external references
for info in registry._types.values():
for info in registry.all_types():
if info.file_id != file_id:
continue
@@ -196,14 +191,12 @@ def _resolve_cross_imports(
all_imports: list[str] = []
for tname in sorted(type_names):
ref_info = None
for info in registry._types.values():
for info in registry.all_types():
if info.scoped_name == tname and info.file_id != file_id:
ref_info = info
break
all_imports.append(tname)
if ref_info and ref_info.kind == "struct":
all_imports.extend([f"{tname}Builder", f"{tname}Reader"])
emitter.add_import(
f"from .{module_name} import {', '.join(all_imports)}"
@@ -226,7 +219,7 @@ def _collect_type_references(info: TypeInfo, registry: TypeRegistry) -> set[int]
return refs
def _collect_field_references(field: dict, refs: set[int], registry: TypeRegistry) -> None:
def _collect_field_references(field: NodeDict, refs: set[int], _registry: TypeRegistry) -> None:
"""Recursively collect type IDs from a field's type."""
if "slot" in field:
_collect_type_references_from_dict(field["slot"].get("type", {}), refs)
@@ -234,7 +227,7 @@ def _collect_field_references(field: dict, refs: set[int], registry: TypeRegistr
refs.add(type_id_to_int(field["group"]["typeId"]))
def _collect_type_references_from_dict(type_dict: dict, refs: set[int]) -> None:
def _collect_type_references_from_dict(type_dict: NodeDict, refs: set[int]) -> None:
"""Recursively collect type IDs from a type dict."""
if "struct" in type_dict:
refs.add(type_id_to_int(type_dict["struct"]["typeId"]))
View File
+8 -13
View File
@@ -6,27 +6,22 @@ building a ``TypeRegistry`` that maps type IDs to ``TypeInfo``.
from __future__ import annotations
from typing import TYPE_CHECKING
from .models import TypeInfo, TypeRegistry
from .utils import get_display_name, get_node_kind, type_id_to_int
if TYPE_CHECKING:
from typing import Any
from .utils import NodeDict, get_display_name, get_node_kind, type_id_to_int
_GENERATABLE_KINDS = frozenset({"struct", "enum", "interface", "const"})
def build_type_registry(
nodes: list[dict],
requested_files: list[dict],
nodes: list[NodeDict],
requested_files: list[NodeDict],
) -> TypeRegistry:
"""Build a ``TypeRegistry`` from the to-dict CodeGeneratorRequest nodes."""
registry = TypeRegistry()
# Build ID → node lookup
nodes_by_id: dict[int, dict] = {}
nodes_by_id: dict[int, NodeDict] = {}
for node in nodes:
nodes_by_id[type_id_to_int(node["id"])] = node
@@ -71,8 +66,8 @@ def build_type_registry(
def _build_file_associations(
nodes_by_id: dict[int, dict],
requested_files: list[dict],
nodes_by_id: dict[int, NodeDict],
_requested_files: list[NodeDict],
registry: TypeRegistry,
) -> None:
for node_id in nodes_by_id:
@@ -82,7 +77,7 @@ def _build_file_associations(
def _find_containing_file(
node_id: int, nodes_by_id: dict[int, dict]
node_id: int, nodes_by_id: dict[int, NodeDict]
) -> int | None:
visited: set[int] = set()
current_id = node_id
@@ -109,7 +104,7 @@ def _find_containing_file(
def _build_scoped_name(
type_id: int, nodes_by_id: dict[int, dict]
type_id: int, nodes_by_id: dict[int, NodeDict]
) -> str:
parts: list[str] = []
visited: set[int] = set()
+6 -11
View File
@@ -5,13 +5,8 @@ Maps type dicts (from ``CodeGeneratorRequest.to_dict()``) to ``PythonType``.
from __future__ import annotations
from typing import TYPE_CHECKING
from .models import PythonType, TypeRegistry
from .utils import type_id_to_int
if TYPE_CHECKING:
from typing import Any
from .utils import NodeDict, type_id_to_int
# Cap'n Proto type key → Python type name
@@ -34,7 +29,7 @@ CAPNP_TO_PYTHON: dict[str, str] = {
def resolve_type(
type_dict: dict,
type_dict: NodeDict,
registry: TypeRegistry,
) -> PythonType | None:
"""Resolve a type dict to a ``PythonType``.
@@ -84,7 +79,7 @@ def resolve_type(
return PythonType(name="Any")
def _type_which(type_dict: dict) -> str | None:
def _type_which(type_dict: NodeDict) -> str | None:
"""Return the union variant key for a type dict."""
for key in CAPNP_TO_PYTHON:
if key in type_dict:
@@ -95,7 +90,7 @@ def _type_which(type_dict: dict) -> str | None:
return None
def _resolve_list_type(list_dict: dict, registry: TypeRegistry) -> PythonType:
def _resolve_list_type(list_dict: NodeDict, registry: TypeRegistry) -> PythonType:
"""Resolve a List(T) type."""
depth = 1
inner = list_dict.get("elementType", {})
@@ -129,7 +124,7 @@ def _resolve_named_type(type_id: int, registry: TypeRegistry) -> PythonType:
def _resolve_any_pointer(
ap_dict: dict, registry: TypeRegistry
ap_dict: NodeDict, registry: TypeRegistry
) -> PythonType | None:
"""Resolve an AnyPointer type."""
if "unconstrained" in ap_dict:
@@ -137,7 +132,7 @@ def _resolve_any_pointer(
if "parameter" in ap_dict:
# Look up the generic parameter name from the parent scope
scope_id = type_id_to_int(ap_dict["parameter"]["scopeId"])
param_index = ap_dict["parameter"]["parameterIndex"]
param_index: int = ap_dict["parameter"]["parameterIndex"]
info = registry.get(scope_id)
if info is not None and param_index < len(info.generic_params):
return PythonType(name=info.generic_params[param_index])
+9 -10
View File
@@ -1,6 +1,11 @@
"""Utility functions for capnp-stubgen."""
import os
from typing import Any
#: A dict from capnp's ``CodeGeneratorRequest.to_dict()``.
#: Keys are strings, values are arbitrary nested capnp data.
NodeDict = dict[str, Any]
def filename_to_module_name(filename: str) -> str:
@@ -14,18 +19,12 @@ def filename_to_module_name(filename: str) -> str:
return base + "_capnp"
def get_display_name(node: dict) -> str:
"""Extract the short display name from a node dict.
The node's ``displayName`` field contains a path like
``path/to/file.capnp:TypeName``. This removes the file prefix.
"""
name: str = node["displayName"]
prefix_len: int = node["displayNamePrefixLength"]
return name[prefix_len:]
def get_display_name(node: NodeDict) -> str:
"""Extract the short display name from a node dict."""
return str(node["displayName"])[int(node["displayNamePrefixLength"]):]
def get_node_kind(node: dict) -> str:
def get_node_kind(node: NodeDict) -> str:
"""Determine the kind of a node from its dict representation.
Returns one of: "file", "struct", "enum", "interface", "const", "annotation".
+1
View File
@@ -0,0 +1 @@
*_capnp.pyi
+58
View File
@@ -0,0 +1,58 @@
"""Regenerate .pyi stubs before running stub tests."""
from __future__ import annotations
import subprocess
import sys
from pathlib import Path
import pytest
SCHEMAS_DIR = Path(__file__).parent
TEST_SCHEMAS = [
"test_simple.capnp",
"test_nested.capnp",
"test_generics.capnp",
"test_interface.capnp",
"addressbook.capnp",
"dummy.capnp",
]
def _find_capnpc_py() -> str:
venv_bin = Path(sys.prefix) / "bin" / "capnpc-py"
if venv_bin.exists():
return str(venv_bin)
return "capnpc-py"
@pytest.fixture(scope="session", autouse=True)
def _regenerate_stubs() -> None: # pyright: ignore[reportUnusedFunction]
"""Delete old .pyi files and regenerate from .capnp schemas."""
capnpc = _find_capnpc_py()
# Clean + generate main schemas
for schema in TEST_SCHEMAS:
pyi = SCHEMAS_DIR / schema.replace(".capnp", "_capnp.pyi")
pyi.unlink(missing_ok=True)
result = subprocess.run(
["capnp", "compile", "-I.", f"-o{capnpc}", schema],
capture_output=True, text=True, cwd=str(SCHEMAS_DIR),
)
assert result.returncode == 0, (
f"capnp compile failed for {schema}:\n{result.stderr}"
)
# Multi-file schemas (consumer imports base)
multi_dir = SCHEMAS_DIR / "multi"
for pat in ("*_capnp.pyi",):
for f in multi_dir.glob(pat):
f.unlink(missing_ok=True)
for schema in ("base.capnp", "consumer.capnp"):
result = subprocess.run(
["capnp", "compile", f"-o{capnpc}", schema],
capture_output=True, text=True, cwd=str(multi_dir),
)
assert result.returncode == 0, (
f"capnp compile failed for multi/{schema}:\n{result.stderr}"
)
+105
View File
@@ -0,0 +1,105 @@
"""Verify generated .pyi stubs are consistent with pycapnp runtime.
Test file lives alongside the ``.capnp`` schemas + generated ``.pyi``
stubs. pycapnp's import hook handles runtime loading.
basedpyright type-checks this file together with the project.
"""
from __future__ import annotations
import sys
from pathlib import Path
# Ensure this directory is importable
_HERE = Path(__file__).parent
if str(_HERE) not in sys.path:
sys.path.insert(0, str(_HERE))
import capnp # pyright: ignore[reportMissingTypeStubs, reportUnusedImport] # activates import hook
import test_generics_capnp
import test_interface_capnp
import test_nested_capnp
import test_simple_capnp
class TestSimpleStub:
"""test_simple.capnp — basic struct."""
def test_new_message(self) -> None:
p = test_simple_capnp.Person.new_message(
name="Alice", age=30, email="alice@example.com"
)
assert p.name == "Alice"
assert p.age == 30
def test_field_assignment(self) -> None:
p = test_simple_capnp.Person.new_message()
p.name = "Bob"
p.age = 25
assert p.name == "Bob"
def test_to_dict(self) -> None:
p = test_simple_capnp.Person.new_message(name="Carol", age=28)
d = p.to_dict()
assert d["name"] == "Carol"
def test_from_dict(self) -> None:
p = test_simple_capnp.Person.new_message()
_ = p.from_dict({"name": "Dave", "age": 35})
assert p.name == "Dave"
def test_serialize_roundtrip(self) -> None:
p = test_simple_capnp.Person.new_message(name="Eve", age=22)
data = p.to_bytes()
p2 = test_simple_capnp.Person.from_bytes(data)
assert p2 is not None
class TestNestedStub:
"""test_nested.capnp — nested types + lists."""
def test_nested_message(self) -> None:
ab = test_nested_capnp.AddressBook.new_message()
assert ab.people is not None
class TestGenericsStub:
"""test_generics.capnp — generic structs."""
def test_holder_loads(self) -> None:
c = test_generics_capnp.Container.new_message()
assert c is not None
class TestInterfaceStub:
"""test_interface.capnp — interfaces + method params."""
def test_calcserver_loads(self) -> None:
cs = test_interface_capnp.CalcServer.new_message()
assert cs is not None
class TestAddressBookStub:
"""addressbook.capnp — classic example."""
def test_person_new_message(self) -> None:
import addressbook_capnp
p = addressbook_capnp.Person.new_message(
name="Alice", email="alice@example.com"
)
assert p.name == "Alice"
def test_addressbook_new_message(self) -> None:
import addressbook_capnp
ab = addressbook_capnp.AddressBook.new_message()
assert ab.people is not None
# dummy.capnp runtime tests skipped — its annotation imports (c.capnp,
# c++.capnp) cause pycapnp's C++ SchemaParser to SIGABRT. The generated
# .pyi passes static type-checking and `capnp compile -opy` succeeds.
# Note: multi/consumer.capnp cross-file import test skipped —
# pycapnp's SchemaParser crashes when resolving annotation imports
# across schema files. The .pyi stubs are still generated and kept.
+2 -2
View File
@@ -85,11 +85,11 @@ class TestTypeRegistry:
assert reg.get_or_raise(1) is info
import pytest
with pytest.raises(KeyError):
reg.get_or_raise(999)
_ = reg.get_or_raise(999)
def test_contains(self):
reg = TypeRegistry()
info = TypeInfo(1, "Foo", "Foo", {}, None, "struct")
info: TypeInfo = TypeInfo(1, "Foo", "Foo", {}, None, "struct")
reg.register(info)
assert 1 in reg
assert 2 not in reg
-1
View File
@@ -1,6 +1,5 @@
"""Tests for capnp_stubgen.utils."""
import pytest
from capnp_stubgen.utils import (
filename_to_module_name,
get_display_name,