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` | | Arch Linux | `sudo pacman -S capnproto` |
| Ubuntu/Debian | `sudo apt install capnproto` | | Ubuntu/Debian | `sudo apt install capnproto` |
| macOS | `brew install capnp` | | macOS | `brew install capnp` |
| Windows | Download from [capnproto.org](https://capnproto.org) | | Windows | Download from [capnproto.org](https://capnproto.org/install.html) |
## Installation ## 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. 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 ```capnp
struct Person { struct Person {
name @0 :Text; name @0 :Text;
age @1 :Int32; 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 ```python
"""Auto-generated type stub for myschema.capnp""" """Auto-generated type stub for addressbook.capnp"""
from __future__ import annotations from __future__ import annotations
from collections.abc import Iterator, Sequence from collections.abc import Iterator, Sequence
from contextlib import contextmanager from contextlib import contextmanager
from typing import IO, Literal from typing import IO, Literal
Gender = Literal["male", "female"]
class Person: class Person:
name: str name: str
age: int age: int
phones: Sequence[Person.PhoneNumber | Person.PhoneNumberBuilder | Person.PhoneNumberReader]
class PhoneNumber:
Type = Literal["mobile", "home", "work"]
type: Type
number: str
...
@staticmethod @staticmethod
def new_message(**kwargs) -> PersonBuilder: ... 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: ... def to_dict(self) -> dict: ...
class PersonReader(Person): class PersonReader(Person):
def as_builder(self) -> PersonBuilder: ... def as_builder(self) -> PersonBuilder: ...
class PersonBuilder(Person): class PersonBuilder(Person):
@staticmethod
def from_dict(dictionary: dict) -> PersonBuilder: ...
def copy(self) -> PersonBuilder: ...
def to_bytes(self) -> bytes: ... def to_bytes(self) -> bytes: ...
def to_bytes_packed(self) -> bytes: ...
def to_segments(self) -> list[bytes]: ...
def as_reader(self) -> PersonReader: ... 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 ## Supported Features
| Feature | Status | | Feature | Status |
|---------|--------| |---------|--------|
| Structs (Reader/Builder classes) | ✅ | | Structs (Reader / Builder classes) | ✅ |
| Enums (Literal aliases) | ✅ | | Enums (`Literal[...]` aliases) | ✅ |
| Nested types | ✅ | | Nested types | ✅ |
| Lists (Sequence[T]) | ✅ | | Lists (`Sequence[T]`) | ✅ |
| Unions (which/init) | ✅ | | Unions (`which()` / `@overload init()`) | ✅ |
| Groups | ✅ |
| Self-referencing structs | ✅ | | Self-referencing structs | ✅ |
| Interfaces (RPC) | ⏳ Phase 3 | | Constants (primitives, enums) | ✅ |
| Constants | ⏳ Phase 2 | | Generics (`TypeVar`, `Generic[T]`, brand resolution) | ✅ |
| Generics | ⏳ Phase 2 | | 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 ## License
Generated
+37 -2
View File
@@ -5,10 +5,24 @@
groups = ["default", "dev"] groups = ["default", "dev"]
strategy = ["inherit_metadata"] strategy = ["inherit_metadata"]
lock_version = "4.5.0" lock_version = "4.5.0"
content_hash = "sha256:b8c78f3679a91711144cac4d1d558cad80470ddb4c350d7409f419b3db462e22" content_hash = "sha256:9760696434f3c5496929539bd0d968d60732150e34e59713be5e2e8d9b1be47c"
[[metadata.targets]] [[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]] [[package]]
name = "colorama" name = "colorama"
@@ -149,6 +163,27 @@ files = [
{file = "markupsafe-3.0.3.tar.gz", hash = "sha256:722695808f4b6457b320fdc131280796bdceb04ab50fe1795cd540799ebe1698"}, {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]] [[package]]
name = "packaging" name = "packaging"
version = "26.2" version = "26.2"
+14 -1
View File
@@ -19,7 +19,8 @@ capnpc-py = "capnp_stubgen.plugin:main"
[dependency-groups] [dependency-groups]
dev = [ dev = [
"pytest>=8", "pytest>=8",
"basedpyright>=1.39.9",
] ]
[build-system] [build-system]
@@ -31,6 +32,18 @@ distribution = true
[tool.pytest.ini_options] [tool.pytest.ini_options]
testpaths = ["tests"] testpaths = ["tests"]
markers = ["slow: slow tests (dummy.capnp, generics, interface)"]
[tool.setuptools.packages.find] [tool.setuptools.packages.find]
where = ["src"] 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: def __init__(self, source_filename: str = "") -> None:
self._source_filename = source_filename self._source_filename: str = source_filename
self._header_lines: list[str] = [] self._header_lines: list[str] = []
self._body_lines: list[str] = [] self._body_lines: list[str] = []
self._indent_level: int = 0 self._indent_level: int = 0
+10 -10
View File
@@ -11,7 +11,7 @@ import math
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .type_resolver import resolve_type from .type_resolver import resolve_type
from .utils import type_id_to_int from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING: if TYPE_CHECKING:
from .emitter import Emitter from .emitter import Emitter
@@ -31,10 +31,10 @@ def generate_const(
globalText: str = "foobar" globalText: str = "foobar"
voidConst: None = None voidConst: None = None
""" """
node: dict = type_info.node node: NodeDict = type_info.node
const_body: dict = node.get("const", {}) const_body: NodeDict = node.get("const", {})
type_dict: dict = const_body.get("type", {}) type_dict: NodeDict = const_body.get("type", {})
value_dict: dict = const_body.get("value", {}) value_dict: NodeDict = const_body.get("value", {})
python_type_str = _resolve_const_type(type_dict, registry) python_type_str = _resolve_const_type(type_dict, registry)
python_value_str = _render_const_value(value_dict, type_dict, registry) python_value_str = _render_const_value(value_dict, type_dict, registry)
@@ -49,7 +49,7 @@ def generate_const(
# ── type resolution ──────────────────────────────────────────────────────── # ── 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.""" """Resolve a constant's type to a Python type annotation string."""
# Handle enum references # Handle enum references
if "enum" in type_dict: if "enum" in type_dict:
@@ -83,8 +83,8 @@ def _resolve_const_type(type_dict: dict, registry: TypeRegistry) -> str:
def _render_const_value( def _render_const_value(
value_dict: dict, value_dict: NodeDict,
type_dict: dict, type_dict: NodeDict,
registry: TypeRegistry, registry: TypeRegistry,
) -> str: ) -> str:
"""Render a constant's value as a Python literal.""" """Render a constant's value as a Python literal."""
@@ -131,7 +131,7 @@ def _render_const_value(
def _resolve_enum_value( def _resolve_enum_value(
type_dict: dict, type_dict: NodeDict,
ordinal: int, ordinal: int,
registry: TypeRegistry, registry: TypeRegistry,
) -> str: ) -> str:
@@ -144,7 +144,7 @@ def _resolve_enum_value(
if info is None: if info is None:
return repr(ordinal) return repr(ordinal)
enum_node: dict = info.node enum_node: NodeDict = info.node
for e in enum_node.get("enum", {}).get("enumerants", []): for e in enum_node.get("enum", {}).get("enumerants", []):
if e.get("ordinal") == ordinal or e.get("codeOrder") == ordinal: if e.get("ordinal") == ordinal or e.get("codeOrder") == ordinal:
return f'"{e["name"]}"' return f'"{e["name"]}"'
+10 -5
View File
@@ -7,6 +7,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .utils import NodeDict
if TYPE_CHECKING: if TYPE_CHECKING:
from .emitter import Emitter from .emitter import Emitter
from .models import TypeInfo from .models import TypeInfo
@@ -22,14 +24,17 @@ def generate_enum(
Gender = Literal["male", "female", "other"] Gender = Literal["male", "female", "other"]
""" """
node: dict = type_info.node node: NodeDict = type_info.node
enum_body: dict = node.get("enum", {}) enum_body: NodeDict = node.get("enum", {})
enumerants: list[dict] = enum_body.get("enumerants", []) enumerants: list[NodeDict] = enum_body.get("enumerants", [])
literals: list[str] = [] literals: list[str] = []
for e in enumerants: for e in enumerants:
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)
+4 -2
View File
@@ -8,6 +8,8 @@ from __future__ import annotations
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
from .utils import NodeDict
if TYPE_CHECKING: if TYPE_CHECKING:
from .emitter import Emitter from .emitter import Emitter
from .models import TypeInfo, TypeRegistry 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. 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", []) params = node.get("parameters", [])
tv_names: list[str] = [] tv_names: list[str] = []
@@ -38,7 +40,7 @@ def generate_type_vars(
def resolve_brand( def resolve_brand(
brand_dict: dict, brand_dict: NodeDict,
registry: TypeRegistry, registry: TypeRegistry,
) -> list[str] | None: ) -> list[str] | None:
"""Resolve brand bindings to Python type strings. """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 typing import TYPE_CHECKING
from .utils import type_id_to_int from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING: if TYPE_CHECKING:
from .emitter import Emitter 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 the interface by setting their ``parent_type_id`` and adjusting their
``scoped_name``. ``scoped_name``.
""" """
for type_info in list(registry._types.values()): for type_info in registry.all_types():
if type_info.kind != "interface": if type_info.kind != "interface":
continue continue
node: dict = type_info.node node: NodeDict = type_info.node
for method in node.get("interface", {}).get("methods", []): for method in node.get("interface", {}).get("methods", []):
mname = method["name"] mname = method["name"]
param_id = type_id_to_int(method["paramStructType"]) param_id = type_id_to_int(method["paramStructType"])
@@ -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(
@@ -58,13 +56,13 @@ def generate_interface(
registry: TypeRegistry, registry: TypeRegistry,
) -> None: ) -> None:
"""Generate interface classes for a Cap'n Proto interface.""" """Generate interface classes for a Cap'n Proto interface."""
node: dict = type_info.node node: NodeDict = type_info.node
iface_body: dict = node.get("interface", {}) iface_body: NodeDict = node.get("interface", {})
name = type_info.name name = type_info.name
methods: list[dict] = iface_body.get("methods", []) methods: list[NodeDict] = iface_body.get("methods", [])
# Build method info list # Build method info list
method_infos: list[dict] = [] method_infos: list[NodeDict] = []
for method in methods: for method in methods:
param_id = type_id_to_int(method["paramStructType"]) param_id = type_id_to_int(method["paramStructType"])
result_id = type_id_to_int(method["resultStructType"]) result_id = type_id_to_int(method["resultStructType"])
@@ -106,7 +104,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", f"server: {name}Server"],
return_type=client_name, return_type=client_name,
) )
emitter.end_class() emitter.end_class()
@@ -121,30 +119,30 @@ 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: object"],
return_type=None, return_type="None",
) )
emitter.end_class() emitter.end_class()
def _method_param_type( def _method_param_type(
mi: dict, mi: NodeDict,
iface_name: str, iface_name: str,
) -> str: ) -> str:
"""Get the type string for a method's params.""" """Get the type string for a method's params."""
param_info = mi.get("param_info") param_info = mi.get("param_info")
if param_info is not None: 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" return f"{iface_name}.{mi['name']}ParamsReader"
def _method_result_type( def _method_result_type(
mi: dict, mi: NodeDict,
iface_name: str, iface_name: str,
) -> str: ) -> str:
"""Get the type string for a method's results.""" """Get the type string for a method's results."""
result_info = mi.get("result_info") result_info = mi.get("result_info")
if result_info is not None: 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" 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_enum import generate_enum
from .gen_generic import generate_type_vars from .gen_generic import generate_type_vars
from .type_resolver import resolve_type from .type_resolver import resolve_type
from .utils import type_id_to_int from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING: if TYPE_CHECKING:
from .emitter import Emitter from .emitter import Emitter
@@ -32,9 +32,9 @@ def generate_struct(
registry: TypeRegistry, registry: TypeRegistry,
) -> None: ) -> None:
"""Generate the class triple for a Cap'n Proto struct.""" """Generate the class triple for a Cap'n Proto struct."""
node: dict = type_info.node node: NodeDict = type_info.node
struct_body: dict = node.get("struct", {}) struct_body: NodeDict = node.get("struct", {})
name = type_info._display_name or type_info.name name = type_info.display_name or type_info.name
discriminant_count = struct_body.get("discriminantCount", 0) discriminant_count = struct_body.get("discriminantCount", 0)
nested_types = registry.get_children(type_info.type_id) 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)) non_union_fields.append((field_name, field_type_str))
elif "group" in field: elif "group" in field:
group_type_id = type_id_to_int(field["group"]["typeId"]) # Groups become inner classes — skip field annotation
group_info = registry.get(group_type_id) # (the group class serves as both type and accessor)
if group_info is not None: pass
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))
# ── class: main type ─────────────────────────────────────────────── # ── class: main type ───────────────────────────────────────────────
# Generic type: generate TypeVars and Generic[T] base # Generic type: generate TypeVars and Generic[T] base
@@ -76,61 +71,72 @@ 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, non_union_fields)
emitter.add_method("to_dict", params=["self"], return_type="dict") emitter.add_typing_import("Any")
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()
@@ -138,7 +144,7 @@ def generate_struct(
def _resolve_field_type( def _resolve_field_type(
type_dict: dict, registry: TypeRegistry, emitter: Emitter, type_dict: NodeDict, registry: TypeRegistry, emitter: Emitter,
) -> str | None: ) -> str | None:
"""Resolve a field's type to a Python type annotation string.""" """Resolve a field's type to a Python type annotation string."""
which = _type_which(type_dict) which = _type_which(type_dict)
@@ -148,7 +154,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":
@@ -159,7 +168,7 @@ def _resolve_field_type(
def _resolve_list_field_type( def _resolve_list_field_type(
list_dict: dict, registry: TypeRegistry, emitter: Emitter, list_dict: NodeDict, registry: TypeRegistry, emitter: Emitter,
) -> str: ) -> str:
"""Resolve a ``List(T)`` field type.""" """Resolve a ``List(T)`` field type."""
emitter.add_typing_import("Sequence", "collections.abc") emitter.add_typing_import("Sequence", "collections.abc")
@@ -173,11 +182,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":
@@ -195,7 +202,7 @@ def _resolve_list_field_type(
return result 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.""" """Return the union variant key for a type dict."""
for key in ("void", "bool", "int8", "int16", "int32", "int64", for key in ("void", "bool", "int8", "int16", "int32", "int64",
"uint8", "uint16", "uint32", "uint64", "uint8", "uint16", "uint32", "uint64",
@@ -206,33 +213,6 @@ def _type_which(type_dict: dict) -> 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 ──────────────────────────────────────────────────────
@@ -240,7 +220,6 @@ def _generate_union_methods(
emitter: Emitter, union_fields: list[tuple[str, str]], emitter: Emitter, union_fields: list[tuple[str, str]],
) -> None: ) -> None:
emitter.add_typing_import("Literal") emitter.add_typing_import("Literal")
emitter.add_typing_import("overload")
literal_values = ", ".join(f'"{name}"' for name, _ in union_fields) literal_values = ", ".join(f'"{name}"' for name, _ in union_fields)
emitter.add_method( emitter.add_method(
@@ -248,8 +227,12 @@ def _generate_union_methods(
return_type=f"Literal[{literal_values}]", return_type=f"Literal[{literal_values}]",
) )
if len(union_fields) > 1:
emitter.add_typing_import("overload")
for field_name, type_str in union_fields: for field_name, type_str in union_fields:
emitter.add_decorator("overload") if len(union_fields) > 1:
emitter.add_decorator("overload")
emitter.add_method( emitter.add_method(
"init", "init",
params=["self", f'name: Literal["{field_name}"]'], params=["self", f'name: Literal["{field_name}"]'],
@@ -257,16 +240,24 @@ 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,
non_union_fields: list[tuple[str, str]],
) -> None:
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()
# 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( emitter.add_static_method(
"new_message", params=["**kwargs"], return_type=builder_name, "new_message", params=kwargs, return_type=qbuilder,
) )
emitter.add_blank_line() emitter.add_blank_line()
@@ -279,7 +270,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 +281,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,10 +301,10 @@ 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")
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_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).""" """The ID of the file node that contains this type (for import resolution)."""
# May be overridden for method param/result structs # 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).""" """Override for the display name (used for method param/result structs)."""
@@ -172,5 +172,9 @@ class TypeRegistry:
def __len__(self) -> int: def __len__(self) -> int:
return len(self._types) 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: def __contains__(self, type_id: int) -> bool:
return type_id in self._types return type_id in self._types
+14 -21
View File
@@ -13,19 +13,16 @@ from __future__ import annotations
import os import os
import sys import sys
from typing import TYPE_CHECKING from typing import Any
from .emitter import Emitter from .emitter import Emitter
from .gen_const import generate_const from .gen_const import generate_const
from .gen_enum import generate_enum from .gen_enum import generate_enum
from .gen_interface import fixup_interface_methods, generate_interface from .gen_interface import fixup_interface_methods, generate_interface
from .gen_struct import generate_struct from .gen_struct import generate_struct
from .models import TypeRegistry from .models import TypeInfo, TypeRegistry
from .schema_walker import build_type_registry from .schema_walker import build_type_registry
from .utils import filename_to_module_name, type_id_to_int from .utils import NodeDict, filename_to_module_name, type_id_to_int
if TYPE_CHECKING:
from typing import Any
def main() -> int: def main() -> int:
@@ -38,15 +35,13 @@ def main() -> int:
def _main_impl() -> int: def _main_impl() -> int:
import capnp
# ── 1. Read CodeGeneratorRequest from stdin ──────────────────────── # ── 1. Read CodeGeneratorRequest from stdin ────────────────────────
request = _read_request() request = _read_request()
# Convert to dict for processing (avoids C++ DynamicStruct quirks) # Convert to dict for processing (avoids C++ DynamicStruct quirks)
request_dict = request.to_dict() request_dict = request.to_dict()
nodes: list[dict] = request_dict["nodes"] nodes: list[NodeDict] = request_dict["nodes"]
requested_files: list[dict] = request_dict["requestedFiles"] requested_files: list[NodeDict] = request_dict["requestedFiles"]
# ── 2. Build type registry ───────────────────────────────────────── # ── 2. Build type registry ─────────────────────────────────────────
registry = build_type_registry(nodes, requested_files) registry = build_type_registry(nodes, requested_files)
@@ -79,7 +74,7 @@ def _main_impl() -> int:
output_path = f"{module_name}.pyi" output_path = f"{module_name}.pyi"
with open(output_path, "w") as f: with open(output_path, "w") as f:
f.write(content) _ = f.write(content)
except Exception as exc: except Exception as exc:
print( print(
@@ -100,13 +95,13 @@ def _read_request() -> Any:
Uses pycapnp's ``SchemaParser`` to dynamically load ``schema.capnp`` Uses pycapnp's ``SchemaParser`` to dynamically load ``schema.capnp``
(shipped with pycapnp), then casts the stdin message. (shipped with pycapnp), then casts the stdin message.
""" """
import capnp import capnp # pyright: ignore[reportMissingTypeStubs]
capnp_dir = os.path.dirname(capnp.__file__) capnp_dir = os.path.dirname(capnp.__file__)
schema_path = os.path.join(capnp_dir, "schema.capnp") schema_path = os.path.join(capnp_dir, "schema.capnp")
import_base = os.path.dirname(capnp_dir) 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]) schema_module = parser.load(schema_path, imports=[import_base])
return schema_module.CodeGeneratorRequest.read(sys.stdin.buffer) return schema_module.CodeGeneratorRequest.read(sys.stdin.buffer)
@@ -130,7 +125,7 @@ def _generate_file(
_ORDER = {"enum": 0, "struct": 1, "const": 2, "interface": 3} _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) return (_ORDER.get(info.kind, 99), info.name)
top_level.sort(key=_sort_key) top_level.sort(key=_sort_key)
@@ -163,12 +158,12 @@ def _resolve_cross_imports(
# Collect all types belonging to this file # Collect all types belonging to this file
own_types: set[int] = set() 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"): if info.file_id == file_id and info.kind in ("struct", "enum", "interface"):
own_types.add(info.type_id) own_types.add(info.type_id)
# Walk types and find external references # Walk types and find external references
for info in registry._types.values(): for info in registry.all_types():
if info.file_id != file_id: if info.file_id != file_id:
continue continue
@@ -196,14 +191,12 @@ def _resolve_cross_imports(
all_imports: list[str] = [] all_imports: list[str] = []
for tname in sorted(type_names): for tname in sorted(type_names):
ref_info = None 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: if info.scoped_name == tname and info.file_id != file_id:
ref_info = info ref_info = info
break break
all_imports.append(tname) all_imports.append(tname)
if ref_info and ref_info.kind == "struct":
all_imports.extend([f"{tname}Builder", f"{tname}Reader"])
emitter.add_import( emitter.add_import(
f"from .{module_name} import {', '.join(all_imports)}" f"from .{module_name} import {', '.join(all_imports)}"
@@ -226,7 +219,7 @@ def _collect_type_references(info: TypeInfo, registry: TypeRegistry) -> set[int]
return refs 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.""" """Recursively collect type IDs from a field's type."""
if "slot" in field: if "slot" in field:
_collect_type_references_from_dict(field["slot"].get("type", {}), refs) _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"])) 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.""" """Recursively collect type IDs from a type dict."""
if "struct" in type_dict: if "struct" in type_dict:
refs.add(type_id_to_int(type_dict["struct"]["typeId"])) 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 __future__ import annotations
from typing import TYPE_CHECKING
from .models import TypeInfo, TypeRegistry from .models import TypeInfo, TypeRegistry
from .utils import get_display_name, get_node_kind, type_id_to_int from .utils import NodeDict, get_display_name, get_node_kind, type_id_to_int
if TYPE_CHECKING:
from typing import Any
_GENERATABLE_KINDS = frozenset({"struct", "enum", "interface", "const"}) _GENERATABLE_KINDS = frozenset({"struct", "enum", "interface", "const"})
def build_type_registry( def build_type_registry(
nodes: list[dict], nodes: list[NodeDict],
requested_files: list[dict], requested_files: list[NodeDict],
) -> TypeRegistry: ) -> TypeRegistry:
"""Build a ``TypeRegistry`` from the to-dict CodeGeneratorRequest nodes.""" """Build a ``TypeRegistry`` from the to-dict CodeGeneratorRequest nodes."""
registry = TypeRegistry() registry = TypeRegistry()
# Build ID → node lookup # Build ID → node lookup
nodes_by_id: dict[int, dict] = {} nodes_by_id: dict[int, NodeDict] = {}
for node in nodes: for node in nodes:
nodes_by_id[type_id_to_int(node["id"])] = node nodes_by_id[type_id_to_int(node["id"])] = node
@@ -71,8 +66,8 @@ def build_type_registry(
def _build_file_associations( def _build_file_associations(
nodes_by_id: dict[int, dict], nodes_by_id: dict[int, NodeDict],
requested_files: list[dict], _requested_files: list[NodeDict],
registry: TypeRegistry, registry: TypeRegistry,
) -> None: ) -> None:
for node_id in nodes_by_id: for node_id in nodes_by_id:
@@ -82,7 +77,7 @@ def _build_file_associations(
def _find_containing_file( def _find_containing_file(
node_id: int, nodes_by_id: dict[int, dict] node_id: int, nodes_by_id: dict[int, NodeDict]
) -> int | None: ) -> int | None:
visited: set[int] = set() visited: set[int] = set()
current_id = node_id current_id = node_id
@@ -109,7 +104,7 @@ def _find_containing_file(
def _build_scoped_name( def _build_scoped_name(
type_id: int, nodes_by_id: dict[int, dict] type_id: int, nodes_by_id: dict[int, NodeDict]
) -> str: ) -> str:
parts: list[str] = [] parts: list[str] = []
visited: set[int] = set() 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 __future__ import annotations
from typing import TYPE_CHECKING
from .models import PythonType, TypeRegistry from .models import PythonType, TypeRegistry
from .utils import type_id_to_int from .utils import NodeDict, type_id_to_int
if TYPE_CHECKING:
from typing import Any
# Cap'n Proto type key → Python type name # Cap'n Proto type key → Python type name
@@ -34,7 +29,7 @@ CAPNP_TO_PYTHON: dict[str, str] = {
def resolve_type( def resolve_type(
type_dict: dict, type_dict: NodeDict,
registry: TypeRegistry, registry: TypeRegistry,
) -> PythonType | None: ) -> PythonType | None:
"""Resolve a type dict to a ``PythonType``. """Resolve a type dict to a ``PythonType``.
@@ -84,7 +79,7 @@ def resolve_type(
return PythonType(name="Any") 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.""" """Return the union variant key for a type dict."""
for key in CAPNP_TO_PYTHON: for key in CAPNP_TO_PYTHON:
if key in type_dict: if key in type_dict:
@@ -95,7 +90,7 @@ def _type_which(type_dict: dict) -> str | None:
return 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.""" """Resolve a List(T) type."""
depth = 1 depth = 1
inner = list_dict.get("elementType", {}) inner = list_dict.get("elementType", {})
@@ -129,7 +124,7 @@ def _resolve_named_type(type_id: int, registry: TypeRegistry) -> PythonType:
def _resolve_any_pointer( def _resolve_any_pointer(
ap_dict: dict, registry: TypeRegistry ap_dict: NodeDict, registry: TypeRegistry
) -> PythonType | None: ) -> PythonType | None:
"""Resolve an AnyPointer type.""" """Resolve an AnyPointer type."""
if "unconstrained" in ap_dict: if "unconstrained" in ap_dict:
@@ -137,7 +132,7 @@ def _resolve_any_pointer(
if "parameter" in ap_dict: if "parameter" in ap_dict:
# Look up the generic parameter name from the parent scope # Look up the generic parameter name from the parent scope
scope_id = type_id_to_int(ap_dict["parameter"]["scopeId"]) 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) info = registry.get(scope_id)
if info is not None and param_index < len(info.generic_params): if info is not None and param_index < len(info.generic_params):
return PythonType(name=info.generic_params[param_index]) return PythonType(name=info.generic_params[param_index])
+9 -10
View File
@@ -1,6 +1,11 @@
"""Utility functions for capnp-stubgen.""" """Utility functions for capnp-stubgen."""
import os 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: def filename_to_module_name(filename: str) -> str:
@@ -14,18 +19,12 @@ def filename_to_module_name(filename: str) -> str:
return base + "_capnp" return base + "_capnp"
def get_display_name(node: dict) -> str: def get_display_name(node: NodeDict) -> str:
"""Extract the short display name from a node dict. """Extract the short display name from a node dict."""
return str(node["displayName"])[int(node["displayNamePrefixLength"]):]
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_node_kind(node: dict) -> str: def get_node_kind(node: NodeDict) -> str:
"""Determine the kind of a node from its dict representation. """Determine the kind of a node from its dict representation.
Returns one of: "file", "struct", "enum", "interface", "const", "annotation". 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 assert reg.get_or_raise(1) is info
import pytest import pytest
with pytest.raises(KeyError): with pytest.raises(KeyError):
reg.get_or_raise(999) _ = reg.get_or_raise(999)
def test_contains(self): def test_contains(self):
reg = TypeRegistry() reg = TypeRegistry()
info = TypeInfo(1, "Foo", "Foo", {}, None, "struct") info: TypeInfo = TypeInfo(1, "Foo", "Foo", {}, None, "struct")
reg.register(info) reg.register(info)
assert 1 in reg assert 1 in reg
assert 2 not in reg assert 2 not in reg
-1
View File
@@ -1,6 +1,5 @@
"""Tests for capnp_stubgen.utils.""" """Tests for capnp_stubgen.utils."""
import pytest
from capnp_stubgen.utils import ( from capnp_stubgen.utils import (
filename_to_module_name, filename_to_module_name,
get_display_name, get_display_name,