use std::fs;
use std::path::Path;
use anyhow::Result;
use rspyts::ir::{BufferElement, FieldDef, ScalarValue, Target, TypeDef, TypeRef, TypeShape};
use serde_json::Value;
use crate::config::PythonConfig;
use crate::resolve::ResolvedContract;
use super::util::{
TypeNames, error_definition, ordered_types, python_doc, python_string_literal,
python_tag_field, python_tagged_variant_name, type_allows_null, type_definition, type_names,
};
use super::{write, write_json};
pub fn emit(
root: &Path,
config: &PythonConfig,
contract: &ResolvedContract,
fingerprint: &str,
) -> Result<()> {
let manifest = &contract.manifest;
let python_root = root.join("python");
let mut package = python_root.clone();
let package_parts = config.package.split('.').collect::<Vec<_>>();
for part in package_parts {
package.push(part);
fs::create_dir_all(&package)?;
}
let names = type_names(contract);
write(&package.join("models.py"), &models(contract, &names))?;
write(&package.join("codecs.py"), &codecs(contract))?;
write(&package.join("errors.py"), &errors(contract))?;
write(&package.join("functions.py"), &functions(contract, &names))?;
write(&package.join("resources.py"), &resources(contract, &names))?;
write(
&package.join("constants.py"),
&constants(contract, fingerprint, &names),
)?;
write(&package.join("__init__.py"), &init_module(contract))?;
let typed_marker = package.join("py.typed");
if !typed_marker.exists() {
write(&typed_marker, "")?;
}
write_json(
&package.join("contract.json"),
&serde_json::json!({
"schemaVersion": crate::LOCK_VERSION,
"fingerprint": fingerprint,
"hosts": contract.hosts,
"dependencies": contract.dependencies,
"manifest": manifest,
}),
)?;
Ok(())
}
fn models(contract: &ResolvedContract, names: &TypeNames) -> String {
let manifest = &contract.manifest;
let mut output = String::from(concat!(
"# Generated by rspyts 0.4. Do not edit.\n",
"from __future__ import annotations\n\n",
"import builtins as __rspyts_builtins__\n",
"import datetime as __rspyts_datetime__\n",
"import math as __rspyts_math__\n",
"from enum import Enum as __rspyts_Enum__\n",
"from typing import Annotated as __rspyts_Annotated__\n",
"from typing import Literal as __rspyts_Literal__\n",
"from typing import TypeAlias as __rspyts_TypeAlias__\n\n",
"from pydantic import AfterValidator as __rspyts_AfterValidator__\n",
"from pydantic import AwareDatetime as __rspyts_AwareDatetime__\n",
"from pydantic import BaseModel as __rspyts_BaseModel__\n",
"from pydantic import BeforeValidator as __rspyts_BeforeValidator__\n",
"from pydantic import ConfigDict as __rspyts_ConfigDict__\n",
"from pydantic import Field as __rspyts_Field__\n",
"from pydantic import JsonValue as __rspyts_PydanticJsonValue__\n\n",
"def __rspyts_validate_json_value__(\n",
" __rspyts_value: __rspyts_PydanticJsonValue__,\n",
" __rspyts_path: str = \"$\",\n",
") -> __rspyts_PydanticJsonValue__:\n",
" if (\n",
" __rspyts_builtins__.isinstance(__rspyts_value, __rspyts_builtins__.bool)\n",
" or __rspyts_value is None\n",
" or __rspyts_builtins__.isinstance(__rspyts_value, __rspyts_builtins__.str)\n",
" ):\n",
" return __rspyts_value\n",
" if __rspyts_builtins__.isinstance(__rspyts_value, __rspyts_builtins__.int):\n",
" if __rspyts_value < -9_007_199_254_740_991 or __rspyts_value > 9_007_199_254_740_991:\n",
" raise __rspyts_builtins__.ValueError(\n",
" f\"JSON integer at {__rspyts_path} is outside JavaScript's safe integer range\"\n",
" )\n",
" return __rspyts_value\n",
" if __rspyts_builtins__.isinstance(__rspyts_value, __rspyts_builtins__.float):\n",
" if not __rspyts_math__.isfinite(__rspyts_value):\n",
" raise __rspyts_builtins__.ValueError(\n",
" f\"JSON number at {__rspyts_path} must be finite\"\n",
" )\n",
" if __rspyts_value.is_integer() and (\n",
" __rspyts_value < -9_007_199_254_740_991\n",
" or __rspyts_value > 9_007_199_254_740_991\n",
" ):\n",
" raise __rspyts_builtins__.ValueError(\n",
" f\"JSON integer at {__rspyts_path} is outside JavaScript's safe integer range\"\n",
" )\n",
" return __rspyts_value\n",
" if __rspyts_builtins__.isinstance(__rspyts_value, __rspyts_builtins__.list):\n",
" for __rspyts_index, __rspyts_item in __rspyts_builtins__.enumerate(__rspyts_value):\n",
" __rspyts_validate_json_value__(\n",
" __rspyts_item, f\"{__rspyts_path}[{__rspyts_index}]\"\n",
" )\n",
" return __rspyts_value\n",
" for __rspyts_key, __rspyts_item in __rspyts_value.items():\n",
" __rspyts_validate_json_value__(\n",
" __rspyts_item, f\"{__rspyts_path}[{__rspyts_key!r}]\"\n",
" )\n",
" return __rspyts_value\n\n",
"def __rspyts_validate_datetime_input__(__rspyts_value: object) -> object:\n",
" if __rspyts_value is None or __rspyts_builtins__.isinstance(\n",
" __rspyts_value, (__rspyts_builtins__.str, __rspyts_datetime__.datetime)\n",
" ):\n",
" return __rspyts_value\n",
" raise __rspyts_builtins__.ValueError(\n",
" \"datetime input must be an ISO 8601 string or datetime\"\n",
" )\n\n",
"def __rspyts_validate_string_enum_input__(__rspyts_value: object) -> object:\n",
" if __rspyts_value is None or __rspyts_builtins__.isinstance(\n",
" __rspyts_value, __rspyts_builtins__.str\n",
" ):\n",
" return __rspyts_value\n",
" raise __rspyts_builtins__.ValueError(\n",
" \"string enum input must be a string or enum member\"\n",
" )\n\n",
"def __rspyts_validate_integer_literal_input__(__rspyts_value: object) -> object:\n",
" if __rspyts_value is None or __rspyts_builtins__.type(__rspyts_value) is __rspyts_builtins__.int:\n",
" return __rspyts_value\n",
" raise __rspyts_builtins__.ValueError(\n",
" \"integer literal input must be an exact integer\"\n",
" )\n\n",
"def __rspyts_validate_boolean_literal_input__(__rspyts_value: object) -> object:\n",
" if __rspyts_value is None or __rspyts_builtins__.type(__rspyts_value) is __rspyts_builtins__.bool:\n",
" return __rspyts_value\n",
" raise __rspyts_builtins__.ValueError(\n",
" \"boolean literal input must be an exact boolean\"\n",
" )\n\n",
"def __rspyts_validate_string_literal_input__(__rspyts_value: object) -> object:\n",
" if __rspyts_value is None or __rspyts_builtins__.isinstance(\n",
" __rspyts_value, __rspyts_builtins__.str\n",
" ):\n",
" return __rspyts_value\n",
" raise __rspyts_builtins__.ValueError(\n",
" \"string literal input must be an exact string\"\n",
" )\n\n",
"JsonValue: __rspyts_TypeAlias__ = __rspyts_Annotated__[\n",
" __rspyts_PydanticJsonValue__,\n",
" __rspyts_AfterValidator__(__rspyts_validate_json_value__),\n",
"]\n",
));
if models_use_buffer(contract) {
output.push_str(
"\nimport numpy as __rspyts_np__\n\
from numpy.typing import NDArray as __rspyts_NDArray__\n",
);
}
output.push_str(&foreign_type_imports(contract));
output.push('\n');
let mut exports = vec!["JsonValue".to_owned()];
exports.extend(
contract
.foreign_types
.values()
.map(|item| item.name.clone()),
);
for item in &manifest.types {
exports.push(item.name.clone());
if let TypeShape::TaggedEnum { variants, .. } = &item.shape {
exports.extend(
variants
.iter()
.map(|variant| python_tagged_variant_name(&item.name, &variant.rust_name)),
);
}
}
output.push_str(&all_assignment(&exports));
for item in ordered_types(manifest) {
output.push_str(&python_type(item, names, contract));
output.push('\n');
}
output
}
fn python_type(item: &TypeDef, names: &TypeNames, contract: &ResolvedContract) -> String {
let mut output = String::new();
match &item.shape {
TypeShape::Struct { fields } => {
output.push_str(&format!("class {}(__rspyts_BaseModel__):\n", item.name));
output.push_str(&python_doc(item.docs.as_deref(), " "));
output.push_str(
" model_config = __rspyts_ConfigDict__(populate_by_name=True, arbitrary_types_allowed=True, extra=\"forbid\", frozen=True, validate_default=True, allow_inf_nan=False, strict=True)\n",
);
if fields.is_empty() {
output.push_str(" pass\n");
}
for field in fields {
output.push_str(&python_field(field, names, contract));
output.push('\n');
}
}
TypeShape::StringEnum { variants } => {
output.push_str(&format!(
"class {}(__rspyts_builtins__.str, __rspyts_Enum__):\n",
item.name
));
output.push_str(&python_doc(item.docs.as_deref(), " "));
if variants.is_empty() {
output.push_str(" pass\n");
}
for variant in variants {
output.push_str(&format!(
" {} = {}\n",
variant.rust_name,
python_string_literal(&variant.wire_name)
));
}
for variant in variants {
if let Some(docs) = variant.docs.as_deref() {
output.push_str(&format!(
"{}.{}.__doc__ = {}\n",
item.name,
variant.rust_name,
python_string_literal(docs)
));
}
}
}
TypeShape::TaggedEnum { tag, variants } => {
let tag_field = python_tag_field(tag, variants);
let mut variant_names = Vec::new();
for variant in variants {
let name = python_tagged_variant_name(&item.name, &variant.rust_name);
variant_names.push(name.clone());
output.push_str(&format!("class {name}(__rspyts_BaseModel__):\n"));
output.push_str(&python_doc(variant.docs.as_deref(), " "));
output.push_str(
" model_config = __rspyts_ConfigDict__(populate_by_name=True, arbitrary_types_allowed=True, extra=\"forbid\", frozen=True, validate_default=True, allow_inf_nan=False, strict=True)\n",
);
if tag_field == *tag {
let wire_name = python_string_literal(&variant.wire_name);
output.push_str(&format!(
" {tag_field}: __rspyts_Literal__[{wire_name}] = {wire_name}\n"
));
} else {
let wire_name = python_string_literal(&variant.wire_name);
let tag = python_string_literal(tag);
output.push_str(&format!(
" {tag_field}: __rspyts_Literal__[{wire_name}] = __rspyts_Field__(default={wire_name}, alias={tag})\n"
));
}
for field in &variant.fields {
output.push_str(&python_field(field, names, contract));
output.push('\n');
}
output.push('\n');
}
output.push_str(&format!(
"{}: __rspyts_TypeAlias__ = {}\n",
item.name,
variant_names.join(" | ")
));
}
TypeShape::Alias { target } => {
output.push_str(&format!(
"{}: __rspyts_TypeAlias__ = {}\n",
item.name,
python_ref(target, names)
));
}
}
output
}
fn python_field(field: &FieldDef, names: &TypeNames, contract: &ResolvedContract) -> String {
let mut ty = match &field.constraints.literal {
Some(value) => {
let literal = python_literal(value, &field.ty, contract);
if type_allows_null(&field.ty, contract) {
format!("__rspyts_Literal__[{literal}] | None")
} else {
format!("__rspyts_Literal__[{literal}]")
}
}
None => python_field_ref(
&field.ty,
names,
contract,
field.constraints.ge,
field.constraints.le,
),
};
if !field.required && field.default.is_none() && !matches!(field.ty, TypeRef::Option { .. }) {
ty = format!("{ty} | None");
}
let metadata = match field.constraints.literal.as_ref() {
Some(value) => vec![format!(
"__rspyts_BeforeValidator__({})",
python_literal_input_validator(value)
)],
None => Vec::new(),
};
if !metadata.is_empty() {
ty = format!("__rspyts_Annotated__[{ty}, {}]", metadata.join(", "));
}
let mut args = Vec::new();
if let Some(default) = &field.default {
args.push(format!("default={}", python_scalar(default, &field.ty)));
} else if !field.required {
args.push("default=None".to_owned());
}
if let Some(minimum) = field.constraints.min_length {
args.push(format!("min_length={minimum}"));
}
if let Some(maximum) = field.constraints.max_length {
args.push(format!("max_length={maximum}"));
}
if field.rust_name != field.wire_name {
args.push(format!("alias={}", python_string_literal(&field.wire_name)));
}
if let Some(docs) = field.docs.as_deref() {
args.push(format!("description={}", python_string_literal(docs)));
}
let mut output = format!(" {}: {ty}", field.rust_name);
if !args.is_empty() {
output.push_str(&format!(" = __rspyts_Field__({})", args.join(", ")));
}
output
}
fn python_literal_input_validator(value: &ScalarValue) -> &'static str {
match value {
ScalarValue::Bool(_) => "__rspyts_validate_boolean_literal_input__",
ScalarValue::I64(_) => "__rspyts_validate_integer_literal_input__",
ScalarValue::String(_) => "__rspyts_validate_string_literal_input__",
}
}
fn python_literal(value: &ScalarValue, ty: &TypeRef, contract: &ResolvedContract) -> String {
match ty {
TypeRef::Option { item } => python_literal(value, item, contract),
TypeRef::Named { identity } => match type_definition(contract, identity) {
Some(TypeDef {
name,
shape: TypeShape::StringEnum { variants },
..
}) => match value {
ScalarValue::String(value) => variants
.iter()
.find(|variant| variant.wire_name == *value)
.map(|variant| format!("{name}.{}", variant.rust_name))
.unwrap_or_else(|| python_scalar(&ScalarValue::String(value.clone()), ty)),
_ => python_scalar(value, ty),
},
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) => python_literal(value, target, contract),
_ => python_scalar(value, ty),
},
_ => python_scalar(value, ty),
}
}
fn codecs(contract: &ResolvedContract) -> String {
let mut output = String::from(concat!(
"# Generated by rspyts 0.4. Do not edit.\n",
"from __future__ import annotations\n\n",
"import builtins as __rspyts_builtins__\n\n",
));
if !python_uses_buffer(contract) {
output.push_str("__all__: __rspyts_builtins__.list[__rspyts_builtins__.str] = []\n");
return output;
}
output.push_str(concat!(
"import numpy as __rspyts_np__\n\n",
"__all__: __rspyts_builtins__.list[__rspyts_builtins__.str] = []\n\n",
));
output.push_str(&format!(
"from . import {} as __rspyts_native__\n\n",
contract.manifest.module_name
));
output.push_str(concat!(
"__rspyts_DTYPES__ = {\n",
" \"u8\": __rspyts_np__.dtype(\"u1\"),\n",
" \"i8\": __rspyts_np__.dtype(\"i1\"),\n",
" \"u16\": __rspyts_np__.dtype(\"<u2\"),\n",
" \"i16\": __rspyts_np__.dtype(\"<i2\"),\n",
" \"u32\": __rspyts_np__.dtype(\"<u4\"),\n",
" \"i32\": __rspyts_np__.dtype(\"<i4\"),\n",
" \"u64\": __rspyts_np__.dtype(\"<u8\"),\n",
" \"i64\": __rspyts_np__.dtype(\"<i8\"),\n",
" \"f32\": __rspyts_np__.dtype(\"<f4\"),\n",
" \"f64\": __rspyts_np__.dtype(\"<f8\"),\n",
"}\n\n",
"def __rspyts_encode_buffer__(__rspyts_value, __rspyts_dtype: str):\n",
" __rspyts_expected = __rspyts_DTYPES__[__rspyts_dtype]\n",
" if not __rspyts_builtins__.isinstance(__rspyts_value, __rspyts_np__.ndarray):\n",
" raise __rspyts_builtins__.TypeError(\n",
" f\"expected {__rspyts_dtype} buffer as numpy.ndarray\"\n",
" )\n",
" if __rspyts_value.dtype != __rspyts_expected:\n",
" raise __rspyts_builtins__.TypeError(\n",
" f\"expected {__rspyts_dtype} buffer dtype {__rspyts_expected}, \"\n",
" f\"received {__rspyts_value.dtype}\"\n",
" )\n",
" __rspyts_array = __rspyts_np__.ascontiguousarray(__rspyts_value)\n",
" return __rspyts_native__.BufferPayload(\n",
" __rspyts_dtype, __rspyts_array.tobytes(order=\"C\")\n",
" )\n\n",
"def __rspyts_decode_buffer__(__rspyts_payload, __rspyts_dtype: str):\n",
" if __rspyts_payload.dtype != __rspyts_dtype:\n",
" raise __rspyts_builtins__.TypeError(\n",
" f\"expected {__rspyts_dtype} buffer, received {__rspyts_payload.dtype}\"\n",
" )\n",
" return __rspyts_np__.frombuffer(\n",
" __rspyts_payload.data,\n",
" dtype=__rspyts_DTYPES__[__rspyts_dtype],\n",
" count=__rspyts_payload.length,\n",
" ).copy()\n",
));
output
}
fn errors(contract: &ResolvedContract) -> String {
let manifest = &contract.manifest;
let mut output = String::from(concat!(
"# Generated by rspyts 0.4. Do not edit.\n",
"import builtins as __rspyts_builtins__\n\n",
));
output.push_str(&foreign_error_imports(contract));
let mut exports = vec!["ContractError".to_owned(), "ResourceClosedError".to_owned()];
exports.extend(
contract
.foreign_errors
.values()
.map(|item| item.name.clone()),
);
exports.extend(manifest.errors.iter().map(|item| item.name.clone()));
output.push_str(&all_assignment(&exports));
output.push_str(concat!(
"class ContractError(__rspyts_builtins__.Exception):\n",
" code: __rspyts_builtins__.str\n\n",
" def __init__(\n",
" self,\n",
" message: __rspyts_builtins__.str,\n",
" code: __rspyts_builtins__.str,\n",
" ) -> None:\n",
" __rspyts_builtins__.Exception.__init__(self, message)\n",
" self.code = code\n\n",
"class ResourceClosedError(__rspyts_builtins__.RuntimeError):\n",
" pass\n\n",
"def __rspyts_translate_error__(__rspyts_error, __rspyts_error_type):\n",
" if __rspyts_builtins__.len(__rspyts_error.args) >= 2:\n",
" __rspyts_code, __rspyts_message = (\n",
" __rspyts_error.args[0], __rspyts_error.args[1]\n",
" )\n",
" elif (\n",
" __rspyts_error.args\n",
" and __rspyts_builtins__.isinstance(\n",
" __rspyts_error.args[0], __rspyts_builtins__.tuple\n",
" )\n",
" and __rspyts_builtins__.len(__rspyts_error.args[0]) >= 2\n",
" ):\n",
" __rspyts_code, __rspyts_message = (\n",
" __rspyts_error.args[0][0], __rspyts_error.args[0][1]\n",
" )\n",
" else:\n",
" raise __rspyts_error\n",
" return __rspyts_error_type(\n",
" __rspyts_builtins__.str(__rspyts_message),\n",
" __rspyts_builtins__.str(__rspyts_code),\n",
" )\n\n",
"def __rspyts_translate_boundary_error__(__rspyts_error, __rspyts_error_type):\n",
" return __rspyts_error_type(\n",
" __rspyts_builtins__.str(__rspyts_error), \"invalid_argument\"\n",
" )\n\n",
));
for error in &manifest.errors {
output.push_str(&format!("class {}(ContractError):\n", error.name));
output.push_str(&python_doc(error.docs.as_deref(), " "));
output.push_str(" pass\n\n");
}
output
}
fn functions(contract: &ResolvedContract, names: &TypeNames) -> String {
let manifest = &contract.manifest;
let mut output = format!(
"# Generated by rspyts 0.4. Do not edit.\n\
from __future__ import annotations\n\n\
import builtins as __rspyts_builtins__\n\n\
from . import {} as __rspyts_native__\n\
from . import errors as __rspyts_errors__\n\
from . import models as __rspyts_models__\n\
from .errors import * # noqa: F403\n\
from .models import * # noqa: F403\n\n",
manifest.module_name
);
if functions_use_adapter(contract) {
output.push_str(
"from pydantic import AwareDatetime as __rspyts_AwareDatetime__\n\
from pydantic import TypeAdapter as __rspyts_TypeAdapter__\n\n",
);
}
if functions_use_inline_annotations(contract) {
output.push_str(
"from typing import Annotated as __rspyts_Annotated__\n\
from pydantic import Field as __rspyts_Field__\n\n",
);
}
if functions_use_buffer(contract) {
output.push_str(
"from .codecs import __rspyts_decode_buffer__, __rspyts_encode_buffer__\n\
import numpy as __rspyts_np__\n\
from numpy.typing import NDArray as __rspyts_NDArray__\n\n",
);
}
output.push_str(&all_assignment(
&manifest
.functions
.iter()
.filter(|function| matches!(function.target, Target::Both | Target::Python))
.map(|function| function.rust_name.clone())
.collect::<Vec<_>>(),
));
for function in manifest
.functions
.iter()
.filter(|function| matches!(function.target, Target::Both | Target::Python))
{
let params = function
.params
.iter()
.map(|param| format!("{}: {}", param.rust_name, python_ref(¶m.ty, names)))
.collect::<Vec<_>>();
let args = function
.params
.iter()
.map(|param| param.rust_name.as_str())
.collect::<Vec<_>>();
let call = format!(
"__rspyts_native__.{}({})",
function.host_name,
function
.params
.iter()
.zip(args)
.map(|(param, argument)| python_to_native(¶m.ty, argument, contract))
.collect::<Vec<_>>()
.join(", ")
);
let body = match error_name(function.error.as_ref(), contract) {
Some(error) => format!(
" try:\n __rspyts_result = {call}\n except __rspyts_builtins__.RuntimeError as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n except (\n __rspyts_builtins__.TypeError,\n __rspyts_builtins__.ValueError,\n __rspyts_builtins__.OverflowError,\n ) as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_boundary_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n"
),
None => format!(" __rspyts_result = {call}\n"),
};
output.push_str(&format!(
"def {}({}) -> {}:\n{}{} return {}\n\n",
function.rust_name,
params.join(", "),
python_ref(&function.returns, names),
python_doc(function.docs.as_deref(), " "),
body,
python_from_native(&function.returns, "__rspyts_result", contract)
));
}
output
}
fn resources(contract: &ResolvedContract, names: &TypeNames) -> String {
let manifest = &contract.manifest;
let mut output = format!(
"# Generated by rspyts 0.4. Do not edit.\n\
from __future__ import annotations\n\n\
import builtins as __rspyts_builtins__\n\n\
from . import {} as __rspyts_native__\n\
from . import errors as __rspyts_errors__\n\
from . import models as __rspyts_models__\n\n",
manifest.module_name
);
output.push_str(
"from .errors import * # noqa: F403\n\
from .models import * # noqa: F403\n\n",
);
if resources_use_adapter(contract) {
output.push_str(
"from pydantic import AwareDatetime as __rspyts_AwareDatetime__\n\
from pydantic import TypeAdapter as __rspyts_TypeAdapter__\n",
);
}
if resources_use_inline_annotations(contract) {
output.push_str(
"from typing import Annotated as __rspyts_Annotated__\n\
from pydantic import Field as __rspyts_Field__\n",
);
}
if resources_use_buffer(contract) {
output.push_str(
"from .codecs import __rspyts_decode_buffer__, __rspyts_encode_buffer__\n\
import numpy as __rspyts_np__\n\
from numpy.typing import NDArray as __rspyts_NDArray__\n",
);
}
output.push('\n');
output.push_str(&all_assignment(
&manifest
.resources
.iter()
.filter(|resource| matches!(resource.target, Target::Both | Target::Python))
.map(|resource| resource.name.clone())
.collect::<Vec<_>>(),
));
for resource in manifest
.resources
.iter()
.filter(|resource| matches!(resource.target, Target::Both | Target::Python))
{
output.push_str(&format!("class {}:\n", resource.name));
output.push_str(&python_doc(resource.docs.as_deref(), " "));
let constructors = resource
.constructors
.iter()
.filter(|constructor| matches!(constructor.target, Target::Both | Target::Python))
.collect::<Vec<_>>();
let primary = constructors
.iter()
.copied()
.find(|constructor| constructor.rust_name == "new")
.or_else(|| constructors.first().copied());
if let Some(constructor) = primary {
let params = constructor
.params
.iter()
.map(|param| format!("{}: {}", param.rust_name, python_ref(¶m.ty, names)))
.collect::<Vec<_>>();
let args = constructor
.params
.iter()
.map(|param| python_to_native(¶m.ty, ¶m.rust_name, contract))
.collect::<Vec<_>>();
let call = format!("__rspyts_native__.{}({})", resource.name, args.join(", "));
let body = match error_name(constructor.error.as_ref(), contract) {
Some(error) => format!(
" try:\n self.handle = {call}\n except __rspyts_builtins__.RuntimeError as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n except (\n __rspyts_builtins__.TypeError,\n __rspyts_builtins__.ValueError,\n __rspyts_builtins__.OverflowError,\n ) as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_boundary_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n"
),
None => format!(" self.handle = {call}\n"),
};
output.push_str(&format!(
" def __init__(self{}{}) -> None:\n{}{body}",
if params.is_empty() { "" } else { ", " },
params.join(", "),
python_doc(constructor.docs.as_deref(), " ")
));
} else {
output.push_str(" def __init__(self) -> None:\n raise __rspyts_builtins__.TypeError(\"resource has no public constructor\")\n");
}
for constructor in constructors
.into_iter()
.filter(|constructor| Some(*constructor) != primary)
{
let params = constructor
.params
.iter()
.map(|param| format!("{}: {}", param.rust_name, python_ref(¶m.ty, names)))
.collect::<Vec<_>>();
let args = constructor
.params
.iter()
.map(|param| python_to_native(¶m.ty, ¶m.rust_name, contract))
.collect::<Vec<_>>();
let call = format!(
"__rspyts_native__.{}.{}({})",
resource.name,
constructor.host_name,
args.join(", ")
);
let body = match error_name(constructor.error.as_ref(), contract) {
Some(error) => format!(
" try:\n __rspyts_resource.handle = {call}\n except __rspyts_builtins__.RuntimeError as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n except (\n __rspyts_builtins__.TypeError,\n __rspyts_builtins__.ValueError,\n __rspyts_builtins__.OverflowError,\n ) as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_boundary_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n"
),
None => format!(" __rspyts_resource.handle = {call}\n"),
};
output.push_str(&format!(
" @__rspyts_builtins__.classmethod\n def {}(cls{}{}) -> {}:\n{} __rspyts_resource = cls.__new__(cls)\n{body} return __rspyts_resource\n",
constructor.rust_name,
if params.is_empty() { "" } else { ", " },
params.join(", "),
resource.name,
python_doc(constructor.docs.as_deref(), " "),
));
}
for method in resource
.methods
.iter()
.filter(|method| matches!(method.target, Target::Both | Target::Python))
{
let params = method
.params
.iter()
.map(|param| format!("{}: {}", param.rust_name, python_ref(¶m.ty, names)))
.collect::<Vec<_>>();
let args = method
.params
.iter()
.map(|param| python_to_native(¶m.ty, ¶m.rust_name, contract))
.collect::<Vec<_>>();
let call = format!("self.handle.{}({})", method.host_name, args.join(", "));
let body = match error_name(method.error.as_ref(), contract) {
Some(error) => format!(
" try:\n __rspyts_result = {call}\n except __rspyts_builtins__.RuntimeError as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n except (\n __rspyts_builtins__.TypeError,\n __rspyts_builtins__.ValueError,\n __rspyts_builtins__.OverflowError,\n ) as __rspyts_error:\n raise __rspyts_errors__.__rspyts_translate_boundary_error__(\n __rspyts_error, __rspyts_errors__.{error}\n ) from None\n"
),
None => format!(" __rspyts_result = {call}\n"),
};
output.push_str(&format!(
" def {}(self{}{}) -> {}:\n{} if self.handle is None:\n raise __rspyts_errors__.ResourceClosedError(\"{} is closed\")\n{} return {}\n",
method.rust_name,
if params.is_empty() { "" } else { ", " },
params.join(", "),
python_ref(&method.returns, names),
python_doc(method.docs.as_deref(), " "),
resource.name,
body,
python_from_native(&method.returns, "__rspyts_result", contract)
));
}
output.push_str(concat!(
" def close(self) -> None:\n",
" if self.handle is not None:\n",
" self.handle.close()\n",
" self.handle = None\n\n",
" def __enter__(self):\n",
" return self\n\n",
" def __exit__(self, __rspyts_exception_type, __rspyts_exception, __rspyts_traceback) -> None:\n",
" self.close()\n\n",
));
}
output
}
fn constants(contract: &ResolvedContract, fingerprint: &str, names: &TypeNames) -> String {
let manifest = &contract.manifest;
let constants = manifest
.constants
.iter()
.filter(|constant| matches!(constant.target, Target::Both | Target::Python))
.collect::<Vec<_>>();
let fingerprint = python_string_literal(fingerprint);
let mut output = format!(
"# Generated by rspyts 0.4. Do not edit.\n\
from __future__ import annotations\n\n\
import builtins as __rspyts_builtins__\n\n\
from . import models as __rspyts_models__\n\
from .models import * # noqa: F403\n\n\
CONTRACT_FINGERPRINT = {fingerprint}\n"
);
let uses_datetime = constants
.iter()
.any(|constant| reference_contains_datetime(&constant.ty, contract));
let uses_adapter = constants
.iter()
.any(|constant| reference_needs_adapter(&constant.ty, contract));
if uses_datetime {
output.push_str(
"from pydantic import AwareDatetime as __rspyts_AwareDatetime__\n\
from pydantic import TypeAdapter as __rspyts_TypeAdapter__\n\n",
);
} else if uses_adapter {
output.push_str("from pydantic import TypeAdapter as __rspyts_TypeAdapter__\n\n");
}
if constants
.iter()
.any(|constant| reference_contains_buffer(&constant.ty, contract))
{
output.push_str(
"import numpy as __rspyts_np__\n\
from numpy.typing import NDArray as __rspyts_NDArray__\n\n",
);
}
if constants
.iter()
.any(|constant| contains_inline_annotations(&constant.ty))
{
output.push_str(
"from typing import Annotated as __rspyts_Annotated__\n\
from pydantic import Field as __rspyts_Field__\n\n",
);
}
output.push_str(&all_assignment(
&std::iter::once("CONTRACT_FINGERPRINT".to_owned())
.chain(constants.iter().map(|constant| constant.host_name.clone()))
.collect::<Vec<_>>(),
));
for constant in constants {
output.push_str(&format!(
"{}: {} = {}\n",
constant.host_name,
python_ref(&constant.ty, names),
python_value(&constant.value, &constant.ty, contract)
));
}
output
}
fn init_module(contract: &ResolvedContract) -> String {
let manifest = &contract.manifest;
let mut output = String::from(concat!(
"# Generated by rspyts 0.4. Do not edit.\n",
"import builtins as __rspyts_builtins__\n\n",
));
for (alias, dependency) in &contract.dependencies {
let package = dependency
.python
.as_deref()
.expect("Python dependency mapping is validated before emission");
let local = format!("__rspyts_{alias}_contract_fingerprint__");
let expected = python_string_literal(&dependency.fingerprint);
output.push_str(&format!(
concat!(
"from {package} import CONTRACT_FINGERPRINT as {local}\n",
"if {local} != {expected}:\n",
" raise ImportError(f\"contract fingerprint mismatch for Python dependency ",
"{package}: expected {expected_display}, received {{{local}}}\")\n",
"del {local}\n",
),
package = package,
local = local,
expected = expected,
expected_display = dependency.fingerprint,
));
}
if !contract.dependencies.is_empty() {
output.push('\n');
}
output.push_str(
"from .models import * # noqa: F403\n\
from .errors import * # noqa: F403\n\
from .functions import * # noqa: F403\n\
from .resources import * # noqa: F403\n\
from .constants import * # noqa: F403\n\n",
);
let mut exports = vec!["JsonValue".to_owned()];
exports.extend(
contract
.foreign_types
.values()
.map(|item| item.name.clone()),
);
for item in &manifest.types {
exports.push(item.name.clone());
if let TypeShape::TaggedEnum { variants, .. } = &item.shape {
exports.extend(
variants
.iter()
.map(|variant| python_tagged_variant_name(&item.name, &variant.rust_name)),
);
}
}
exports.extend(["ContractError".to_owned(), "ResourceClosedError".to_owned()]);
exports.extend(
contract
.foreign_errors
.values()
.map(|item| item.name.clone()),
);
exports.extend(manifest.errors.iter().map(|item| item.name.clone()));
exports.extend(
manifest
.functions
.iter()
.filter(|function| matches!(function.target, Target::Both | Target::Python))
.map(|function| function.rust_name.clone()),
);
exports.extend(
manifest
.resources
.iter()
.filter(|resource| matches!(resource.target, Target::Both | Target::Python))
.map(|resource| resource.name.clone()),
);
exports.push("CONTRACT_FINGERPRINT".to_owned());
exports.extend(
manifest
.constants
.iter()
.filter(|constant| matches!(constant.target, Target::Both | Target::Python))
.map(|constant| constant.host_name.clone()),
);
output.push_str(&all_assignment(&exports));
output.push_str(&format!(
"__rspyts_contract_module__ = {}\n",
python_string_literal(&manifest.module_name)
));
output
}
fn foreign_type_imports(contract: &ResolvedContract) -> String {
let mut output = String::new();
for dependency in contract.dependencies.values() {
let names = contract
.foreign_types
.iter()
.filter(|(identity, _)| identity.owner == dependency.owner)
.map(|(_, definition)| definition.name.as_str())
.collect::<Vec<_>>();
if !names.is_empty() {
let package = dependency
.python
.as_deref()
.expect("Python dependency mapping is validated before emission");
output.push_str(&format!("from {package} import {}\n", names.join(", ")));
}
}
output
}
fn foreign_error_imports(contract: &ResolvedContract) -> String {
let mut output = String::new();
for dependency in contract.dependencies.values() {
let names = contract
.foreign_errors
.iter()
.filter(|(identity, _)| identity.owner == dependency.owner)
.map(|(_, definition)| definition.name.as_str())
.collect::<Vec<_>>();
if !names.is_empty() {
let package = dependency
.python
.as_deref()
.expect("Python dependency mapping is validated before emission");
output.push_str(&format!("from {package} import {}\n", names.join(", ")));
}
}
output
}
fn all_assignment(names: &[String]) -> String {
let mut seen = std::collections::BTreeSet::new();
let names = names
.iter()
.filter(|name| seen.insert(name.as_str()))
.collect::<Vec<_>>();
if names.is_empty() {
return "__all__: __rspyts_builtins__.list[__rspyts_builtins__.str] = []\n\n".into();
}
format!(
"__all__ = [\n{}]\n\n",
names
.iter()
.map(|name| format!(" {},\n", python_string_literal(name)))
.collect::<String>()
)
}
fn python_ref(reference: &TypeRef, names: &TypeNames) -> String {
match reference {
TypeRef::Unit => "None".into(),
TypeRef::Bool => "__rspyts_builtins__.bool".into(),
TypeRef::Int { signed, bits } => {
let (minimum, maximum) = intrinsic_integer_bounds(*signed, *bits);
format!(
"__rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge={minimum}, le={maximum})]"
)
}
TypeRef::Float { bits: 32 } => concat!(
"__rspyts_Annotated__[__rspyts_builtins__.float, ",
"__rspyts_Field__(ge=-3.4028234663852886e38, ",
"le=3.4028234663852886e38, allow_inf_nan=False)]"
)
.into(),
TypeRef::Float { bits: 64 } => concat!(
"__rspyts_Annotated__[__rspyts_builtins__.float, ",
"__rspyts_Field__(ge=-1.7976931348623157e308, ",
"le=1.7976931348623157e308, allow_inf_nan=False)]"
)
.into(),
TypeRef::Float { bits } => unreachable!("validated float width {bits}"),
TypeRef::String => "__rspyts_builtins__.str".into(),
TypeRef::DateTime => "__rspyts_AwareDatetime__".into(),
TypeRef::Json => "JsonValue".into(),
TypeRef::Option { item } => format!("{} | None", python_ref(item, names)),
TypeRef::List { item } => {
format!("__rspyts_builtins__.list[{}]", python_ref(item, names))
}
TypeRef::Map { value } => format!(
"__rspyts_builtins__.dict[__rspyts_builtins__.str, {}]",
python_ref(value, names)
),
TypeRef::Tuple { items } => format!(
"__rspyts_builtins__.tuple[{}]",
items
.iter()
.map(|item| python_ref(item, names))
.collect::<Vec<_>>()
.join(", ")
),
TypeRef::Named { identity } => names
.get(identity)
.cloned()
.unwrap_or_else(|| identity.to_string()),
TypeRef::Bytes => "__rspyts_builtins__.bytes".into(),
TypeRef::FixedBytes { length } => {
format!(
"__rspyts_Annotated__[__rspyts_builtins__.bytes, __rspyts_Field__(min_length={length}, max_length={length})]"
)
}
TypeRef::Buffer { element } => {
format!(
"__rspyts_NDArray__[__rspyts_np__.{}]",
numpy_scalar(*element)
)
}
}
}
fn python_field_ref(
reference: &TypeRef,
names: &TypeNames,
contract: &ResolvedContract,
minimum: Option<i64>,
maximum: Option<i64>,
) -> String {
match reference {
TypeRef::Int { signed, bits } => {
let (intrinsic_minimum, intrinsic_maximum) = intrinsic_integer_bounds(*signed, *bits);
let minimum = minimum
.map(i128::from)
.unwrap_or(intrinsic_minimum)
.max(intrinsic_minimum);
let maximum = maximum
.map(i128::from)
.unwrap_or(intrinsic_maximum)
.min(intrinsic_maximum);
format!(
"__rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge={minimum}, le={maximum})]"
)
}
TypeRef::Float { bits: 32 } => {
let minimum = minimum
.map(|value| value as f64)
.unwrap_or(-3.4028234663852886e38)
.max(-3.4028234663852886e38);
let maximum = maximum
.map(|value| value as f64)
.unwrap_or(3.4028234663852886e38)
.min(3.4028234663852886e38);
format!(
"__rspyts_Annotated__[__rspyts_builtins__.float, __rspyts_Field__(ge={minimum}, le={maximum}, allow_inf_nan=False)]"
)
}
TypeRef::Float { bits: 64 } => {
let minimum = minimum
.map(|value| value as f64)
.unwrap_or(-1.7976931348623157e308)
.max(-1.7976931348623157e308);
let maximum = maximum
.map(|value| value as f64)
.unwrap_or(1.7976931348623157e308)
.min(1.7976931348623157e308);
format!(
"__rspyts_Annotated__[__rspyts_builtins__.float, __rspyts_Field__(ge={minimum}, le={maximum}, allow_inf_nan=False)]"
)
}
TypeRef::Float { bits } => unreachable!("validated float width {bits}"),
TypeRef::DateTime => concat!(
"__rspyts_Annotated__[__rspyts_AwareDatetime__, ",
"__rspyts_Field__(strict=False), ",
"__rspyts_BeforeValidator__(__rspyts_validate_datetime_input__)]"
)
.to_owned(),
TypeRef::Option { item } => format!(
"{} | None",
python_field_ref(item, names, contract, minimum, maximum)
),
TypeRef::List { item } => format!(
"__rspyts_builtins__.list[{}]",
python_field_ref(item, names, contract, None, None)
),
TypeRef::Map { value } => format!(
"__rspyts_builtins__.dict[__rspyts_builtins__.str, {}]",
python_field_ref(value, names, contract, None, None)
),
TypeRef::Tuple { items } => format!(
"__rspyts_builtins__.tuple[{}]",
items
.iter()
.map(|item| python_field_ref(item, names, contract, None, None))
.collect::<Vec<_>>()
.join(", ")
),
TypeRef::Named { identity } => match type_definition(contract, identity) {
Some(TypeDef {
name,
shape: TypeShape::StringEnum { .. },
..
}) => format!(
"__rspyts_Annotated__[{name}, __rspyts_Field__(strict=False), __rspyts_BeforeValidator__(__rspyts_validate_string_enum_input__)]"
),
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) if minimum.is_some()
|| maximum.is_some()
|| reference_contains_wire_string(target, contract) =>
{
python_field_ref(target, names, contract, minimum, maximum)
}
_ => python_ref(reference, names),
},
_ => python_ref(reference, names),
}
}
fn reference_contains_wire_string(reference: &TypeRef, contract: &ResolvedContract) -> bool {
match reference {
TypeRef::DateTime => true,
TypeRef::Option { item } | TypeRef::List { item } => {
reference_contains_wire_string(item, contract)
}
TypeRef::Map { value } => reference_contains_wire_string(value, contract),
TypeRef::Tuple { items } => items
.iter()
.any(|item| reference_contains_wire_string(item, contract)),
TypeRef::Named { identity } => match type_definition(contract, identity) {
Some(TypeDef {
shape: TypeShape::StringEnum { .. },
..
}) => true,
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) => reference_contains_wire_string(target, contract),
_ => false,
},
_ => false,
}
}
fn intrinsic_integer_bounds(signed: bool, bits: u16) -> (i128, i128) {
match (signed, bits) {
(true, 8 | 16 | 32 | 64) => {
let magnitude = 1_i128 << (bits - 1);
(-magnitude, magnitude - 1)
}
(false, 8 | 16 | 32 | 64) => (0, (1_i128 << bits) - 1),
_ => unreachable!("validated integer width {bits}"),
}
}
fn python_to_native(reference: &TypeRef, expression: &str, contract: &ResolvedContract) -> String {
match reference {
TypeRef::Named { identity } => {
match type_definition(contract, identity).map(|item| &item.shape) {
Some(TypeShape::Struct { fields }) => format!(
"{{{}}}",
fields
.iter()
.map(|field| format!(
"{}: {}",
python_string_literal(&field.wire_name),
python_to_native(
&field.ty,
&format!("{expression}.{}", field.rust_name),
contract
)
))
.collect::<Vec<_>>()
.join(", ")
),
Some(TypeShape::TaggedEnum { tag, variants }) => {
python_tagged_enum_to_native(tag, variants, expression, contract)
}
Some(TypeShape::StringEnum { .. }) => format!("{expression}.value"),
Some(TypeShape::Alias { target }) => python_to_native(target, expression, contract),
None => expression.to_owned(),
}
}
TypeRef::Option { item } => format!(
"None if {expression} is None else {}",
python_to_native(item, expression, contract)
),
TypeRef::List { item } => format!(
"[{} for __rspyts_item in {expression}]",
python_to_native(item, "__rspyts_item", contract)
),
TypeRef::Map { value } => format!(
"{{__rspyts_key: {} for __rspyts_key, __rspyts_item in {expression}.items()}}",
python_to_native(value, "__rspyts_item", contract)
),
TypeRef::Tuple { items } => {
let converted = items
.iter()
.enumerate()
.map(|(index, item)| {
python_to_native(item, &format!("({expression})[{index}]"), contract)
})
.collect::<Vec<_>>()
.join(", ");
let trailing_comma = if items.len() == 1 { "," } else { "" };
let message = python_string_literal(&format!(
"expected tuple of length {}, received a different length",
items.len()
));
format!(
"(({converted}{trailing_comma}) if ({expression}).__len__() == {} else (__rspyts_item for __rspyts_item in ()).throw(__rspyts_builtins__.ValueError({message})))",
items.len()
)
}
TypeRef::Buffer { element } => {
format!(
"__rspyts_encode_buffer__({expression}, {})",
python_string_literal(buffer_name(*element))
)
}
TypeRef::DateTime => format!("{expression}.isoformat()"),
_ => expression.to_owned(),
}
}
fn python_tagged_enum_to_native(
tag: &str,
variants: &[rspyts::ir::EnumVariantDef],
expression: &str,
contract: &ResolvedContract,
) -> String {
let tag_field = python_tag_field(tag, variants);
let variant_value = |variant: &rspyts::ir::EnumVariantDef| {
let fields = std::iter::once(format!(
"{}: {}",
python_string_literal(tag),
python_string_literal(&variant.wire_name)
))
.chain(variant.fields.iter().map(|field| {
format!(
"{}: {}",
python_string_literal(&field.wire_name),
python_to_native(
&field.ty,
&format!("({expression}).{}", field.rust_name),
contract,
)
)
}))
.collect::<Vec<_>>()
.join(", ");
format!("{{{fields}}}")
};
let invalid_message = python_string_literal(&format!(
"unknown discriminator for tagged enum tag `{tag}`"
));
let mut rendered = format!(
"(__rspyts_item for __rspyts_item in ()).throw(__rspyts_builtins__.ValueError({invalid_message}))"
);
for variant in variants.iter().rev() {
let wire_name = python_string_literal(&variant.wire_name);
rendered = format!(
"({} if ({expression}).{tag_field} == {wire_name} else {rendered})",
variant_value(variant),
);
}
rendered
}
fn python_from_native(
reference: &TypeRef,
expression: &str,
contract: &ResolvedContract,
) -> String {
match reference {
TypeRef::Named { identity } => match type_definition(contract, identity) {
Some(TypeDef {
name,
shape: TypeShape::Struct { fields },
..
}) => format!(
"__rspyts_models__.{name}.model_validate({})",
python_object_from_native(fields, expression, contract)
),
Some(TypeDef {
name,
shape: TypeShape::TaggedEnum { tag, variants },
..
}) => python_tagged_enum_from_native(name, tag, variants, expression, contract),
Some(TypeDef {
name,
shape: TypeShape::StringEnum { .. },
..
}) => format!("__rspyts_models__.{name}({expression})"),
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) => python_from_native(target, expression, contract),
_ => expression.to_owned(),
},
TypeRef::Option { item } => format!(
"None if {expression} is None else {}",
python_from_native(item, expression, contract)
),
TypeRef::List { item } => format!(
"[{} for __rspyts_item in {expression}]",
python_from_native(item, "__rspyts_item", contract)
),
TypeRef::Map { value } => format!(
"{{__rspyts_key: {} for __rspyts_key, __rspyts_item in {expression}.items()}}",
python_from_native(value, "__rspyts_item", contract)
),
TypeRef::Tuple { items } => format!(
"({})",
items
.iter()
.enumerate()
.map(|(index, item)| {
python_from_native(item, &format!("{expression}[{index}]"), contract)
})
.collect::<Vec<_>>()
.join(", ")
),
TypeRef::Buffer { element } => {
format!(
"__rspyts_decode_buffer__({expression}, {})",
python_string_literal(buffer_name(*element))
)
}
TypeRef::DateTime => {
format!(
"__rspyts_TypeAdapter__(__rspyts_AwareDatetime__).validate_python({expression})"
)
}
_ => expression.to_owned(),
}
}
fn python_tagged_enum_from_native(
name: &str,
tag: &str,
variants: &[rspyts::ir::EnumVariantDef],
expression: &str,
contract: &ResolvedContract,
) -> String {
let tag_key = python_string_literal(tag);
let variant_value = |variant: &rspyts::ir::EnumVariantDef| {
python_object_from_native(&variant.fields, expression, contract)
};
let (last, remaining) = variants
.split_last()
.expect("tagged enum variants are validated before Python emission");
let mut rendered = variant_value(last);
for variant in remaining.iter().rev() {
let wire_name = python_string_literal(&variant.wire_name);
rendered = format!(
"({} if {expression}[{tag_key}] == {wire_name} else {rendered})",
variant_value(variant),
);
}
format!("__rspyts_TypeAdapter__(__rspyts_models__.{name}).validate_python({rendered})")
}
fn python_object_from_native(
fields: &[FieldDef],
expression: &str,
contract: &ResolvedContract,
) -> String {
let overrides = fields
.iter()
.map(|field| {
let key = python_string_literal(&field.wire_name);
format!(
"**({{{key}: {}}} if {key} in {expression} else {{}})",
python_from_native(&field.ty, &format!("{expression}[{key}]"), contract)
)
})
.collect::<Vec<_>>()
.join(", ");
if overrides.is_empty() {
format!("{{**{expression}}}")
} else {
format!("{{**{expression}, {overrides}}}")
}
}
fn error_name<'a>(
identity: Option<&rspyts::ir::DefinitionId>,
contract: &'a ResolvedContract,
) -> Option<&'a str> {
error_definition(contract, identity?).map(|error| error.name.as_str())
}
fn models_use_buffer(contract: &ResolvedContract) -> bool {
contract
.manifest
.types
.iter()
.chain(contract.foreign_types.values())
.any(|item| match &item.shape {
TypeShape::Struct { fields } => fields.iter().any(|field| contains_buffer(&field.ty)),
TypeShape::TaggedEnum { variants, .. } | TypeShape::StringEnum { variants } => variants
.iter()
.flat_map(|variant| &variant.fields)
.any(|field| contains_buffer(&field.ty)),
TypeShape::Alias { target } => contains_buffer(target),
})
}
fn functions_use_buffer(contract: &ResolvedContract) -> bool {
contract.manifest.functions.iter().any(|function| {
matches!(function.target, Target::Both | Target::Python)
&& (reference_contains_buffer(&function.returns, contract)
|| function
.params
.iter()
.any(|param| reference_contains_buffer(¶m.ty, contract)))
})
}
fn resources_use_buffer(contract: &ResolvedContract) -> bool {
contract.manifest.resources.iter().any(|resource| {
matches!(resource.target, Target::Both | Target::Python)
&& (resource.constructors.iter().any(|constructor| {
matches!(constructor.target, Target::Both | Target::Python)
&& constructor
.params
.iter()
.any(|param| reference_contains_buffer(¶m.ty, contract))
}) || resource.methods.iter().any(|method| {
matches!(method.target, Target::Both | Target::Python)
&& (reference_contains_buffer(&method.returns, contract)
|| method
.params
.iter()
.any(|param| reference_contains_buffer(¶m.ty, contract)))
}))
})
}
fn python_uses_buffer(contract: &ResolvedContract) -> bool {
models_use_buffer(contract)
|| functions_use_buffer(contract)
|| resources_use_buffer(contract)
|| contract.manifest.constants.iter().any(|constant| {
matches!(constant.target, Target::Both | Target::Python)
&& reference_contains_buffer(&constant.ty, contract)
})
}
fn contains_buffer(reference: &TypeRef) -> bool {
match reference {
TypeRef::Buffer { .. } => true,
TypeRef::Option { item } | TypeRef::List { item } => contains_buffer(item),
TypeRef::Map { value } => contains_buffer(value),
TypeRef::Tuple { items } => items.iter().any(contains_buffer),
_ => false,
}
}
fn reference_contains_buffer(reference: &TypeRef, contract: &ResolvedContract) -> bool {
reference_contains(
reference,
contract,
&mut std::collections::BTreeSet::new(),
&|reference| matches!(reference, TypeRef::Buffer { .. }),
)
}
fn reference_contains_datetime(reference: &TypeRef, contract: &ResolvedContract) -> bool {
reference_contains(
reference,
contract,
&mut std::collections::BTreeSet::new(),
&|reference| matches!(reference, TypeRef::DateTime),
)
}
fn functions_use_adapter(contract: &ResolvedContract) -> bool {
contract.manifest.functions.iter().any(|function| {
matches!(function.target, Target::Both | Target::Python)
&& (reference_needs_adapter(&function.returns, contract)
|| function
.params
.iter()
.any(|param| reference_needs_adapter(¶m.ty, contract)))
})
}
fn functions_use_inline_annotations(contract: &ResolvedContract) -> bool {
contract.manifest.functions.iter().any(|function| {
matches!(function.target, Target::Both | Target::Python)
&& (contains_inline_annotations(&function.returns)
|| function
.params
.iter()
.any(|param| contains_inline_annotations(¶m.ty)))
})
}
fn resources_use_adapter(contract: &ResolvedContract) -> bool {
contract.manifest.resources.iter().any(|resource| {
matches!(resource.target, Target::Both | Target::Python)
&& (resource.constructors.iter().any(|constructor| {
matches!(constructor.target, Target::Both | Target::Python)
&& constructor
.params
.iter()
.any(|param| reference_needs_adapter(¶m.ty, contract))
}) || resource.methods.iter().any(|method| {
matches!(method.target, Target::Both | Target::Python)
&& (reference_needs_adapter(&method.returns, contract)
|| method
.params
.iter()
.any(|param| reference_needs_adapter(¶m.ty, contract)))
}))
})
}
fn resources_use_inline_annotations(contract: &ResolvedContract) -> bool {
contract.manifest.resources.iter().any(|resource| {
matches!(resource.target, Target::Both | Target::Python)
&& (resource.constructors.iter().any(|constructor| {
matches!(constructor.target, Target::Both | Target::Python)
&& constructor
.params
.iter()
.any(|param| contains_inline_annotations(¶m.ty))
}) || resource.methods.iter().any(|method| {
matches!(method.target, Target::Both | Target::Python)
&& (contains_inline_annotations(&method.returns)
|| method
.params
.iter()
.any(|param| contains_inline_annotations(¶m.ty)))
}))
})
}
fn contains_inline_annotations(reference: &TypeRef) -> bool {
match reference {
TypeRef::Int { .. } | TypeRef::Float { .. } | TypeRef::FixedBytes { .. } => true,
TypeRef::Option { item } | TypeRef::List { item } => contains_inline_annotations(item),
TypeRef::Map { value } => contains_inline_annotations(value),
TypeRef::Tuple { items } => items.iter().any(contains_inline_annotations),
_ => false,
}
}
fn reference_needs_adapter(reference: &TypeRef, contract: &ResolvedContract) -> bool {
reference_needs_adapter_at(reference, contract, &mut std::collections::BTreeSet::new())
}
fn reference_needs_adapter_at(
reference: &TypeRef,
contract: &ResolvedContract,
visited: &mut std::collections::BTreeSet<rspyts::ir::DefinitionId>,
) -> bool {
match reference {
TypeRef::DateTime => true,
TypeRef::Named { identity } if visited.insert(identity.clone()) => {
match type_definition(contract, identity).map(|item| &item.shape) {
Some(TypeShape::TaggedEnum { .. }) => true,
Some(TypeShape::Struct { fields }) => fields
.iter()
.any(|field| reference_needs_adapter_at(&field.ty, contract, visited)),
Some(TypeShape::StringEnum { .. }) | None => false,
Some(TypeShape::Alias { target }) => {
reference_needs_adapter_at(target, contract, visited)
}
}
}
TypeRef::Option { item } | TypeRef::List { item } => {
reference_needs_adapter_at(item, contract, visited)
}
TypeRef::Map { value } => reference_needs_adapter_at(value, contract, visited),
TypeRef::Tuple { items } => items
.iter()
.any(|item| reference_needs_adapter_at(item, contract, visited)),
_ => false,
}
}
fn reference_contains(
reference: &TypeRef,
contract: &ResolvedContract,
visited: &mut std::collections::BTreeSet<rspyts::ir::DefinitionId>,
predicate: &impl Fn(&TypeRef) -> bool,
) -> bool {
if predicate(reference) {
return true;
}
match reference {
TypeRef::Named { identity } if visited.insert(identity.clone()) => {
match type_definition(contract, identity).map(|item| &item.shape) {
Some(TypeShape::Struct { fields }) => fields
.iter()
.any(|field| reference_contains(&field.ty, contract, visited, predicate)),
Some(TypeShape::TaggedEnum { variants, .. })
| Some(TypeShape::StringEnum { variants }) => variants
.iter()
.flat_map(|variant| &variant.fields)
.any(|field| reference_contains(&field.ty, contract, visited, predicate)),
Some(TypeShape::Alias { target }) => {
reference_contains(target, contract, visited, predicate)
}
None => false,
}
}
TypeRef::Option { item } | TypeRef::List { item } => {
reference_contains(item, contract, visited, predicate)
}
TypeRef::Map { value } => reference_contains(value, contract, visited, predicate),
TypeRef::Tuple { items } => items
.iter()
.any(|item| reference_contains(item, contract, visited, predicate)),
_ => false,
}
}
fn python_scalar(value: &ScalarValue, _ty: &TypeRef) -> String {
match value {
ScalarValue::Bool(value) => if *value { "True" } else { "False" }.to_owned(),
ScalarValue::I64(value) => value.to_string(),
ScalarValue::String(value) => python_string_literal(value),
}
}
fn buffer_name(element: BufferElement) -> &'static str {
match element {
BufferElement::U8 => "u8",
BufferElement::I8 => "i8",
BufferElement::U16 => "u16",
BufferElement::I16 => "i16",
BufferElement::U32 => "u32",
BufferElement::I32 => "i32",
BufferElement::U64 => "u64",
BufferElement::I64 => "i64",
BufferElement::F32 => "f32",
BufferElement::F64 => "f64",
}
}
fn numpy_scalar(element: BufferElement) -> &'static str {
match element {
BufferElement::U8 => "uint8",
BufferElement::I8 => "int8",
BufferElement::U16 => "uint16",
BufferElement::I16 => "int16",
BufferElement::U32 => "uint32",
BufferElement::I32 => "int32",
BufferElement::U64 => "uint64",
BufferElement::I64 => "int64",
BufferElement::F32 => "float32",
BufferElement::F64 => "float64",
}
}
fn numpy_dtype(element: BufferElement) -> &'static str {
match element {
BufferElement::U8 => "u1",
BufferElement::I8 => "i1",
BufferElement::U16 => "<u2",
BufferElement::I16 => "<i2",
BufferElement::U32 => "<u4",
BufferElement::I32 => "<i4",
BufferElement::U64 => "<u8",
BufferElement::I64 => "<i8",
BufferElement::F32 => "<f4",
BufferElement::F64 => "<f8",
}
}
fn python_value(value: &Value, ty: &TypeRef, contract: &ResolvedContract) -> String {
match (value, ty) {
(Value::Null, _) => "None".into(),
(Value::Bool(value), _) => if *value { "True" } else { "False" }.into(),
(Value::Number(value), _) => value.to_string(),
(Value::String(value), TypeRef::Named { identity }) => {
match type_definition(contract, identity) {
Some(TypeDef {
name,
shape: TypeShape::StringEnum { .. },
..
}) => format!("__rspyts_models__.{name}({})", python_string_literal(value)),
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) => python_value(&Value::String(value.clone()), target, contract),
_ => python_string_literal(value),
}
}
(Value::String(value), TypeRef::DateTime) => {
format!(
"__rspyts_TypeAdapter__(__rspyts_AwareDatetime__).validate_python({})",
python_string_literal(value)
)
}
(Value::String(value), _) => python_string_literal(value),
(Value::Array(items), TypeRef::List { item }) => format!(
"[{}]",
items
.iter()
.map(|value| python_value(value, item, contract))
.collect::<Vec<_>>()
.join(", ")
),
(Value::Array(items), TypeRef::Tuple { items: types }) => format!(
"({}{})",
items
.iter()
.zip(types)
.map(|(value, ty)| python_value(value, ty, contract))
.collect::<Vec<_>>()
.join(", "),
if items.len() == 1 { "," } else { "" }
),
(Value::Array(items), TypeRef::Bytes | TypeRef::FixedBytes { .. }) => format!(
"__rspyts_builtins__.bytes([{}])",
items
.iter()
.map(Value::to_string)
.collect::<Vec<_>>()
.join(", ")
),
(Value::Array(items), TypeRef::Buffer { element }) => format!(
"__rspyts_np__.array([{}], dtype=__rspyts_np__.dtype({}))",
items
.iter()
.map(Value::to_string)
.collect::<Vec<_>>()
.join(", "),
python_string_literal(numpy_dtype(*element)),
),
(Value::Object(items), TypeRef::Map { value }) => format!(
"{{{}}}",
items
.iter()
.map(|(key, item)| format!(
"{}: {}",
python_string_literal(key),
python_value(item, value, contract)
))
.collect::<Vec<_>>()
.join(", ")
),
(value, TypeRef::Option { item }) => python_value(value, item, contract),
(value, TypeRef::Named { identity }) => match type_definition(contract, identity) {
Some(TypeDef {
name,
shape: TypeShape::Struct { fields },
..
}) => format!(
"__rspyts_models__.{name}.model_validate({})",
python_object_value(value, fields, contract)
),
Some(TypeDef {
name,
shape: TypeShape::TaggedEnum { tag, variants },
..
}) => {
let fields = value
.as_object()
.and_then(|items| items.get(tag))
.and_then(Value::as_str)
.and_then(|wire_name| {
variants
.iter()
.find(|variant| variant.wire_name == wire_name)
})
.map(|variant| variant.fields.as_slice())
.unwrap_or_default();
format!(
"__rspyts_TypeAdapter__(__rspyts_models__.{name}).validate_python({})",
python_object_value(value, fields, contract)
)
}
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) => python_value(value, target, contract),
_ => serde_json::to_string(value).expect("JSON value is serializable"),
},
(Value::Array(items), _) => format!(
"[{}]",
items.iter().map(python_json).collect::<Vec<_>>().join(", ")
),
(Value::Object(items), _) => format!(
"{{{}}}",
items
.iter()
.map(|(key, value)| format!(
"{}: {}",
python_string_literal(key),
python_json(value)
))
.collect::<Vec<_>>()
.join(", ")
),
}
}
fn python_object_value(value: &Value, fields: &[FieldDef], contract: &ResolvedContract) -> String {
let Value::Object(items) = value else {
return python_json(value);
};
format!(
"{{{}}}",
items
.iter()
.map(|(key, value)| {
let rendered = fields
.iter()
.find(|field| field.wire_name == *key)
.map_or_else(
|| python_json(value),
|field| python_value(value, &field.ty, contract),
);
format!("{}: {rendered}", python_string_literal(key))
})
.collect::<Vec<_>>()
.join(", ")
)
}
fn python_json(value: &Value) -> String {
match value {
Value::Null => "None".into(),
Value::Bool(value) => if *value { "True" } else { "False" }.into(),
Value::Number(value) => value.to_string(),
Value::String(value) => python_string_literal(value),
Value::Array(values) => format!(
"[{}]",
values
.iter()
.map(python_json)
.collect::<Vec<_>>()
.join(", ")
),
Value::Object(values) => format!(
"{{{}}}",
values
.iter()
.map(|(key, value)| format!(
"{}: {}",
python_string_literal(key),
python_json(value)
))
.collect::<Vec<_>>()
.join(", ")
),
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeMap;
use std::fs;
use std::process::Command;
use std::time::{SystemTime, UNIX_EPOCH};
use super::*;
use rspyts::ir::Manifest;
fn owner() -> rspyts::ir::CargoPackageId {
rspyts::ir::CargoPackageId::new("sample")
}
fn identity(id: &str) -> rspyts::ir::DefinitionId {
rspyts::ir::DefinitionId {
owner: owner(),
id: id.into(),
}
}
fn resolved(manifest: Manifest) -> ResolvedContract {
ResolvedContract {
manifest,
hosts: crate::LockedHosts {
python: Some("sample".into()),
typescript: None,
},
dependencies: BTreeMap::new(),
foreign_types: BTreeMap::new(),
foreign_errors: BTreeMap::new(),
}
}
fn integer_fields_manifest() -> Manifest {
let u32_alias = identity("sample::U32Alias");
let integer_field = |name: &str, signed: bool, bits: u16| FieldDef {
rust_name: name.into(),
wire_name: name.into(),
docs: None,
ty: TypeRef::Int { signed, bits },
required: true,
default: None,
constraints: Default::default(),
};
Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: u32_alias.id.clone(),
name: "U32Alias".into(),
docs: None,
shape: TypeShape::Alias {
target: TypeRef::Int {
signed: false,
bits: 32,
},
},
},
TypeDef {
owner: owner(),
id: "sample::IntegerFields".into(),
name: "IntegerFields".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
integer_field("i8_value", true, 8),
integer_field("i16_value", true, 16),
integer_field("i32_value", true, 32),
integer_field("i64_value", true, 64),
integer_field("u8_value", false, 8),
integer_field("u16_value", false, 16),
integer_field("u32_value", false, 32),
integer_field("u64_value", false, 64),
FieldDef {
rust_name: "bounded_u32".into(),
wire_name: "bounded_u32".into(),
docs: None,
ty: TypeRef::Int {
signed: false,
bits: 32,
},
required: false,
default: Some(ScalarValue::I64(10)),
constraints: rspyts::ir::FieldConstraints {
ge: Some(10),
le: Some(200),
..Default::default()
},
},
FieldDef {
rust_name: "literal_u8".into(),
wire_name: "literal_u8".into(),
docs: None,
ty: TypeRef::Int {
signed: false,
bits: 8,
},
required: true,
default: None,
constraints: rspyts::ir::FieldConstraints {
literal: Some(ScalarValue::I64(7)),
..Default::default()
},
},
FieldDef {
rust_name: "default_i16".into(),
wire_name: "default_i16".into(),
docs: None,
ty: TypeRef::Int {
signed: true,
bits: 16,
},
required: false,
default: Some(ScalarValue::I64(-3)),
constraints: Default::default(),
},
FieldDef {
rust_name: "alias_u32".into(),
wire_name: "alias_u32".into(),
docs: None,
ty: TypeRef::Named {
identity: u32_alias,
},
required: true,
default: None,
constraints: Default::default(),
},
FieldDef {
rust_name: "optional_i8".into(),
wire_name: "optional_i8".into(),
docs: None,
ty: TypeRef::Option {
item: Box::new(TypeRef::Int {
signed: true,
bits: 8,
}),
},
required: false,
default: None,
constraints: Default::default(),
},
],
},
},
],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
}
}
#[test]
fn integer_fields_emit_intrinsic_rust_bounds() {
let contract = resolved(integer_fields_manifest());
let generated = models(&contract, &type_names(&contract));
for expected in [
"U32Alias: __rspyts_TypeAlias__ = __rspyts_Annotated__",
"i8_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=-128, le=127)]",
"i16_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=-32768, le=32767)]",
"i32_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=-2147483648, le=2147483647)]",
"i64_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=-9223372036854775808, le=9223372036854775807)]",
"u8_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=0, le=255)]",
"u16_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=0, le=65535)]",
"u32_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=0, le=4294967295)]",
"u64_value: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=0, le=18446744073709551615)]",
"bounded_u32: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=10, le=200)] = __rspyts_Field__(default=10)",
"literal_u8: __rspyts_Annotated__[__rspyts_Literal__[7], __rspyts_BeforeValidator__(__rspyts_validate_integer_literal_input__)]",
"default_i16: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=-32768, le=32767)] = __rspyts_Field__(default=-3)",
"alias_u32: U32Alias",
"optional_i8: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=-128, le=127)] | None = __rspyts_Field__(default=None)",
] {
assert!(
generated.contains(expected),
"missing generated integer constraint: {expected}\n{generated}"
);
}
}
#[test]
fn pydantic_enforces_intrinsic_rust_integer_bounds() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let contract = resolved(integer_fields_manifest());
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-integers-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
fs::create_dir_all(&root).unwrap();
fs::write(
root.join("models.py"),
models(&contract, &type_names(&contract)),
)
.unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", &root)
.arg("-c")
.arg(
r#"
from pydantic import ValidationError
from models import IntegerFields
limits = {
"i8_value": (-128, 127),
"i16_value": (-32768, 32767),
"i32_value": (-2147483648, 2147483647),
"i64_value": (-9223372036854775808, 9223372036854775807),
"u8_value": (0, 255),
"u16_value": (0, 65535),
"u32_value": (0, 4294967295),
"u64_value": (0, 18446744073709551615),
}
base = {name: 0 for name in limits}
base["literal_u8"] = 7
base["alias_u32"] = 0
for name, (minimum, maximum) in limits.items():
for valid in (minimum, maximum):
values = dict(base)
values[name] = valid
assert getattr(IntegerFields(**values), name) == valid
for invalid in (minimum - 1, maximum + 1):
values = dict(base)
values[name] = invalid
try:
IntegerFields(**values)
except ValidationError:
pass
else:
raise AssertionError(f"{name} accepted out-of-range value {invalid}")
model = IntegerFields(**base)
assert model.bounded_u32 == 10
assert model.default_i16 == -3
assert model.optional_i8 is None
for field, invalid in (
("bounded_u32", 9),
("bounded_u32", 4294967296),
("literal_u8", 8),
("alias_u32", -1),
("alias_u32", 4294967296),
("optional_i8", -129),
("optional_i8", 128),
):
values = dict(base)
values[field] = invalid
try:
IntegerFields(**values)
except ValidationError:
pass
else:
raise AssertionError(f"{field} accepted invalid value {invalid}")
for invalid_literal in (7.0, True, "7"):
values = dict(base)
values["literal_u8"] = invalid_literal
try:
IntegerFields(**values)
except ValidationError:
pass
else:
raise AssertionError(f"literal_u8 accepted coercible value {invalid_literal!r}")
"#,
)
.output()
.expect("Python is required to verify generated Pydantic integer models");
assert!(
result.status.success(),
"integer bounds failed real Pydantic validation:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn pydantic_recursively_enforces_scalar_ranges_without_constraining_buffers() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, numpy, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let u8_alias = identity("sample::U8Alias");
let f32_alias = identity("sample::F32Alias");
let nested_alias = identity("sample::NestedAlias");
let numeric_envelope = identity("sample::NumericEnvelope");
let field = |name: &str, ty: TypeRef| FieldDef {
rust_name: name.into(),
wire_name: name.into(),
docs: None,
ty,
required: true,
default: None,
constraints: Default::default(),
};
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: u8_alias.id.clone(),
name: "U8Alias".into(),
docs: None,
shape: TypeShape::Alias {
target: TypeRef::Int {
signed: false,
bits: 8,
},
},
},
TypeDef {
owner: owner(),
id: f32_alias.id.clone(),
name: "F32Alias".into(),
docs: None,
shape: TypeShape::Alias {
target: TypeRef::Float { bits: 32 },
},
},
TypeDef {
owner: owner(),
id: nested_alias.id.clone(),
name: "NestedAlias".into(),
docs: None,
shape: TypeShape::Alias {
target: TypeRef::List {
item: Box::new(TypeRef::Option {
item: Box::new(TypeRef::Named {
identity: u8_alias.clone(),
}),
}),
},
},
},
TypeDef {
owner: owner(),
id: numeric_envelope.id.clone(),
name: "NumericEnvelope".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field(
"optional_values",
TypeRef::Option {
item: Box::new(TypeRef::Named {
identity: nested_alias,
}),
},
),
field(
"by_name",
TypeRef::Map {
value: Box::new(TypeRef::Int {
signed: true,
bits: 16,
}),
},
),
field(
"pair",
TypeRef::Tuple {
items: vec![
TypeRef::Int {
signed: false,
bits: 32,
},
TypeRef::Named {
identity: f32_alias,
},
TypeRef::Option {
item: Box::new(TypeRef::Float { bits: 64 }),
},
],
},
),
field(
"nested",
TypeRef::List {
item: Box::new(TypeRef::Map {
value: Box::new(TypeRef::Tuple {
items: vec![
TypeRef::Int {
signed: true,
bits: 8,
},
TypeRef::Float { bits: 32 },
],
}),
}),
},
),
field(
"samples_f32",
TypeRef::Buffer {
element: BufferElement::F32,
},
),
field(
"samples_f64",
TypeRef::Buffer {
element: BufferElement::F64,
},
),
],
},
},
],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_LIMITS".into(),
host_name: "DEFAULT_LIMITS".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Tuple {
items: vec![
TypeRef::Named { identity: u8_alias },
TypeRef::Float { bits: 32 },
],
},
value: serde_json::json!([7, 1.5]),
}],
};
let contract = resolved(manifest);
let generated_models = models(&contract, &type_names(&contract));
let generated_constants = constants(&contract, "sha256:test", &type_names(&contract));
for expected in [
"U8Alias: __rspyts_TypeAlias__ = __rspyts_Annotated__",
"F32Alias: __rspyts_TypeAlias__ = __rspyts_Annotated__",
"NestedAlias: __rspyts_TypeAlias__ = __rspyts_builtins__.list[U8Alias | None]",
"by_name: __rspyts_builtins__.dict[__rspyts_builtins__.str, __rspyts_Annotated__",
"__rspyts_builtins__.tuple[__rspyts_Annotated__",
"allow_inf_nan=False",
] {
assert!(
generated_models.contains(expected),
"missing recursive scalar annotation {expected}:\n{generated_models}"
);
}
assert!(generated_constants.contains("Annotated as __rspyts_Annotated__"));
assert!(generated_constants.contains("Field as __rspyts_Field__"));
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-recursive-numerics-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
"sha256:test",
)
.unwrap();
fs::write(root.join("python/sample_contract/native.py"), "").unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.arg("-c")
.arg(
r#"
from copy import deepcopy
from typing import get_type_hints
import numpy as np
from pydantic import ValidationError
import sample_contract.constants as constants_module
from sample_contract import NumericEnvelope
nonfinite_f32 = np.array([np.nan, np.inf, -np.inf], dtype=np.float32)
nonfinite_f64 = np.array([np.nan, np.inf, -np.inf], dtype=np.float64)
base = {
"optional_values": [0, 255, None],
"by_name": {"low": -32768, "high": 32767},
"pair": (4294967295, 3.4028234663852886e38, 1.7976931348623157e308),
"nested": [{"item": (-128, -3.4028234663852886e38)}],
"samples_f32": nonfinite_f32,
"samples_f64": nonfinite_f64,
}
model = NumericEnvelope.model_validate(base)
assert model.samples_f32 is nonfinite_f32
assert model.samples_f64 is nonfinite_f64
assert np.isnan(model.samples_f32[0]) and np.isposinf(model.samples_f32[1])
assert np.isneginf(model.samples_f64[2])
def rejected(update):
values = deepcopy(base)
values.update(update)
try:
NumericEnvelope.model_validate(values)
except ValidationError:
return
raise AssertionError(f"accepted invalid recursive scalar value: {update!r}")
rejected({"optional_values": [256]})
rejected({"optional_values": [-1]})
rejected({"by_name": {"bad": 32768}})
rejected({"by_name": {"bad": -32769}})
rejected({"pair": (4294967296, 0.0, 0.0)})
rejected({"pair": (0, 3.5e38, 0.0)})
for nonfinite in (float("nan"), float("inf"), float("-inf")):
rejected({"pair": (0, nonfinite, 0.0)})
rejected({"pair": (0, 0.0, nonfinite)})
rejected({"nested": [{"bad": (128, 0.0)}]})
rejected({"nested": [{"bad": (0, -3.5e38)}]})
model_hints = get_type_hints(NumericEnvelope, include_extras=True)
constant_hints = get_type_hints(constants_module, include_extras=True)
assert "Annotated" in repr(model_hints["by_name"])
assert "Annotated" in repr(model_hints["pair"])
assert "Annotated" in repr(constant_hints["DEFAULT_LIMITS"])
assert constants_module.DEFAULT_LIMITS == (7, 1.5)
"#,
)
.output()
.expect("Python is required to verify recursive numeric model constraints");
assert!(
result.status.success(),
"recursive numeric Pydantic constraints failed:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn fixed_bytes_are_length_checked_in_nested_pydantic_models_and_boundaries() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: "sample::Packet".into(),
name: "Packet".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
FieldDef {
rust_name: "digest".into(),
wire_name: "digest".into(),
docs: None,
ty: TypeRef::FixedBytes { length: 4 },
required: true,
default: None,
constraints: Default::default(),
},
FieldDef {
rust_name: "chunks".into(),
wire_name: "chunks".into(),
docs: None,
ty: TypeRef::List {
item: Box::new(TypeRef::FixedBytes { length: 2 }),
},
required: true,
default: None,
constraints: Default::default(),
},
],
},
}],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "echo_digest".into(),
host_name: "echoDigest".into(),
docs: None,
target: Target::Python,
params: vec![rspyts::ir::ParamDef {
rust_name: "digest".into(),
host_name: "digest".into(),
ty: TypeRef::FixedBytes { length: 4 },
}],
returns: TypeRef::List {
item: Box::new(TypeRef::FixedBytes { length: 2 }),
},
error: None,
}],
resources: vec![],
constants: vec![rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "EMPTY".into(),
host_name: "EMPTY".into(),
docs: None,
target: Target::Python,
ty: TypeRef::FixedBytes { length: 0 },
value: serde_json::json!([]),
}],
};
let contract = resolved(manifest);
let names = type_names(&contract);
let generated_models = models(&contract, &names);
let generated_functions = functions(&contract, &names);
let generated_constants = constants(&contract, "sha256:test", &names);
assert!(
generated_models
.contains("digest: __rspyts_Annotated__[__rspyts_builtins__.bytes, __rspyts_Field__(min_length=4, max_length=4)]")
);
assert!(
generated_models
.contains("chunks: __rspyts_builtins__.list[__rspyts_Annotated__[__rspyts_builtins__.bytes, __rspyts_Field__(min_length=2, max_length=2)]]")
);
assert!(
generated_functions
.contains("digest: __rspyts_Annotated__[__rspyts_builtins__.bytes, __rspyts_Field__(min_length=4, max_length=4)]")
);
assert!(
generated_functions
.contains("__rspyts_builtins__.list[__rspyts_Annotated__[__rspyts_builtins__.bytes, __rspyts_Field__(min_length=2, max_length=2)]]")
);
assert!(generated_functions.contains("Annotated as __rspyts_Annotated__"));
assert!(generated_functions.contains("Field as __rspyts_Field__"));
assert!(
generated_constants
.contains("EMPTY: __rspyts_Annotated__[__rspyts_builtins__.bytes, __rspyts_Field__(min_length=0, max_length=0)] = __rspyts_builtins__.bytes([])")
);
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-fixed-bytes-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
fs::create_dir_all(&root).unwrap();
fs::write(root.join("models.py"), generated_models).unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", &root)
.arg("-c")
.arg(
r#"
from pydantic import ValidationError
from models import Packet
packet = Packet(digest=b"abcd", chunks=[b"xy", b"zz"])
assert packet.digest == b"abcd"
assert packet.chunks == [b"xy", b"zz"]
for values in (
{"digest": b"abc", "chunks": [b"xy"]},
{"digest": b"abcde", "chunks": [b"xy"]},
{"digest": b"abcd", "chunks": [b"x"]},
{"digest": b"abcd", "chunks": [b"xyz"]},
):
try:
Packet(**values)
except ValidationError:
pass
else:
raise AssertionError(f"accepted invalid fixed bytes: {values!r}")
"#,
)
.output()
.expect("Python is required to verify generated fixed-bytes models");
assert!(
result.status.success(),
"fixed bytes failed real Pydantic validation:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn authored_python_strings_and_docs_compile_import_and_round_trip() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let special = "quote=\" slash=\\ controls=\0\u{1f}\u{7f} lines=\n\r\t separators=\u{2028}\u{2029} unicode=é🦀";
let special_alias = "wire\0\u{1f}\u{7f}\"\\\u{2028}\u{2029}🦀";
let special_tag = "tag\0\u{7f}\u{2028}\u{2029}";
let special_variant = "variant\0\u{1f}\u{7f}\"\\\u{2028}\u{2029}🦀";
let status_identity = identity("sample::Status");
let event_identity = identity("sample::Event");
let reader_identity = identity("sample::Reader");
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: status_identity.id.clone(),
name: "Status".into(),
docs: Some(special.into()),
shape: TypeShape::StringEnum {
variants: vec![rspyts::ir::EnumVariantDef {
rust_name: "Control".into(),
wire_name: special_variant.into(),
docs: Some(special.into()),
fields: vec![],
}],
},
},
TypeDef {
owner: owner(),
id: event_identity.id.clone(),
name: "Event".into(),
docs: Some(special.into()),
shape: TypeShape::TaggedEnum {
tag: special_tag.into(),
variants: vec![
rspyts::ir::EnumVariantDef {
rust_name: "Data".into(),
wire_name: special_variant.into(),
docs: Some(special.into()),
fields: vec![FieldDef {
rust_name: "payload".into(),
wire_name: special_alias.into(),
docs: Some(special.into()),
ty: TypeRef::String,
required: true,
default: None,
constraints: Default::default(),
}],
},
rspyts::ir::EnumVariantDef {
rust_name: "Empty".into(),
wire_name: "empty".into(),
docs: None,
fields: vec![],
},
],
},
},
TypeDef {
owner: owner(),
id: "sample::Record".into(),
name: "Record".into(),
docs: Some(special.into()),
shape: TypeShape::Struct {
fields: vec![
FieldDef {
rust_name: "default_value".into(),
wire_name: "default_value".into(),
docs: Some(special.into()),
ty: TypeRef::String,
required: false,
default: Some(ScalarValue::String(special.into())),
constraints: Default::default(),
},
FieldDef {
rust_name: "aliased_value".into(),
wire_name: special_alias.into(),
docs: Some(special.into()),
ty: TypeRef::String,
required: true,
default: None,
constraints: Default::default(),
},
FieldDef {
rust_name: "literal_value".into(),
wire_name: "literal_value".into(),
docs: None,
ty: TypeRef::String,
required: true,
default: None,
constraints: rspyts::ir::FieldConstraints {
literal: Some(ScalarValue::String(special.into())),
..Default::default()
},
},
],
},
},
],
errors: vec![rspyts::ir::ErrorDef {
owner: owner(),
id: "sample::DocumentedError".into(),
name: "DocumentedError".into(),
docs: Some(special.into()),
variants: vec![],
}],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "read_message".into(),
host_name: "read_message".into(),
docs: Some(special.into()),
target: Target::Python,
params: vec![],
returns: TypeRef::String,
error: None,
}],
resources: vec![rspyts::ir::ResourceDef {
owner: owner(),
id: reader_identity.id.clone(),
name: "Reader".into(),
docs: Some(special.into()),
target: Target::Python,
constructors: vec![
rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "new".into(),
host_name: "new".into(),
docs: Some(special.into()),
target: Target::Python,
params: vec![],
returns: TypeRef::Named {
identity: reader_identity.clone(),
},
error: None,
},
rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "from_default".into(),
host_name: "from_default".into(),
docs: Some(special.into()),
target: Target::Python,
params: vec![],
returns: TypeRef::Named {
identity: reader_identity,
},
error: None,
},
],
methods: vec![rspyts::ir::MethodDef {
rust_name: "read".into(),
host_name: "read".into(),
docs: Some(special.into()),
target: Target::Python,
mutable: false,
params: vec![],
returns: TypeRef::String,
error: None,
}],
}],
constants: vec![
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "SPECIAL".into(),
host_name: "SPECIAL".into(),
docs: Some(special.into()),
target: Target::Python,
ty: TypeRef::String,
value: Value::String(special.into()),
},
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_STATUS".into(),
host_name: "DEFAULT_STATUS".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: status_identity,
},
value: Value::String(special_variant.into()),
},
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_EVENT".into(),
host_name: "DEFAULT_EVENT".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: event_identity,
},
value: Value::Object(serde_json::Map::from_iter([
(special_tag.into(), Value::String(special_variant.into())),
(special_alias.into(), Value::String(special.into())),
])),
},
],
};
let contract = resolved(manifest);
let root = std::env::temp_dir().join(format!(
"rspyts-python-string-literals-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
special,
)
.unwrap();
let package = root.join("python/sample_contract");
fs::write(
package.join("native.py"),
concat!(
"def read_message():\n",
" return \"function result\"\n\n",
"class Reader:\n",
" def __init__(self):\n",
" self.closed = False\n\n",
" @classmethod\n",
" def from_default(cls):\n",
" return cls()\n\n",
" def read(self):\n",
" return \"resource result\"\n\n",
" def close(self):\n",
" self.closed = True\n",
),
)
.unwrap();
for path in fs::read_dir(&package).unwrap() {
let path = path.unwrap().path();
if path.extension().is_some_and(|extension| extension == "py") {
let source = fs::read_to_string(path).unwrap();
assert!(!source.contains('\0'));
assert!(!source.contains("\\u{"));
}
}
let special = python_string_literal(special);
let special_alias = python_string_literal(special_alias);
let special_tag = python_string_literal(special_tag);
let special_variant = python_string_literal(special_variant);
let script = format!(
r#"
import inspect
import sample_contract
from sample_contract import DocumentedError, EventData, Reader, Record, Status, read_message
special = {special}
special_alias = {special_alias}
special_tag = {special_tag}
special_variant = {special_variant}
def assert_special_doc(value):
assert value in (special, inspect.cleandoc(special))
assert_special_doc(Record.__doc__)
assert Record.model_fields["default_value"].description == special
assert Record.model_fields["aliased_value"].description == special
assert_special_doc(Status.__doc__)
assert_special_doc(Status.Control.__doc__)
assert_special_doc(EventData.__doc__)
assert EventData.model_fields["payload"].description == special
assert_special_doc(DocumentedError.__doc__)
assert_special_doc(read_message.__doc__)
assert_special_doc(Reader.__doc__)
assert_special_doc(Reader.__init__.__doc__)
assert_special_doc(Reader.from_default.__doc__)
assert_special_doc(Reader.read.__doc__)
assert read_message() == "function result"
reader = Reader()
assert reader.read() == "resource result"
reader.close()
alternate = Reader.from_default()
assert alternate.read() == "resource result"
alternate.close()
assert sample_contract.CONTRACT_FINGERPRINT == special
assert sample_contract.SPECIAL == special
assert sample_contract.DEFAULT_STATUS is Status.Control
assert sample_contract.DEFAULT_STATUS.value == special_variant
record = Record.model_validate({{
special_alias: special,
"literal_value": special,
}})
assert record.default_value == special
assert record.aliased_value == special
assert record.literal_value == special
assert record.model_dump(by_alias=True)[special_alias] == special
event = sample_contract.DEFAULT_EVENT
assert isinstance(event, EventData)
assert event.variant == special_variant
assert event.payload == special
dumped = event.model_dump(by_alias=True)
assert dumped[special_tag] == special_variant
assert dumped[special_alias] == special
"#
);
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.arg("-c")
.arg(script)
.output()
.expect("Python is required to verify generated string literals");
assert!(
result.status.success(),
"generated Python strings failed compile/import/runtime validation:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn package_import_eagerly_builds_recursive_little_endian_buffer_constants() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, numpy, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let packet_identity = identity("sample::Packet");
let event_identity = identity("sample::Event");
let field = |name: &str, ty: TypeRef| FieldDef {
rust_name: name.into(),
wire_name: name.into(),
docs: None,
ty,
required: true,
default: None,
constraints: Default::default(),
};
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: packet_identity.id.clone(),
name: "Packet".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field(
"samples",
TypeRef::Buffer {
element: BufferElement::F32,
},
),
field(
"channels",
TypeRef::List {
item: Box::new(TypeRef::Buffer {
element: BufferElement::I16,
}),
},
),
],
},
},
TypeDef {
owner: owner(),
id: event_identity.id.clone(),
name: "Event".into(),
docs: None,
shape: TypeShape::TaggedEnum {
tag: "kind".into(),
variants: vec![
rspyts::ir::EnumVariantDef {
rust_name: "Data".into(),
wire_name: "data".into(),
docs: None,
fields: vec![
field(
"packet",
TypeRef::Named {
identity: packet_identity.clone(),
},
),
field(
"direct",
TypeRef::Buffer {
element: BufferElement::U16,
},
),
],
},
rspyts::ir::EnumVariantDef {
rust_name: "Empty".into(),
wire_name: "empty".into(),
docs: None,
fields: vec![],
},
],
},
},
],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_PACKET".into(),
host_name: "DEFAULT_PACKET".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: packet_identity,
},
value: serde_json::json!({
"samples": [1.5, -2.25],
"channels": [[-32768, 32767], [0, 7]],
}),
},
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_EVENT".into(),
host_name: "DEFAULT_EVENT".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: event_identity,
},
value: serde_json::json!({
"kind": "data",
"packet": {
"samples": [3.25, 4.5],
"channels": [[-1, 2]],
},
"direct": [0, 65535],
}),
},
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_RAW".into(),
host_name: "DEFAULT_RAW".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Buffer {
element: BufferElement::F64,
},
value: serde_json::json!([1.0, -2.5]),
},
],
};
let contract = resolved(manifest);
let names = type_names(&contract);
let generated = constants(&contract, "sha256:test", &names);
assert!(generated.contains("import numpy as __rspyts_np__"));
for dtype in ["<f4", "<i2", "<u2", "<f8"] {
assert!(
generated.contains(&format!("dtype=__rspyts_np__.dtype(\"{dtype}\")")),
"missing {dtype} array construction:\n{generated}"
);
}
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-buffer-constants-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
"sha256:test",
)
.unwrap();
fs::write(root.join("python/sample_contract/native.py"), "").unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.arg("-c")
.arg(
r#"
import numpy as np
import sample_contract
from sample_contract import EventData
packet = sample_contract.DEFAULT_PACKET
assert isinstance(packet.samples, np.ndarray)
assert packet.samples.dtype == np.dtype("<f4")
assert packet.samples.tolist() == [1.5, -2.25]
assert [channel.dtype for channel in packet.channels] == [np.dtype("<i2"), np.dtype("<i2")]
assert [channel.tolist() for channel in packet.channels] == [[-32768, 32767], [0, 7]]
event = sample_contract.DEFAULT_EVENT
assert isinstance(event, EventData)
assert event.packet.samples.dtype == np.dtype("<f4")
assert event.packet.samples.tolist() == [3.25, 4.5]
assert event.packet.channels[0].dtype == np.dtype("<i2")
assert event.packet.channels[0].tolist() == [-1, 2]
assert event.direct.dtype == np.dtype("<u2")
assert event.direct.tolist() == [0, 65535]
assert isinstance(sample_contract.DEFAULT_RAW, np.ndarray)
assert sample_contract.DEFAULT_RAW.dtype == np.dtype("<f8")
assert sample_contract.DEFAULT_RAW.tolist() == [1.0, -2.5]
"#,
)
.output()
.expect("Python is required to verify eager NumPy constant imports");
assert!(
result.status.success(),
"generated buffer constants failed eager package import:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn package_import_eagerly_builds_recursive_fixed_bytes_and_tagged_enum_constants() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let packet_identity = identity("sample::Packet");
let choice_identity = identity("sample::Choice");
let field = |name: &str, ty: TypeRef| FieldDef {
rust_name: name.into(),
wire_name: name.into(),
docs: None,
ty,
required: true,
default: None,
constraints: Default::default(),
};
let packet_value = serde_json::json!({
"digest": [1, 2, 3, 4],
"chunks": [[5, 6], [7, 8]],
});
let choice_value = serde_json::json!({
"kind": "inline",
"packet": packet_value.clone(),
"token": [9, 10, 11],
});
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: packet_identity.id.clone(),
name: "Packet".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field("digest", TypeRef::FixedBytes { length: 4 }),
field(
"chunks",
TypeRef::List {
item: Box::new(TypeRef::FixedBytes { length: 2 }),
},
),
],
},
},
TypeDef {
owner: owner(),
id: choice_identity.id.clone(),
name: "Choice".into(),
docs: None,
shape: TypeShape::TaggedEnum {
tag: "kind".into(),
variants: vec![
rspyts::ir::EnumVariantDef {
rust_name: "Inline".into(),
wire_name: "inline".into(),
docs: None,
fields: vec![
field(
"packet",
TypeRef::Named {
identity: packet_identity.clone(),
},
),
field("token", TypeRef::FixedBytes { length: 3 }),
],
},
rspyts::ir::EnumVariantDef {
rust_name: "Empty".into(),
wire_name: "empty".into(),
docs: None,
fields: vec![],
},
],
},
},
TypeDef {
owner: owner(),
id: "sample::Envelope".into(),
name: "Envelope".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field(
"choice",
TypeRef::Named {
identity: choice_identity.clone(),
},
),
field(
"digests",
TypeRef::Map {
value: Box::new(TypeRef::FixedBytes { length: 2 }),
},
),
],
},
},
],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_PACKET".into(),
host_name: "DEFAULT_PACKET".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: packet_identity,
},
value: packet_value,
},
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_CHOICE".into(),
host_name: "DEFAULT_CHOICE".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: choice_identity,
},
value: choice_value.clone(),
},
rspyts::ir::ConstantDef {
owner: owner(),
rust_name: "DEFAULT_ENVELOPE".into(),
host_name: "DEFAULT_ENVELOPE".into(),
docs: None,
target: Target::Python,
ty: TypeRef::Named {
identity: identity("sample::Envelope"),
},
value: serde_json::json!({
"choice": choice_value,
"digests": {
"first": [12, 13],
"second": [14, 15],
},
}),
},
],
};
let contract = resolved(manifest);
let names = type_names(&contract);
let generated_constants = constants(&contract, "sha256:test", &names);
assert!(generated_constants.contains("TypeAdapter as __rspyts_TypeAdapter__"));
assert!(
generated_constants
.contains("__rspyts_TypeAdapter__(__rspyts_models__.Choice).validate_python(")
);
assert!(!generated_constants.contains("Choice.model_validate("));
for expected in [
"__rspyts_builtins__.bytes([1, 2, 3, 4])",
"__rspyts_builtins__.bytes([5, 6])",
"__rspyts_builtins__.bytes([7, 8])",
"__rspyts_builtins__.bytes([9, 10, 11])",
"__rspyts_builtins__.bytes([12, 13])",
"__rspyts_builtins__.bytes([14, 15])",
] {
assert!(
generated_constants.contains(expected),
"missing recursively rendered constant value {expected}:\n{generated_constants}"
);
}
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-constants-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
"sha256:test",
)
.unwrap();
fs::write(root.join("python/sample_contract/native.py"), "").unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.arg("-c")
.arg(
r#"
import sample_contract
from sample_contract import ChoiceInline
assert sample_contract.DEFAULT_PACKET.digest == b"\x01\x02\x03\x04"
assert sample_contract.DEFAULT_PACKET.chunks == [b"\x05\x06", b"\x07\x08"]
assert isinstance(sample_contract.DEFAULT_CHOICE, ChoiceInline)
assert sample_contract.DEFAULT_CHOICE.packet.digest == b"\x01\x02\x03\x04"
assert sample_contract.DEFAULT_CHOICE.token == b"\x09\x0a\x0b"
assert isinstance(sample_contract.DEFAULT_ENVELOPE.choice, ChoiceInline)
assert sample_contract.DEFAULT_ENVELOPE.choice.packet.chunks == [b"\x05\x06", b"\x07\x08"]
assert sample_contract.DEFAULT_ENVELOPE.digests == {
"first": b"\x0c\x0d",
"second": b"\x0e\x0f",
}
"#,
)
.output()
.expect("Python is required to verify eager package constant imports");
assert!(
result.status.success(),
"generated package failed eager constant import:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn tagged_enum_arguments_preserve_nested_bytes_and_alias_non_identifier_tags() {
if !Command::new("python3")
.arg("-c")
.arg("import sys; from pydantic import TypeAdapter; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let packet_identity = identity("sample::Packet");
let event_identity = identity("sample::Event");
let field = |name: &str, ty: TypeRef| FieldDef {
rust_name: name.into(),
wire_name: name.into(),
docs: None,
ty,
required: true,
default: None,
constraints: Default::default(),
};
let event_param = || rspyts::ir::ParamDef {
rust_name: "event".into(),
host_name: "event".into(),
ty: TypeRef::Named {
identity: event_identity.clone(),
},
};
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: packet_identity.id.clone(),
name: "Packet".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field("digest", TypeRef::FixedBytes { length: 4 }),
field(
"chunks",
TypeRef::List {
item: Box::new(TypeRef::FixedBytes { length: 2 }),
},
),
],
},
},
TypeDef {
owner: owner(),
id: event_identity.id.clone(),
name: "Event".into(),
docs: None,
shape: TypeShape::TaggedEnum {
tag: "event-type".into(),
variants: vec![
rspyts::ir::EnumVariantDef {
rust_name: "Data".into(),
wire_name: "data".into(),
docs: None,
fields: vec![
field(
"packet",
TypeRef::Named {
identity: packet_identity,
},
),
field("token", TypeRef::FixedBytes { length: 3 }),
],
},
rspyts::ir::EnumVariantDef {
rust_name: "Idle".into(),
wire_name: "idle".into(),
docs: None,
fields: vec![],
},
],
},
},
],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "send_event".into(),
host_name: "sendEvent".into(),
docs: None,
target: Target::Python,
params: vec![event_param()],
returns: TypeRef::Unit,
error: None,
}],
resources: vec![rspyts::ir::ResourceDef {
owner: owner(),
id: "sample::Receiver".into(),
name: "Receiver".into(),
docs: None,
target: Target::Python,
constructors: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "new".into(),
host_name: "new".into(),
docs: None,
target: Target::Python,
params: vec![event_param()],
returns: TypeRef::Named {
identity: identity("sample::Receiver"),
},
error: None,
}],
methods: vec![rspyts::ir::MethodDef {
rust_name: "send".into(),
host_name: "sendEvent".into(),
docs: None,
target: Target::Python,
mutable: true,
params: vec![event_param()],
returns: TypeRef::Unit,
error: None,
}],
}],
constants: vec![],
};
crate::validate::manifest(&manifest).unwrap();
let contract = resolved(manifest);
let names = type_names(&contract);
let generated_models = models(&contract, &names);
let generated_functions = functions(&contract, &names);
let generated_resources = resources(&contract, &names);
assert!(generated_models.contains(
"variant: __rspyts_Literal__[\"data\"] = __rspyts_Field__(default=\"data\", alias=\"event-type\")"
));
assert!(!generated_models.contains("event-type:"));
assert!(!generated_functions.contains("model_dump(mode=\"json\""));
assert!(!generated_resources.contains("model_dump(mode=\"json\""));
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-tagged-arguments-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
"sha256:test",
)
.unwrap();
fs::write(
root.join("python/sample_contract/native.py"),
r#"
function_calls = []
resource_calls = []
def sendEvent(event):
function_calls.append(event)
class Receiver:
def __init__(self, event):
resource_calls.append(("init", event))
def sendEvent(self, event):
resource_calls.append(("send", event))
def close(self):
pass
"#,
)
.unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.current_dir(root.join("python"))
.arg("-c")
.arg(
r#"
import compileall
assert compileall.compile_dir("sample_contract", quiet=1)
from pydantic import TypeAdapter
import sample_contract.native as native
from sample_contract import Event, EventData, Receiver, send_event
event = TypeAdapter(Event).validate_python(
{
"event-type": "data",
"packet": {
"digest": b"\xff\x00\xfe\x80",
"chunks": [b"\x80\xff", b"\x00\xfe"],
},
"token": b"\xff\x00\x80",
}
)
assert isinstance(event, EventData)
assert event.variant == "data"
expected = {
"event-type": "data",
"packet": {
"digest": b"\xff\x00\xfe\x80",
"chunks": [b"\x80\xff", b"\x00\xfe"],
},
"token": b"\xff\x00\x80",
}
send_event(event)
receiver = Receiver(event)
receiver.send(event)
assert native.function_calls == [expected]
assert native.resource_calls == [("init", expected), ("send", expected)]
unknown = EventData.model_construct(
variant="unknown",
packet=event.packet,
token=event.token,
)
for call in (
lambda: send_event(unknown),
lambda: Receiver(unknown),
lambda: receiver.send(unknown),
):
try:
call()
except ValueError as error:
assert "unknown discriminator" in str(error)
else:
raise AssertionError("unknown discriminator reached the native boundary")
assert native.function_calls == [expected]
assert native.resource_calls == [("init", expected), ("send", expected)]
"#,
)
.output()
.expect("Python is required to verify generated tagged-enum arguments");
assert!(
result.status.success(),
"generated tagged-enum arguments failed real Pydantic validation:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn python_boundaries_reject_buffer_and_tuple_shape_mismatches_before_native_calls() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, numpy, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let packet_identity = identity("sample::Packet");
let tuple = |items| TypeRef::Tuple { items };
let buffer = |element| TypeRef::Buffer { element };
let int = |signed, bits| TypeRef::Int { signed, bits };
let field = |name: &str, ty: TypeRef| FieldDef {
rust_name: name.into(),
wire_name: name.into(),
docs: None,
ty,
required: true,
default: None,
constraints: Default::default(),
};
let config_type = || {
tuple(vec![
buffer(BufferElement::F32),
tuple(vec![int(false, 16), buffer(BufferElement::F64)]),
])
};
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: packet_identity.id.clone(),
name: "Packet".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field("samples", buffer(BufferElement::F32)),
field(
"pair",
tuple(vec![
buffer(BufferElement::I16),
tuple(vec![int(false, 8), buffer(BufferElement::U16)]),
]),
),
],
},
}],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "submit_buffers".into(),
host_name: "submitBuffers".into(),
docs: None,
target: Target::Python,
params: vec![
rspyts::ir::ParamDef {
rust_name: "shape".into(),
host_name: "shape".into(),
ty: tuple(vec![
int(false, 8),
tuple(vec![buffer(BufferElement::I16), buffer(BufferElement::U16)]),
]),
},
rspyts::ir::ParamDef {
rust_name: "raw_f32".into(),
host_name: "rawF32".into(),
ty: buffer(BufferElement::F32),
},
rspyts::ir::ParamDef {
rust_name: "raw_f64".into(),
host_name: "rawF64".into(),
ty: buffer(BufferElement::F64),
},
],
returns: TypeRef::Unit,
error: None,
}],
resources: vec![rspyts::ir::ResourceDef {
owner: owner(),
id: "sample::Runner".into(),
name: "Runner".into(),
docs: None,
target: Target::Python,
constructors: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "new".into(),
host_name: "new".into(),
docs: None,
target: Target::Python,
params: vec![rspyts::ir::ParamDef {
rust_name: "configuration".into(),
host_name: "configuration".into(),
ty: config_type(),
}],
returns: TypeRef::Named {
identity: identity("sample::Runner"),
},
error: None,
}],
methods: vec![rspyts::ir::MethodDef {
rust_name: "run".into(),
host_name: "run".into(),
docs: None,
target: Target::Python,
mutable: false,
params: vec![rspyts::ir::ParamDef {
rust_name: "packet".into(),
host_name: "packet".into(),
ty: TypeRef::Named {
identity: packet_identity,
},
}],
returns: TypeRef::Unit,
error: None,
}],
}],
constants: vec![],
};
let contract = resolved(manifest);
let names = type_names(&contract);
let generated_codecs = codecs(&contract);
let generated_functions = functions(&contract, &names);
let generated_resources = resources(&contract, &names);
assert!(generated_codecs.contains("__rspyts_builtins__.isinstance("));
assert!(generated_codecs.contains("__rspyts_value.dtype != __rspyts_expected"));
assert!(
generated_codecs
.contains("__rspyts_array = __rspyts_np__.ascontiguousarray(__rspyts_value)")
);
assert!(!generated_codecs.contains("ascontiguousarray(__rspyts_value, dtype="));
assert!(generated_functions.contains(".__len__() == 2"));
assert!(generated_resources.contains(".__len__() == 2"));
for generated in [&generated_functions, &generated_resources] {
assert!(generated.contains("Annotated as __rspyts_Annotated__"));
assert!(generated.contains("Field as __rspyts_Field__"));
}
let root = std::env::temp_dir().join(format!(
"rspyts-python-strict-boundaries-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
"sha256:test",
)
.unwrap();
fs::write(
root.join("python/sample_contract/native.py"),
r#"
import numpy as np
DTYPES = {
"i16": np.dtype("<i2"),
"u16": np.dtype("<u2"),
"f32": np.dtype("<f4"),
"f64": np.dtype("<f8"),
}
function_calls = []
resource_calls = []
class BufferPayload:
def __init__(self, dtype, data):
self.dtype = dtype
self.data = data
self.length = len(data) // DTYPES[dtype].itemsize
def submitBuffers(shape, raw_f32, raw_f64):
function_calls.append((shape, raw_f32, raw_f64))
class Runner:
def __init__(self, configuration):
resource_calls.append(("init", configuration))
def run(self, packet):
resource_calls.append(("run", packet))
def close(self):
pass
"#,
)
.unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.arg("-c")
.arg(
r#"
from typing import get_type_hints
import numpy as np
import sample_contract.native as native
from sample_contract import Packet, Runner, submit_buffers
i16 = np.array([-32768, 32767], dtype=np.int16)
u16 = np.array([0, 65535], dtype=np.uint16)
f32_source = np.array([np.nan, 1.0, np.inf, 2.0, -np.inf, 3.0], dtype=np.float32)
f64_source = np.array([np.nan, 1.0, np.inf, 2.0, -np.inf, 3.0], dtype=np.float64)
f32 = f32_source[::2]
f64 = f64_source[::2]
assert not f32.flags.c_contiguous and not f64.flags.c_contiguous
shape = (7, (i16, u16))
def expect_failure(call, calls):
before = list(calls)
try:
call()
except (TypeError, ValueError, AttributeError):
pass
else:
raise AssertionError("invalid boundary input reached native code")
assert calls == before
# Direct buffer inputs require ndarray and exact dtype; no coercion is allowed.
expect_failure(lambda: submit_buffers(shape, [1.0], f64), native.function_calls)
expect_failure(
lambda: submit_buffers(shape, np.array([1.0], dtype=np.float64), f64),
native.function_calls,
)
expect_failure(
lambda: submit_buffers(shape, f32, np.array([1.0], dtype=np.float32)),
native.function_calls,
)
# Direct and recursively nested function tuples reject both short and long shapes.
for invalid in (
(),
(7,),
(7, (i16, u16), "extra"),
(7, (i16,)),
(7, (i16, u16, u16)),
):
expect_failure(lambda invalid=invalid: submit_buffers(invalid, f32, f64), native.function_calls)
submit_buffers(shape, f32, f64)
assert len(native.function_calls) == 1
encoded_shape, encoded_f32, encoded_f64 = native.function_calls[0]
assert encoded_shape[0] == 7
assert encoded_shape[1][0].length == 2 and encoded_shape[1][1].length == 2
decoded_f32 = np.frombuffer(encoded_f32.data, dtype=np.dtype("<f4"))
decoded_f64 = np.frombuffer(encoded_f64.data, dtype=np.dtype("<f8"))
assert encoded_f32.length == 3 and encoded_f64.length == 3
assert np.isnan(decoded_f32[0]) and np.isposinf(decoded_f32[1]) and np.isneginf(decoded_f32[2])
assert np.isnan(decoded_f64[0]) and np.isposinf(decoded_f64[1]) and np.isneginf(decoded_f64[2])
valid_configuration = (f32, (65535, f64))
for invalid in (
(),
(f32,),
(f32, (65535, f64), "extra"),
(f32, (65535,)),
(f32, (65535, f64, f64)),
([1.0], (65535, f64)),
(np.array([1.0], dtype=np.float64), (65535, f64)),
):
expect_failure(lambda invalid=invalid: Runner(invalid), native.resource_calls)
runner = Runner(valid_configuration)
assert len(native.resource_calls) == 1 and native.resource_calls[0][0] == "init"
configuration = native.resource_calls[0][1]
assert configuration[0].length == 3 and configuration[1][1].length == 3
valid_packet = Packet(samples=f32, pair=(i16, (255, u16)))
runner.run(valid_packet)
assert len(native.resource_calls) == 2 and native.resource_calls[1][0] == "run"
# model_construct bypasses Pydantic so the generated method boundary itself is proven strict.
for invalid_packet in (
Packet.model_construct(samples=[1.0], pair=(i16, (255, u16))),
Packet.model_construct(
samples=np.array([1.0], dtype=np.float64),
pair=(i16, (255, u16)),
),
Packet.model_construct(samples=f32, pair=()),
Packet.model_construct(samples=f32, pair=(i16,)),
Packet.model_construct(samples=f32, pair=(i16, (255, u16), "extra")),
Packet.model_construct(samples=f32, pair=(i16, (255,))),
Packet.model_construct(samples=f32, pair=(i16, (255, u16, u16))),
Packet.model_construct(samples=f32, pair=([1, 2], (255, u16))),
):
expect_failure(lambda invalid_packet=invalid_packet: runner.run(invalid_packet), native.resource_calls)
function_hints = get_type_hints(submit_buffers, include_extras=True)
constructor_hints = get_type_hints(Runner.__init__, include_extras=True)
method_hints = get_type_hints(Runner.run, include_extras=True)
assert "Annotated" in repr(function_hints["shape"])
assert "numpy.ndarray" in repr(function_hints["raw_f32"])
assert "Annotated" in repr(constructor_hints["configuration"])
assert method_hints["packet"] is Packet
assert "Annotated" in repr(get_type_hints(Packet, include_extras=True)["pair"])
"#,
)
.output()
.expect("Python is required to verify strict buffer and tuple boundaries");
assert!(
result.status.success(),
"strict Python buffer/tuple boundary validation failed:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn tagged_enum_returns_decode_direct_and_nested_buffers_before_validation() {
if !Command::new("python3")
.arg("-c")
.arg(
"import sys, numpy; from pydantic import TypeAdapter; assert sys.version_info >= (3, 11)",
)
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let packet_identity = identity("sample::BufferPacket");
let event_identity = identity("sample::BufferEvent");
let field = |rust_name: &str, wire_name: &str, ty: TypeRef| FieldDef {
rust_name: rust_name.into(),
wire_name: wire_name.into(),
docs: None,
ty,
required: true,
default: None,
constraints: Default::default(),
};
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: packet_identity.id.clone(),
name: "BufferPacket".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field(
"series",
"series",
TypeRef::List {
item: Box::new(TypeRef::Buffer {
element: BufferElement::I16,
}),
},
),
field(
"by_name",
"byName",
TypeRef::Map {
value: Box::new(TypeRef::Buffer {
element: BufferElement::U16,
}),
},
),
field(
"pair",
"pair",
TypeRef::Tuple {
items: vec![
TypeRef::Buffer {
element: BufferElement::U8,
},
TypeRef::Buffer {
element: BufferElement::F64,
},
],
},
),
field(
"optional",
"optional",
TypeRef::Option {
item: Box::new(TypeRef::Buffer {
element: BufferElement::I32,
}),
},
),
],
},
},
TypeDef {
owner: owner(),
id: event_identity.id.clone(),
name: "BufferEvent".into(),
docs: None,
shape: TypeShape::TaggedEnum {
tag: "event-type".into(),
variants: vec![
rspyts::ir::EnumVariantDef {
rust_name: "Data".into(),
wire_name: "data".into(),
docs: None,
fields: vec![
field(
"samples",
"samples",
TypeRef::Buffer {
element: BufferElement::F32,
},
),
field(
"packet",
"packet",
TypeRef::Named {
identity: packet_identity,
},
),
],
},
rspyts::ir::EnumVariantDef {
rust_name: "Idle".into(),
wire_name: "idle".into(),
docs: None,
fields: vec![],
},
],
},
},
],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "read_event".into(),
host_name: "readEvent".into(),
docs: None,
target: Target::Python,
params: vec![],
returns: TypeRef::Named {
identity: event_identity.clone(),
},
error: None,
}],
resources: vec![rspyts::ir::ResourceDef {
owner: owner(),
id: "sample::Reader".into(),
name: "Reader".into(),
docs: None,
target: Target::Python,
constructors: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "new".into(),
host_name: "new".into(),
docs: None,
target: Target::Python,
params: vec![],
returns: TypeRef::Named {
identity: identity("sample::Reader"),
},
error: None,
}],
methods: vec![rspyts::ir::MethodDef {
rust_name: "read_event".into(),
host_name: "readEvent".into(),
docs: None,
target: Target::Python,
mutable: false,
params: vec![],
returns: TypeRef::Named {
identity: event_identity,
},
error: None,
}],
}],
constants: vec![],
};
crate::validate::manifest(&manifest).unwrap();
let contract = resolved(manifest);
let names = type_names(&contract);
let generated_functions = functions(&contract, &names);
let generated_resources = resources(&contract, &names);
for generated in [&generated_functions, &generated_resources] {
assert!(generated.contains("__rspyts_result[\"event-type\"] == \"data\""));
assert!(
generated
.contains("__rspyts_decode_buffer__(__rspyts_result[\"samples\"], \"f32\")")
);
assert!(generated.contains("__rspyts_models__.BufferPacket.model_validate("));
assert!(generated.contains("__rspyts_decode_buffer__(__rspyts_item, \"i16\")"));
assert!(generated.contains("__rspyts_decode_buffer__(__rspyts_item, \"u16\")"));
assert!(generated.contains(
"__rspyts_decode_buffer__(__rspyts_result[\"packet\"][\"pair\"][0], \"u8\")"
));
assert!(generated.contains(
"__rspyts_decode_buffer__(__rspyts_result[\"packet\"][\"pair\"][1], \"f64\")"
));
assert!(generated.contains(
"__rspyts_decode_buffer__(__rspyts_result[\"packet\"][\"optional\"], \"i32\")"
));
}
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-tagged-returns-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
emit(
&root,
&PythonConfig {
package: "sample_contract".into(),
source: None,
},
&contract,
"sha256:test",
)
.unwrap();
fs::write(
root.join("python/sample_contract/native.py"),
r#"
import numpy as np
DTYPES = {
"u8": np.dtype("u1"),
"i16": np.dtype("<i2"),
"u16": np.dtype("<u2"),
"i32": np.dtype("<i4"),
"f32": np.dtype("<f4"),
"f64": np.dtype("<f8"),
}
class BufferPayload:
def __init__(self, dtype, data):
self.dtype = dtype
self.data = data
self.length = len(data) // DTYPES[dtype].itemsize
def payload(dtype, values):
return BufferPayload(dtype, np.asarray(values, dtype=DTYPES[dtype]).tobytes())
def buildEvent(offset):
return {
"event-type": "data",
"samples": payload("f32", [1.25 + offset, -2.5 - offset]),
"packet": {
"series": [payload("i16", [-2 + offset, 3 + offset])],
"byName": {"first": payload("u16", [1 + offset, 65000])},
"pair": [payload("u8", [0, 255]), payload("f64", [0.5 + offset, -1.5])],
"optional": payload("i32", [-100 + offset, 200]),
},
}
def readEvent():
return buildEvent(0)
class Reader:
def readEvent(self):
return buildEvent(10)
def close(self):
pass
"#,
)
.unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", root.join("python"))
.current_dir(root.join("python"))
.arg("-c")
.arg(
r#"
import compileall
assert compileall.compile_dir("sample_contract", quiet=1)
import numpy as np
from sample_contract import BufferEventData, Reader, read_event
def assert_event(event, offset):
assert isinstance(event, BufferEventData)
assert event.variant == "data"
np.testing.assert_allclose(event.samples, [1.25 + offset, -2.5 - offset])
assert event.samples.dtype == np.dtype("<f4")
np.testing.assert_array_equal(event.packet.series[0], [-2 + offset, 3 + offset])
assert event.packet.series[0].dtype == np.dtype("<i2")
np.testing.assert_array_equal(event.packet.by_name["first"], [1 + offset, 65000])
assert event.packet.by_name["first"].dtype == np.dtype("<u2")
assert isinstance(event.packet.pair, tuple)
np.testing.assert_array_equal(event.packet.pair[0], [0, 255])
assert event.packet.pair[0].dtype == np.dtype("u1")
np.testing.assert_allclose(event.packet.pair[1], [0.5 + offset, -1.5])
assert event.packet.pair[1].dtype == np.dtype("<f8")
np.testing.assert_array_equal(event.packet.optional, [-100 + offset, 200])
assert event.packet.optional.dtype == np.dtype("<i4")
assert_event(read_event(), 0)
reader = Reader()
assert_event(reader.read_event(), 10)
"#,
)
.output()
.expect("Python is required to verify generated tagged-enum returns");
assert!(
result.status.success(),
"generated tagged-enum returns failed real Pydantic validation:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn json_is_recursive_and_precise() {
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: "value".into(),
name: "Value".into(),
docs: None,
shape: TypeShape::Alias {
target: TypeRef::Json,
},
}],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = models(&contract, &type_names(&contract));
assert!(generated.contains("AfterValidator as __rspyts_AfterValidator__"));
assert!(generated.contains("def __rspyts_validate_json_value__("));
assert!(generated.contains("JsonValue: __rspyts_TypeAlias__"));
assert!(generated.contains("Value: __rspyts_TypeAlias__ = JsonValue"));
assert!(!generated.contains("Any"));
}
#[test]
fn pydantic_json_value_validates_nested_record_payloads() {
if !Command::new("python3")
.arg("-c")
.arg("import sys, pydantic; assert sys.version_info >= (3, 11)")
.output()
.is_ok_and(|output| output.status.success())
{
return;
}
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: "sample::Record".into(),
name: "Record".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![FieldDef {
rust_name: "metadata".into(),
wire_name: "metadata".into(),
docs: None,
ty: TypeRef::Json,
required: true,
default: None,
constraints: Default::default(),
}],
},
}],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let root = std::env::temp_dir().join(format!(
"rspyts-pydantic-json-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
fs::create_dir_all(&root).unwrap();
fs::write(
root.join("models.py"),
models(&contract, &type_names(&contract)),
)
.unwrap();
let result = Command::new("python3")
.env("PYTHONPATH", &root)
.arg("-c")
.arg(
r#"from pydantic import TypeAdapter, ValidationError
from models import JsonValue, Record
payload = {
"items": [
-9007199254740991,
9007199254740991,
1.25,
{"category": None, "flags": [True, False]},
]
}
model = Record(metadata=payload)
assert model.model_dump() == {"metadata": payload}
assert TypeAdapter(JsonValue).validate_python(payload) == payload
invalid = [
9007199254740992,
-9007199254740992,
9007199254740992.0,
float("nan"),
float("inf"),
-float("inf"),
]
for value in invalid:
nested = {"items": [None, {"value": value}]}
try:
Record(metadata=nested)
except ValidationError:
pass
else:
raise AssertionError(f"accepted invalid JSON number: {value!r}")
"#,
)
.output()
.expect("Python is required to verify generated Pydantic models");
assert!(
result.status.success(),
"nested JsonValue failed real Pydantic validation:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn buffer_free_modules_do_not_import_missing_codec_helpers() {
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "answer".into(),
host_name: "answer".into(),
docs: None,
target: Target::Python,
params: vec![],
returns: TypeRef::Int {
signed: true,
bits: 32,
},
error: None,
}],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated_codecs = codecs(&contract);
let generated_functions = functions(&contract, &type_names(&contract));
assert_eq!(
generated_codecs,
"# Generated by rspyts 0.4. Do not edit.\nfrom __future__ import annotations\n\nimport builtins as __rspyts_builtins__\n\n__all__: __rspyts_builtins__.list[__rspyts_builtins__.str] = []\n"
);
assert!(!generated_functions.contains("from .codecs import"));
assert!(generated_functions.contains("__all__ = [\n \"answer\",\n]"));
}
#[test]
fn buffer_codecs_import_the_configured_native_module() {
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "bridge".into(),
imports: vec![],
types: vec![],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "samples".into(),
host_name: "samples".into(),
docs: None,
target: Target::Python,
params: vec![],
returns: TypeRef::Buffer {
element: BufferElement::F32,
},
error: None,
}],
resources: vec![],
constants: vec![],
};
let generated = codecs(&resolved(manifest));
assert!(generated.contains("from . import bridge as __rspyts_native__"));
assert!(!generated.contains("from . import native as __rspyts_native__"));
}
#[test]
fn declared_errors_normalize_python_boundary_validation_failures() {
let error_identity = identity("sample::DomainError");
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![],
errors: vec![rspyts::ir::ErrorDef {
owner: owner(),
id: error_identity.id.clone(),
name: "DomainError".into(),
docs: None,
variants: vec![],
}],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "process".into(),
host_name: "process".into(),
docs: None,
target: Target::Python,
params: vec![rspyts::ir::ParamDef {
rust_name: "ratio".into(),
host_name: "ratio".into(),
ty: TypeRef::Float { bits: 64 },
}],
returns: TypeRef::Unit,
error: Some(error_identity),
}],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated_errors = errors(&contract);
let generated_functions = functions(&contract, &type_names(&contract));
assert!(generated_errors.contains("def __rspyts_translate_boundary_error__("));
assert!(generated_functions.contains("__rspyts_builtins__.ValueError"));
assert!(
generated_functions.contains("__rspyts_errors__.__rspyts_translate_boundary_error__(")
);
assert!(generated_functions.contains("__rspyts_errors__.DomainError"));
}
#[test]
fn generated_wrapper_internals_cannot_be_shadowed_by_common_parameter_names() {
let model_identity = identity("sample::CollisionModel");
let error_identity = identity("sample::CollisionError");
let params = [
"native",
"result",
"error",
"TypeAdapter",
"Field",
"CollisionModel",
"CollisionError",
"ValueError",
"resource",
"item",
"key",
]
.into_iter()
.map(|name| rspyts::ir::ParamDef {
rust_name: name.into(),
host_name: name.into(),
ty: TypeRef::String,
})
.collect();
let manifest = Manifest {
ir_version: rspyts::ir::IR_VERSION,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: model_identity.id.clone(),
name: "CollisionModel".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![FieldDef {
rust_name: "updated_at".into(),
wire_name: "updatedAt".into(),
docs: None,
ty: TypeRef::DateTime,
required: true,
default: None,
constraints: Default::default(),
}],
},
}],
errors: vec![rspyts::ir::ErrorDef {
owner: owner(),
id: error_identity.id.clone(),
name: "CollisionError".into(),
docs: None,
variants: vec![],
}],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "collide".into(),
host_name: "collide".into(),
docs: None,
target: Target::Python,
params,
returns: TypeRef::Named {
identity: model_identity,
},
error: Some(error_identity),
}],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = functions(&contract, &type_names(&contract));
assert!(generated.contains("from . import native as __rspyts_native__"));
assert!(generated.contains("from . import errors as __rspyts_errors__"));
assert!(generated.contains("from . import models as __rspyts_models__"));
assert!(generated.contains("native: __rspyts_builtins__.str"));
assert!(generated.contains("CollisionError: __rspyts_builtins__.str"));
assert!(generated.contains("__rspyts_result = __rspyts_native__.collide("));
assert!(generated.contains("__rspyts_errors__.CollisionError"));
assert!(generated.contains("__rspyts_builtins__.ValueError"));
assert!(generated.contains("__rspyts_models__.CollisionModel.model_validate("));
assert!(
generated.contains("__rspyts_TypeAdapter__(__rspyts_AwareDatetime__).validate_python(")
);
}
#[test]
fn models_are_frozen_strict_and_export_only_contract_names() {
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: "sample::Value".into(),
name: "Value".into(),
docs: None,
shape: TypeShape::Struct { fields: vec![] },
}],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = models(&contract, &type_names(&contract));
assert!(generated.contains(
"__rspyts_ConfigDict__(populate_by_name=True, arbitrary_types_allowed=True, extra=\"forbid\", frozen=True, validate_default=True, allow_inf_nan=False, strict=True)"
));
assert!(generated.contains("__all__ = [\n \"JsonValue\",\n \"Value\",\n]"));
assert!(!generated.contains(" \"TypeAdapter\","));
assert!(!generated.contains(" \"BaseModel\","));
}
#[test]
fn models_emit_literals_constraints_typed_defaults_and_aware_datetimes() {
let status_identity = identity("sample::Status");
let optional_label_identity = identity("sample::OptionalLabel");
let field = |rust_name: &str, ty: TypeRef, required: bool, default, constraints| FieldDef {
rust_name: rust_name.into(),
wire_name: rust_name.into(),
docs: None,
ty,
required,
default,
constraints,
};
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![
TypeDef {
owner: owner(),
id: status_identity.id.clone(),
name: "Status".into(),
docs: None,
shape: TypeShape::StringEnum {
variants: vec![rspyts::ir::EnumVariantDef {
rust_name: "Unknown".into(),
wire_name: "unknown".into(),
docs: None,
fields: vec![],
}],
},
},
TypeDef {
owner: owner(),
id: optional_label_identity.id.clone(),
name: "OptionalLabel".into(),
docs: None,
shape: TypeShape::Alias {
target: TypeRef::Option {
item: Box::new(TypeRef::String),
},
},
},
TypeDef {
owner: owner(),
id: "sample::Request".into(),
name: "Request".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field(
"revision",
TypeRef::Int {
signed: false,
bits: 32,
},
true,
None,
rspyts::ir::FieldConstraints {
literal: Some(ScalarValue::I64(2)),
..Default::default()
},
),
field(
"items",
TypeRef::List {
item: Box::new(TypeRef::String),
},
true,
None,
rspyts::ir::FieldConstraints {
min_length: Some(1),
max_length: Some(200),
..Default::default()
},
),
field(
"count",
TypeRef::Int {
signed: false,
bits: 32,
},
false,
Some(ScalarValue::I64(1)),
rspyts::ir::FieldConstraints {
ge: Some(1),
le: Some(200),
..Default::default()
},
),
field(
"status",
TypeRef::Named {
identity: status_identity,
},
false,
Some(ScalarValue::String("unknown".into())),
Default::default(),
),
field(
"updated_at",
TypeRef::DateTime,
true,
None,
Default::default(),
),
field(
"label",
TypeRef::Named {
identity: optional_label_identity,
},
true,
None,
rspyts::ir::FieldConstraints {
literal: Some(ScalarValue::String("ready".into())),
..Default::default()
},
),
],
},
},
],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = models(&contract, &type_names(&contract));
assert!(generated.contains(
"revision: __rspyts_Annotated__[__rspyts_Literal__[2], __rspyts_BeforeValidator__(__rspyts_validate_integer_literal_input__)]"
));
assert!(generated.contains("items: __rspyts_builtins__.list[__rspyts_builtins__.str] = __rspyts_Field__(min_length=1, max_length=200)"));
assert!(
generated.contains("count: __rspyts_Annotated__[__rspyts_builtins__.int, __rspyts_Field__(ge=1, le=200)] = __rspyts_Field__(default=1)")
);
assert!(!generated.contains("count: __rspyts_builtins__.int | None"));
assert!(generated.contains(
"status: __rspyts_Annotated__[Status, __rspyts_Field__(strict=False), __rspyts_BeforeValidator__(__rspyts_validate_string_enum_input__)] = __rspyts_Field__(default=\"unknown\")"
));
assert!(generated.contains(
"updated_at: __rspyts_Annotated__[__rspyts_AwareDatetime__, __rspyts_Field__(strict=False), __rspyts_BeforeValidator__(__rspyts_validate_datetime_input__)]"
));
assert!(generated.contains(
"label: __rspyts_Annotated__[__rspyts_Literal__[\"ready\"] | None, __rspyts_BeforeValidator__(__rspyts_validate_string_literal_input__)]"
));
assert!(generated.contains("validate_default=True"));
}
#[test]
fn nested_datetime_values_are_converted_at_the_native_boundary() {
let record_identity = identity("sample::Record");
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: record_identity.id.clone(),
name: "Record".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![FieldDef {
rust_name: "updated_at".into(),
wire_name: "updatedAt".into(),
docs: None,
ty: TypeRef::DateTime,
required: true,
default: None,
constraints: Default::default(),
}],
},
}],
errors: vec![],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "round_trip".into(),
host_name: "round_trip".into(),
docs: None,
target: Target::Python,
params: vec![rspyts::ir::ParamDef {
rust_name: "record".into(),
host_name: "record".into(),
ty: TypeRef::Named {
identity: record_identity.clone(),
},
}],
returns: TypeRef::Named {
identity: record_identity,
},
error: None,
}],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = functions(&contract, &type_names(&contract));
assert!(generated.contains("record.updated_at.isoformat()"));
assert!(
generated.contains("__rspyts_TypeAdapter__(__rspyts_AwareDatetime__).validate_python(__rspyts_result[\"updatedAt\"])")
);
assert!(generated.contains("AwareDatetime as __rspyts_AwareDatetime__"));
assert!(generated.contains("TypeAdapter as __rspyts_TypeAdapter__"));
}
#[test]
fn python_constants_filter_targets_and_keep_all_strict() {
let constant = |name: &str, target| rspyts::ir::ConstantDef {
owner: owner(),
rust_name: name.into(),
host_name: name.into(),
docs: None,
target,
ty: TypeRef::String,
value: Value::String(name.into()),
};
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![
constant("PYTHON_VALUE", Target::Python),
constant("SHARED_VALUE", Target::Both),
constant("TYPESCRIPT_VALUE", Target::Typescript),
constant("STATIC_VALUE", Target::Static),
],
};
let contract = resolved(manifest);
let generated_constants = constants(&contract, "sha256:test", &type_names(&contract));
let generated_init = init_module(&contract);
for generated in [&generated_constants, &generated_init] {
assert!(generated.contains("PYTHON_VALUE"));
assert!(generated.contains("SHARED_VALUE"));
assert!(!generated.contains("TYPESCRIPT_VALUE"));
assert!(!generated.contains("STATIC_VALUE"));
}
assert!(generated_constants.contains(
"__all__ = [\n \"CONTRACT_FINGERPRINT\",\n \"PYTHON_VALUE\",\n \"SHARED_VALUE\",\n]"
));
}
#[test]
fn generated_source_preserves_shared_namespace_packages() {
let root = std::env::temp_dir().join(format!(
"rspyts-python-namespace-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let owner = root.join("owner");
let consumer = root.join("consumer");
for (staging, package, marker) in [
(&owner, "example.owner.contracts", "owner"),
(&consumer, "example.consumer.contracts", "consumer"),
] {
let parent = staging
.join("python")
.join(package.replace(".contracts", "").replace('.', "/"));
fs::create_dir_all(&parent).unwrap();
fs::write(
parent.join("__init__.py"),
format!("PACKAGE_MARKER = {}\n", python_string_literal(marker)),
)
.unwrap();
let config = PythonConfig {
package: package.into(),
source: None,
};
emit(
staging,
&config,
&resolved(Manifest {
ir_version: 4,
crate_name: marker.into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
}),
"sha256:test",
)
.unwrap();
assert!(!staging.join("python/example/__init__.py").exists());
assert_eq!(
fs::read_to_string(parent.join("__init__.py")).unwrap(),
format!("PACKAGE_MARKER = {}\n", python_string_literal(marker))
);
assert!(parent.join("contracts/__init__.py").is_file());
}
let python_path =
std::env::join_paths([owner.join("python"), consumer.join("python")]).unwrap();
let import = Command::new("python3")
.env("PYTHONPATH", python_path)
.arg("-c")
.arg(
"import example.owner as owner; \
import example.consumer as consumer; \
assert owner.PACKAGE_MARKER == 'owner'; \
assert consumer.PACKAGE_MARKER == 'consumer'",
)
.output()
.expect("Python is required to verify namespace package imports");
assert!(
import.status.success(),
"separately staged namespace packages did not coexist:\n{}{}",
String::from_utf8_lossy(&import.stdout),
String::from_utf8_lossy(&import.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn foreign_models_are_imported_from_the_configured_python_package() {
let dependency_owner = rspyts::ir::CargoPackageId::new("owner");
let item_identity = rspyts::ir::DefinitionId::new(dependency_owner.as_str(), "owner::Item");
let foreign = TypeDef {
owner: dependency_owner.clone(),
id: item_identity.id.clone(),
name: "Item".into(),
docs: None,
shape: TypeShape::Struct { fields: vec![] },
};
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: "sample::Request".into(),
name: "Request".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![rspyts::ir::FieldDef {
rust_name: "item".into(),
wire_name: "item".into(),
docs: None,
ty: TypeRef::Named {
identity: item_identity.clone(),
},
required: true,
default: None,
constraints: rspyts::ir::FieldConstraints::default(),
}],
},
}],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
};
let contract = ResolvedContract {
manifest,
hosts: crate::LockedHosts {
python: Some("sample".into()),
typescript: None,
},
dependencies: BTreeMap::from([(
"owner".into(),
crate::LockedDependency {
owner: dependency_owner,
crate_version: "1.2.3".into(),
fingerprint: "sha256:owner".into(),
python: Some("example.owner.contracts".into()),
typescript: None,
types: vec![foreign.clone()],
errors: vec![],
},
)]),
foreign_types: BTreeMap::from([(item_identity, foreign)]),
foreign_errors: BTreeMap::new(),
};
let generated = models(&contract, &type_names(&contract));
let generated_init = init_module(&contract);
assert!(generated.contains("from example.owner.contracts import Item"));
assert!(generated.contains(" item: Item"));
assert!(!generated.contains("class Item"));
assert!(generated_init.contains(
"from example.owner.contracts import CONTRACT_FINGERPRINT as __rspyts_owner_contract_fingerprint__"
));
assert!(generated_init.contains(
"contract fingerprint mismatch for Python dependency example.owner.contracts: expected sha256:owner"
));
assert!(generated_init.contains(
"if __rspyts_owner_contract_fingerprint__ != \"sha256:owner\":\n raise ImportError"
));
assert!(generated_init.contains("del __rspyts_owner_contract_fingerprint__\n"));
let exports = generated_init
.split_once("__all__ =")
.expect("generated package initializer declares explicit exports")
.1;
assert!(!exports.contains("\"rspyts_owner_contract_fingerprint\""));
let root = std::env::temp_dir().join(format!(
"rspyts-python-dependency-guard-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let dependency = root.join("example/owner/contracts");
let package = root.join("sample");
fs::create_dir_all(&dependency).unwrap();
fs::create_dir_all(&package).unwrap();
fs::write(
dependency.join("__init__.py"),
"CONTRACT_FINGERPRINT = \"sha256:owner\"\n",
)
.unwrap();
fs::write(package.join("__init__.py"), generated_init).unwrap();
for module in ["models", "errors", "functions", "resources"] {
fs::write(
package.join(format!("{module}.py")),
"__all__: list[str] = []\n",
)
.unwrap();
}
fs::write(
package.join("constants.py"),
"CONTRACT_FINGERPRINT = \"sha256:sample\"\n\
__all__ = [\"CONTRACT_FINGERPRINT\"]\n",
)
.unwrap();
let import = Command::new("python3")
.env("PYTHONPATH", &root)
.arg("-c")
.arg(
"import sample; \
assert sample.CONTRACT_FINGERPRINT == 'sha256:sample'; \
assert not hasattr(sample, 'rspyts_owner_contract_fingerprint')",
)
.output()
.expect("Python is required to verify the generated dependency guard namespace");
assert!(
import.status.success(),
"generated dependency guard leaked its internal binding:\n{}{}",
String::from_utf8_lossy(&import.stdout),
String::from_utf8_lossy(&import.stderr)
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn resources_emit_named_python_factories_and_filter_other_targets() {
let constructor = |rust_name: &str, host_name: &str, target| rspyts::ir::FunctionDef {
owner: owner(),
rust_name: rust_name.into(),
host_name: host_name.into(),
docs: None,
target,
params: vec![],
returns: TypeRef::Named {
identity: identity("sample::Runner"),
},
error: None,
};
let method = |rust_name: &str, host_name: &str, target| rspyts::ir::MethodDef {
rust_name: rust_name.into(),
host_name: host_name.into(),
docs: None,
target,
mutable: false,
params: vec![],
returns: TypeRef::Unit,
error: None,
};
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![],
errors: vec![],
functions: vec![],
resources: vec![rspyts::ir::ResourceDef {
owner: owner(),
id: "sample::Runner".into(),
name: "Runner".into(),
docs: None,
target: Target::Both,
constructors: vec![
constructor("new", "new", Target::Python),
constructor("from_cache", "fromCache", Target::Both),
constructor("from_browser", "fromBrowser", Target::Typescript),
],
methods: vec![
method("run", "run", Target::Python),
method("browser_only", "browserOnly", Target::Typescript),
],
}],
constants: vec![],
};
let contract = resolved(manifest);
let generated = resources(&contract, &type_names(&contract));
assert!(generated.contains("def __init__(self) -> None"));
assert!(
generated
.contains("@__rspyts_builtins__.classmethod\n def from_cache(cls) -> Runner")
);
assert!(generated.contains("def run(self) -> None"));
assert!(!generated.contains("from_browser"));
assert!(!generated.contains("browser_only"));
}
#[test]
fn every_generated_python_module_compiles() {
let manifest = Manifest {
ir_version: 4,
crate_name: "sample".into(),
crate_version: "1.0.0".into(),
module_name: "native".into(),
imports: vec![],
types: vec![TypeDef {
owner: owner(),
id: "sample::ResultValue".into(),
name: "ResultValue".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![rspyts::ir::FieldDef {
rust_name: "values".into(),
wire_name: "values".into(),
docs: None,
ty: TypeRef::Buffer {
element: BufferElement::F64,
},
required: true,
default: None,
constraints: rspyts::ir::FieldConstraints::default(),
}],
},
}],
errors: vec![rspyts::ir::ErrorDef {
owner: owner(),
id: "sample::SampleError".into(),
name: "SampleError".into(),
docs: None,
variants: vec![],
}],
functions: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "run_sample".into(),
host_name: "runSample".into(),
docs: None,
target: Target::Both,
params: vec![rspyts::ir::ParamDef {
rust_name: "values".into(),
host_name: "values".into(),
ty: TypeRef::Buffer {
element: BufferElement::F64,
},
}],
returns: TypeRef::Named {
identity: identity("sample::ResultValue"),
},
error: Some(identity("sample::SampleError")),
}],
resources: vec![rspyts::ir::ResourceDef {
owner: owner(),
id: "sample::Runner".into(),
name: "Runner".into(),
docs: None,
target: Target::Both,
constructors: vec![rspyts::ir::FunctionDef {
owner: owner(),
rust_name: "new".into(),
host_name: "new".into(),
docs: None,
target: Target::Both,
params: vec![],
returns: TypeRef::Named {
identity: identity("sample::Runner"),
},
error: Some(identity("sample::SampleError")),
}],
methods: vec![rspyts::ir::MethodDef {
rust_name: "run_sample".into(),
host_name: "runSample".into(),
docs: None,
target: Target::Both,
mutable: true,
params: vec![],
returns: TypeRef::Named {
identity: identity("sample::ResultValue"),
},
error: Some(identity("sample::SampleError")),
}],
}],
constants: vec![],
};
let contract = resolved(manifest);
let root = std::env::temp_dir().join(format!(
"rspyts-python-source-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
fs::create_dir_all(&root).unwrap();
let names = type_names(&contract);
for (name, source) in [
("models.py", models(&contract, &names)),
("codecs.py", codecs(&contract)),
("errors.py", errors(&contract)),
("functions.py", functions(&contract, &names)),
("resources.py", resources(&contract, &names)),
("constants.py", constants(&contract, "sha256:test", &names)),
("__init__.py", init_module(&contract)),
] {
fs::write(root.join(name), source).unwrap();
}
let result = Command::new("python3")
.arg("-m")
.arg("compileall")
.arg("-q")
.arg(&root)
.output()
.expect("Python is required to verify generated Python syntax");
assert!(
result.status.success(),
"generated Python did not compile:\n{}{}",
String::from_utf8_lossy(&result.stdout),
String::from_utf8_lossy(&result.stderr)
);
fs::remove_dir_all(root).unwrap();
}
}