use std::fs;
use std::path::Path;
use anyhow::{Context, Result};
use rspyts::ir::{BufferElement, FieldDef, ScalarValue, Target, TypeDef, TypeRef, TypeShape};
use serde_json::Value;
use crate::config::{PythonConfig, PythonMode};
use crate::resolve::ResolvedContract;
use super::util::{
TypeNames, error_definition, ordered_types, pascal_case, python_doc, type_definition,
type_names,
};
use super::{write, write_json};
pub fn emit(
root: &Path,
config: &PythonConfig,
contract: &ResolvedContract,
fingerprint: &str,
native_library: Option<&Path>,
) -> 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 (index, part) in package_parts.iter().enumerate() {
package.push(part);
fs::create_dir_all(&package)?;
let init = package.join("__init__.py");
let leaf = index + 1 == package_parts.len();
if !leaf && config.mode == PythonMode::Standalone && !init.exists() {
write(&init, "")?;
}
}
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,
}),
)?;
if let Some(native_library) = native_library {
let extension = if cfg!(windows) { "pyd" } else { "so" };
let destination = package.join(format!("{}.{}", manifest.module_name, extension));
if destination.exists() || destination.is_symlink() {
anyhow::bail!(
"generated native extension {} collides with authored source",
destination.display()
);
}
fs::copy(native_library, &destination).with_context(|| {
format!(
"failed to copy native extension {} to {}",
native_library.display(),
destination.display()
)
})?;
}
Ok(())
}
fn models(contract: &ResolvedContract, names: &TypeNames) -> String {
let manifest = &contract.manifest;
let mut output = String::from(
"# Generated by rspyts 0.4. Do not edit.\n\
from __future__ import annotations\n\n\
from enum import Enum\n\
from typing import Literal, TypeAlias\n\n\
from pydantic import AwareDatetime, BaseModel, ConfigDict, Field, JsonValue\n",
);
if models_use_buffer(contract) {
output.push_str(
"\nimport numpy as np\n\
from numpy.typing import 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| format!("{}{}", item.name, pascal_case(&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 = python_doc(item.docs.as_deref(), "");
match &item.shape {
TypeShape::Struct { fields } => {
output.push_str(&format!("class {}(BaseModel):\n", item.name));
output.push_str(
" model_config = ConfigDict(populate_by_name=True, arbitrary_types_allowed=True, extra=\"forbid\", frozen=True, validate_default=True)\n",
);
if fields.is_empty() {
output.push_str(" pass\n");
}
for field in fields {
output.push_str(&python_doc(field.docs.as_deref(), " "));
output.push_str(&python_field(field, names, contract));
output.push('\n');
}
}
TypeShape::StringEnum { variants } => {
output.push_str(&format!("class {}(str, Enum):\n", item.name));
if variants.is_empty() {
output.push_str(" pass\n");
}
for variant in variants {
output.push_str(&python_doc(variant.docs.as_deref(), " "));
output.push_str(&format!(
" {} = {:?}\n",
variant.rust_name, variant.wire_name
));
}
}
TypeShape::TaggedEnum { tag, variants } => {
let mut variant_names = Vec::new();
for variant in variants {
let name = format!("{}{}", item.name, pascal_case(&variant.rust_name));
variant_names.push(name.clone());
output.push_str(&format!("class {name}(BaseModel):\n"));
output.push_str(
" model_config = ConfigDict(populate_by_name=True, arbitrary_types_allowed=True, extra=\"forbid\", frozen=True, validate_default=True)\n",
);
output.push_str(&format!(
" {}: Literal[{:?}] = {:?}\n",
tag, variant.wire_name, variant.wire_name
));
for field in &variant.fields {
output.push_str(&python_field(field, names, contract));
output.push('\n');
}
output.push('\n');
}
output.push_str(&format!(
"{}: TypeAlias = {}\n",
item.name,
variant_names.join(" | ")
));
}
TypeShape::Alias { target } => {
output.push_str(&format!(
"{}: 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 matches!(field.ty, TypeRef::Option { .. }) {
format!("Literal[{literal}] | None")
} else {
format!("Literal[{literal}]")
}
}
None => python_ref(&field.ty, names),
};
if !field.required && field.default.is_none() && !matches!(field.ty, TypeRef::Option { .. }) {
ty = format!("{ty} | None");
}
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 let Some(minimum) = field.constraints.ge {
args.push(format!("ge={minimum}"));
}
if field.rust_name != field.wire_name {
args.push(format!("alias={:?}", field.wire_name));
}
let mut output = format!(" {}: {ty}", field.rust_name);
if !args.is_empty() {
output.push_str(&format!(" = Field({})", args.join(", ")));
}
output
}
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("# Generated by rspyts 0.4. Do not edit.\n");
if !python_uses_buffer(contract) {
output.push_str("__all__: list[str] = []\n");
return output;
}
output.push_str(concat!(
"from __future__ import annotations\n\n",
"__all__: list[str] = []\n\n",
"import numpy as np\n\n",
"from . import native\n\n",
"DTYPES = {\n",
" \"u8\": np.dtype(\"u1\"),\n",
" \"i8\": np.dtype(\"i1\"),\n",
" \"u16\": np.dtype(\"<u2\"),\n",
" \"i16\": np.dtype(\"<i2\"),\n",
" \"u32\": np.dtype(\"<u4\"),\n",
" \"i32\": np.dtype(\"<i4\"),\n",
" \"u64\": np.dtype(\"<u8\"),\n",
" \"i64\": np.dtype(\"<i8\"),\n",
" \"f32\": np.dtype(\"<f4\"),\n",
" \"f64\": np.dtype(\"<f8\"),\n",
"}\n\n",
"def encode_buffer(value, dtype: str):\n",
" array = np.ascontiguousarray(value, dtype=DTYPES[dtype])\n",
" return native.BufferPayload(dtype, array.tobytes(order=\"C\"))\n\n",
"def decode_buffer(payload, dtype: str):\n",
" if payload.dtype != dtype:\n",
" raise TypeError(f\"expected {dtype} buffer, received {payload.dtype}\")\n",
" return np.frombuffer(payload.data, dtype=DTYPES[dtype], count=payload.length).copy()\n",
));
output
}
fn errors(contract: &ResolvedContract) -> String {
let manifest = &contract.manifest;
let mut output = String::from("# Generated by rspyts 0.4. Do not edit.\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(Exception):\n",
" code: str\n\n",
" def __init__(self, message: str, code: str) -> None:\n",
" super().__init__(message)\n",
" self.code = code\n\n",
"class ResourceClosedError(RuntimeError):\n",
" pass\n\n",
"def translate_error(error: RuntimeError, error_type):\n",
" if len(error.args) >= 2:\n",
" code, message = error.args[0], error.args[1]\n",
" elif error.args and isinstance(error.args[0], tuple) and len(error.args[0]) >= 2:\n",
" code, message = error.args[0][0], error.args[0][1]\n",
" else:\n",
" raise error\n",
" return error_type(str(message), str(code))\n\n",
));
for error in &manifest.errors {
output.push_str(&python_doc(error.docs.as_deref(), ""));
output.push_str(&format!(
"class {}(ContractError):\n pass\n\n",
error.name
));
}
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\
from . import {} as native\n\
from .errors import * # noqa: F403\n\
from .errors import translate_error\n\
from .models import * # noqa: F403\n\n",
manifest.module_name
);
if functions_use_adapter(contract) {
output.push_str("from pydantic import AwareDatetime, TypeAdapter\n\n");
}
if functions_use_buffer(contract) {
output.push_str("from .codecs import decode_buffer, encode_buffer\n");
output.push_str("import numpy as np\nfrom numpy.typing import 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))
{
output.push_str(&python_doc(function.docs.as_deref(), ""));
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!(
"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 result = {call}\n except RuntimeError as error:\n raise translate_error(error, {error}) from None\n"
),
None => format!(" result = {call}\n"),
};
output.push_str(&format!(
"def {}({}) -> {}:\n{} return {}\n\n",
function.rust_name,
params.join(", "),
python_ref(&function.returns, names),
body,
python_from_native(&function.returns, "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\
from . import {} as native\n\n",
manifest.module_name
);
output.push_str(
"from .errors import * # noqa: F403\n\
from .errors import translate_error\n\
from .models import * # noqa: F403\n\n",
);
if resources_use_adapter(contract) {
output.push_str("from pydantic import AwareDatetime, TypeAdapter\n");
}
if resources_use_buffer(contract) {
output.push_str("from .codecs import decode_buffer, encode_buffer\n");
output.push_str("import numpy as np\nfrom numpy.typing import 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(&python_doc(resource.docs.as_deref(), ""));
output.push_str(&format!("class {}:\n", resource.name));
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!("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 RuntimeError as error:\n raise translate_error(error, {error}) 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(", ")
));
} else {
output.push_str(" def __init__(self) -> None:\n raise TypeError(\"resource has no public constructor\")\n");
}
for constructor in constructors
.into_iter()
.filter(|constructor| Some(*constructor) != primary)
{
output.push_str(&python_doc(constructor.docs.as_deref(), " "));
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!(
"native.{}.{}({})",
resource.name,
constructor.host_name,
args.join(", ")
);
let body = match error_name(constructor.error.as_ref(), contract) {
Some(error) => format!(
" try:\n resource.handle = {call}\n except RuntimeError as error:\n raise translate_error(error, {error}) from None\n"
),
None => format!(" resource.handle = {call}\n"),
};
output.push_str(&format!(
" @classmethod\n def {}(cls{}{}) -> {}:\n resource = cls.__new__(cls)\n{body} return resource\n",
constructor.rust_name,
if params.is_empty() { "" } else { ", " },
params.join(", "),
resource.name,
));
}
for method in resource
.methods
.iter()
.filter(|method| matches!(method.target, Target::Both | Target::Python))
{
output.push_str(&python_doc(method.docs.as_deref(), " "));
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 result = {call}\n except RuntimeError as error:\n raise translate_error(error, {error}) from None\n"
),
None => format!(" result = {call}\n"),
};
output.push_str(&format!(
" def {}(self{}{}) -> {}:\n if self.handle is None:\n raise ResourceClosedError(\"{} is closed\")\n{} return {}\n",
method.rust_name,
if params.is_empty() { "" } else { ", " },
params.join(", "),
python_ref(&method.returns, names),
resource.name,
body,
python_from_native(&method.returns, "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, exception_type, exception, 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 mut output = format!(
"# Generated by rspyts 0.4. Do not edit.\n\
from __future__ import annotations\n\n\
from .models import * # noqa: F403\n\n\
CONTRACT_FINGERPRINT = {fingerprint:?}\n"
);
if constants
.iter()
.any(|constant| reference_contains_datetime(&constant.ty, contract))
{
output.push_str("from pydantic import AwareDatetime, TypeAdapter\n\n");
}
if constants
.iter()
.any(|constant| contains_buffer(&constant.ty))
{
output.push_str("import numpy as np\nfrom numpy.typing import NDArray\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(&python_doc(constant.docs.as_deref(), ""));
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(
"# Generated by rspyts 0.4. Do not edit.\n\
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| format!("{}{}", item.name, pascal_case(&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!(
"__contract_module__ = {:?}\n",
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__: list[str] = []\n\n".into();
}
format!(
"__all__ = [\n{}]\n\n",
names
.iter()
.map(|name| format!(" {name:?},\n"))
.collect::<String>()
)
}
fn python_ref(reference: &TypeRef, names: &TypeNames) -> String {
match reference {
TypeRef::Unit => "None".into(),
TypeRef::Bool => "bool".into(),
TypeRef::Int { .. } => "int".into(),
TypeRef::Float { .. } => "float".into(),
TypeRef::String => "str".into(),
TypeRef::DateTime => "AwareDatetime".into(),
TypeRef::Json => "JsonValue".into(),
TypeRef::Option { item } => format!("{} | None", python_ref(item, names)),
TypeRef::List { item } => format!("list[{}]", python_ref(item, names)),
TypeRef::Map { value } => format!("dict[str, {}]", python_ref(value, names)),
TypeRef::Tuple { items } => format!(
"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 => "bytes".into(),
TypeRef::Buffer { element } => format!("NDArray[np.{}]", numpy_scalar(*element)),
}
}
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!(
"{:?}: {}",
field.wire_name,
python_to_native(
&field.ty,
&format!("{expression}.{}", field.rust_name),
contract
)
))
.collect::<Vec<_>>()
.join(", ")
),
Some(TypeShape::TaggedEnum { .. }) => {
format!("{expression}.model_dump(mode=\"json\", by_alias=True)")
}
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 item in {expression}]",
python_to_native(item, "item", contract)
),
TypeRef::Map { value } => format!(
"{{key: {} for key, item in {expression}.items()}}",
python_to_native(value, "item", contract)
),
TypeRef::Tuple { items } => format!(
"({})",
items
.iter()
.enumerate()
.map(|(index, item)| {
python_to_native(item, &format!("{expression}[{index}]"), contract)
})
.collect::<Vec<_>>()
.join(", ")
),
TypeRef::Buffer { element } => {
format!("encode_buffer({expression}, {:?})", buffer_name(*element))
}
TypeRef::DateTime => format!("{expression}.isoformat()"),
_ => expression.to_owned(),
}
}
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 },
..
}) => {
let overrides = fields
.iter()
.map(|field| {
let key = format!("{:?}", field.wire_name);
format!(
"**({{{key}: {}}} if {key} in {expression} else {{}})",
python_from_native(
&field.ty,
&format!("{expression}[{key}]"),
contract
)
)
})
.collect::<Vec<_>>()
.join(", ");
format!("{name}.model_validate({{**{expression}, {overrides}}})")
}
Some(TypeDef {
name,
shape: TypeShape::TaggedEnum { .. },
..
}) => format!("TypeAdapter({name}).validate_python({expression})"),
Some(TypeDef {
name,
shape: TypeShape::StringEnum { .. },
..
}) => format!("{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 item in {expression}]",
python_from_native(item, "item", contract)
),
TypeRef::Map { value } => format!(
"{{key: {} for key, item in {expression}.items()}}",
python_from_native(value, "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!("decode_buffer({expression}, {:?})", buffer_name(*element))
}
TypeRef::DateTime => {
format!("TypeAdapter(AwareDatetime).validate_python({expression})")
}
_ => expression.to_owned(),
}
}
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 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 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) => {
serde_json::to_string(value).expect("a string is serializable")
}
}
}
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 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!("{name}({value:?})"),
Some(TypeDef {
shape: TypeShape::Alias { target },
..
}) => python_value(&Value::String(value.clone()), target, contract),
_ => format!("{value:?}"),
}
}
(Value::String(value), TypeRef::DateTime) => {
format!("TypeAdapter(AwareDatetime).validate_python({value:?})")
}
(Value::String(value), _) => format!("{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) => format!(
"bytes([{}])",
items
.iter()
.map(Value::to_string)
.collect::<Vec<_>>()
.join(", ")
),
(Value::Array(items), TypeRef::Buffer { .. }) => format!(
"[{}]",
items
.iter()
.map(Value::to_string)
.collect::<Vec<_>>()
.join(", ")
),
(Value::Object(items), TypeRef::Map { value }) => format!(
"{{{}}}",
items
.iter()
.map(|(key, item)| format!("{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 { .. } | TypeShape::TaggedEnum { .. },
..
}) => format!("{name}.model_validate({})", python_json(value)),
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!("{key:?}: {}", python_json(value)))
.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) => format!("{value:?}"),
Value::Array(values) => format!(
"[{}]",
values
.iter()
.map(python_json)
.collect::<Vec<_>>()
.join(", ")
),
Value::Object(values) => format!(
"{{{}}}",
values
.iter()
.map(|(key, value)| format!("{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(),
}
}
#[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("Field, JsonValue"));
assert!(!generated.contains("JsonValue: TypeAlias"));
assert!(generated.contains("Value: TypeAlias = JsonValue"));
assert!(!generated.contains("Any"));
}
#[test]
fn pydantic_json_value_validates_nested_annotation_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::AnnotationLike".into(),
name: "AnnotationLike".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(
"from models import AnnotationLike; \
payload = {'segments': [1, {'stage': None, 'flags': [True, False]}]}; \
model = AnnotationLike(metadata=payload); \
assert model.model_dump() == {'metadata': payload}",
)
.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.\n__all__: list[str] = []\n"
);
assert!(!generated_functions.contains("from .codecs import"));
assert!(generated_functions.contains("__all__ = [\n \"answer\",\n]"));
}
#[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(
"ConfigDict(populate_by_name=True, arbitrary_types_allowed=True, extra=\"forbid\", frozen=True, validate_default=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 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: "sample::Request".into(),
name: "Request".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![
field(
"contract_version",
TypeRef::Int {
signed: false,
bits: 32,
},
true,
None,
rspyts::ir::FieldConstraints {
literal: Some(ScalarValue::I64(2)),
..Default::default()
},
),
field(
"batch",
TypeRef::List {
item: Box::new(TypeRef::String),
},
true,
None,
rspyts::ir::FieldConstraints {
min_length: Some(1),
max_length: Some(200),
..Default::default()
},
),
field(
"quantity",
TypeRef::Int {
signed: false,
bits: 32,
},
false,
Some(ScalarValue::I64(1)),
rspyts::ir::FieldConstraints {
ge: Some(1),
..Default::default()
},
),
field(
"status",
TypeRef::Named {
identity: status_identity,
},
false,
Some(ScalarValue::String("unknown".into())),
Default::default(),
),
field(
"observed_at",
TypeRef::DateTime,
true,
None,
Default::default(),
),
],
},
},
],
errors: vec![],
functions: vec![],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = models(&contract, &type_names(&contract));
assert!(generated.contains("contract_version: Literal[2]"));
assert!(generated.contains("batch: list[str] = Field(min_length=1, max_length=200)"));
assert!(generated.contains("quantity: int = Field(default=1, ge=1)"));
assert!(!generated.contains("quantity: int | None"));
assert!(generated.contains("status: Status = Field(default=\"unknown\")"));
assert!(generated.contains("observed_at: AwareDatetime"));
assert!(generated.contains("validate_default=True"));
}
#[test]
fn nested_datetime_values_are_converted_at_the_native_boundary() {
let event_identity = identity("sample::Event");
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: event_identity.id.clone(),
name: "Event".into(),
docs: None,
shape: TypeShape::Struct {
fields: vec![FieldDef {
rust_name: "observed_at".into(),
wire_name: "observedAt".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: "event".into(),
host_name: "event".into(),
ty: TypeRef::Named {
identity: event_identity.clone(),
},
}],
returns: TypeRef::Named {
identity: event_identity,
},
error: None,
}],
resources: vec![],
constants: vec![],
};
let contract = resolved(manifest);
let generated = functions(&contract, &type_names(&contract));
assert!(generated.contains("event.observed_at.isoformat()"));
assert!(
generated
.contains("TypeAdapter(AwareDatetime).validate_python(result[\"observedAt\"])")
);
assert!(generated.contains("from pydantic import AwareDatetime, 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 source_mode_preserves_shared_namespace_while_standalone_initializes_it() {
let root = std::env::temp_dir().join(format!(
"rspyts-python-namespace-{}-{}",
std::process::id(),
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos()
));
let hardware = root.join("hardware");
let algorithms = root.join("algorithms");
let standalone = root.join("standalone");
for (staging, package, marker) in [
(&hardware, "neurovirtual.hardware.contracts", "hardware"),
(
&algorithms,
"neurovirtual.algorithms.contracts",
"algorithms",
),
] {
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 = {marker:?}\n"),
)
.unwrap();
let config = PythonConfig {
package: package.into(),
source: None,
mode: PythonMode::Source,
};
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",
None,
)
.unwrap();
assert!(!staging.join("python/neurovirtual/__init__.py").exists());
assert_eq!(
fs::read_to_string(parent.join("__init__.py")).unwrap(),
format!("PACKAGE_MARKER = {marker:?}\n")
);
assert!(parent.join("contracts/__init__.py").is_file());
}
let python_path =
std::env::join_paths([hardware.join("python"), algorithms.join("python")]).unwrap();
let import = Command::new("python3")
.env("PYTHONPATH", python_path)
.arg("-c")
.arg(
"import neurovirtual.hardware as hardware; \
import neurovirtual.algorithms as algorithms; \
assert hardware.PACKAGE_MARKER == 'hardware'; \
assert algorithms.PACKAGE_MARKER == 'algorithms'",
)
.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)
);
let config = PythonConfig {
package: "neurovirtual.hardware.contracts".into(),
source: None,
mode: PythonMode::Standalone,
};
emit(
&standalone,
&config,
&resolved(Manifest {
ir_version: 4,
crate_name: "standalone".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",
None,
)
.unwrap();
assert!(standalone.join("python/neurovirtual/__init__.py").is_file());
assert!(
standalone
.join("python/neurovirtual/hardware/__init__.py")
.is_file()
);
assert!(
standalone
.join("python/neurovirtual/hardware/contracts/__init__.py")
.is_file()
);
fs::remove_dir_all(root).unwrap();
}
#[test]
fn foreign_models_are_imported_from_the_configured_python_package() {
let hardware_owner = rspyts::ir::CargoPackageId::new("hardware");
let hardware_identity =
rspyts::ir::DefinitionId::new(hardware_owner.as_str(), "hardware::SignalDefinition");
let foreign = TypeDef {
owner: hardware_owner.clone(),
id: hardware_identity.id.clone(),
name: "SignalDefinition".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: "signal".into(),
wire_name: "signal".into(),
docs: None,
ty: TypeRef::Named {
identity: hardware_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([(
"hardware".into(),
crate::LockedDependency {
owner: hardware_owner,
fingerprint: "sha256:hardware".into(),
python: Some("neurovirtual.hardware.contracts".into()),
typescript: None,
types: vec![foreign.clone()],
errors: vec![],
},
)]),
foreign_types: BTreeMap::from([(hardware_identity, foreign)]),
foreign_errors: BTreeMap::new(),
};
let generated = models(&contract, &type_names(&contract));
assert!(generated.contains("from neurovirtual.hardware.contracts import SignalDefinition"));
assert!(generated.contains(" signal: SignalDefinition"));
assert!(!generated.contains("class SignalDefinition"));
}
#[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("@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();
}
}