use std::collections::{BTreeMap, BTreeSet};
use std::path::{Path, PathBuf};
use heck::{ToShoutySnakeCase, ToSnakeCase, ToUpperCamelCase};
use indexmap::IndexMap;
use crate::error::{Error, Result};
use crate::generator::json_schema::python as python_json;
use crate::generator::proto::python as python_proto;
use crate::generator::render_request_plan;
use crate::generator::{
ExternalModelBackend, GeneratedFileMap, GeneratedFileOrigin, GeneratedFiles, GenerationMode,
};
use crate::language::Language;
use crate::planning::{
PlannedFamily, PlannedOperationResourceFieldBinding, PlannedOperationResourceReturn,
PlannedProtoType, PlannedProtoTypeInfo, PlannedRecordType, PlannedResource,
PlannedResourceMethod, PlannedResourceMethodBindingSpec, PlannedResourceMethodResultKind,
PlannedSpec, PlannedType, message_model_name, operation_input_model,
operation_output_direct_result,
};
use crate::planning::{RequestPlan, ResolvedResourceBindingSource};
use crate::spec::{
AliasTypeSpec, EnumSpec, ExternalTypeSourceSpec, ExternalTypeSpec, FlagsSpec, FunctionArgsSpec,
FunctionFieldSpec, FunctionResultSpec, LanguageImportSpec, LanguageImportStyle,
LanguageStringSpec, ModulePath, OperationSpec, RecordFieldSpec, RecordFieldVisibility,
RecordSpec, SupportFragmentSpec, TypeDeclSpec, TypeReplacementSpec, TypeSpec, VariantSpec,
};
use crate::spec::{ApiSpecBranch, ApiSpecNode};
pub(in crate::generator) const GENERATED_HEADER: &str = concat!(
"# Generated by nexgen v",
env!("CARGO_PKG_VERSION"),
". DO NOT EDIT!"
);
const PYTHON_FORMAT_LINE_LENGTH: usize = 88;
const EXPERIMENTAL_WARNING: &str = "This API is experimental and subject to change.";
pub(in crate::generator) type RootPackageImports = BTreeMap<String, BTreeSet<String>>;
struct PythonGenerationResult {
files: GeneratedFileMap,
warnings: Vec<String>,
root_package_imports: RootPackageImports,
exported_names: BTreeSet<String>,
}
pub(crate) fn generate(
tree: &crate::spec::ApiSpecTree<PlannedFamily>,
support: &crate::SupportFiles,
) -> Result<GeneratedFiles> {
match &tree.root {
ApiSpecNode::Leaf(leaf) => {
let support_fragments = support_fragments_for_plan(&leaf.spec, support);
let generated = generate_leaf(
&leaf.spec,
&support_fragments,
GeneratedFileOrigin::fixed("generated Python package module"),
)?;
Ok(GeneratedFiles {
layout: crate::generator::GeneratedOutputLayout::Directory,
files: generated.files.into_files(),
warnings: generated.warnings,
})
}
ApiSpecNode::Branch(branch) => generate_tree(branch, support),
}
}
fn generate_leaf(
api_plan: &PlannedSpec,
support_fragments: &[SupportFragmentSpec],
module_origin: GeneratedFileOrigin,
) -> Result<PythonGenerationResult> {
reject_support_namespaces(Language::Python, support_fragments)?;
let inline_model_rebuilds = api_plan
.data
.module_imports
.values()
.all(BTreeSet::is_empty);
ApiPlanner::new(api_plan, inline_model_rebuilds, None)?.build(support_fragments, module_origin)
}
fn generate_leaf_with_model_hoists(
api_plan: &PlannedSpec,
support_fragments: &[SupportFragmentSpec],
model_hoists: &PythonModelHoists,
module_origin: GeneratedFileOrigin,
) -> Result<PythonGenerationResult> {
reject_support_namespaces(Language::Python, support_fragments)?;
ApiPlanner::new(api_plan, true, Some(model_hoists))?.build(support_fragments, module_origin)
}
fn generate_tree(
branch: &ApiSpecBranch<PlannedFamily>,
support: &crate::SupportFiles,
) -> Result<GeneratedFiles> {
let model_hoists = tree_model_hoists(branch)?;
let mut files = GeneratedFileMap::default();
let mut warnings = Vec::new();
let mut root_package_imports = RootPackageImports::new();
files.insert_multi(
render_tree_support_files(branch),
GeneratedFileOrigin::fixed(format!(
"generated Python validation runtime for {}",
source_paths_label(&tree_json_source_paths(branch))
)),
)?;
for (path, contents) in model_hoists.files() {
files.insert(
path.clone(),
contents.clone(),
GeneratedFileOrigin::fixed(format!(
"generated Python recursive-model module for {}",
source_paths_label(model_hoists.source_paths())
)),
)?;
}
let mut child_exports = BTreeMap::new();
for (name, node) in &branch.children {
let exported_names = generate_tree_node(
node,
support,
&model_hoists,
&mut files,
&mut warnings,
&mut root_package_imports,
)?;
child_exports.insert(name.clone(), exported_names);
}
insert_branch_index_file(
&mut files,
branch,
&model_hoists,
&child_exports,
&root_package_imports,
)?;
Ok(GeneratedFiles {
layout: crate::generator::GeneratedOutputLayout::Directory,
files: files.into_files(),
warnings,
})
}
fn generate_tree_node(
node: &ApiSpecNode<PlannedFamily>,
support: &crate::SupportFiles,
model_hoists: &PythonModelHoists,
files: &mut GeneratedFileMap,
warnings: &mut Vec<String>,
root_package_imports: &mut RootPackageImports,
) -> Result<BTreeSet<String>> {
match node {
ApiSpecNode::Leaf(leaf) => {
let support_fragments = support_fragments_for_plan(&leaf.spec, support);
let generated = generate_leaf_with_model_hoists(
&leaf.spec,
&support_fragments,
model_hoists,
GeneratedFileOrigin::input_module(Language::Python, &leaf.source_path),
)?;
extend_root_package_imports(root_package_imports, generated.root_package_imports);
warnings.extend(generated.warnings);
let prefix = leaf.module_path.to_path_buf();
files.extend(generated.files.prefix(prefix)?)?;
Ok(generated.exported_names)
}
ApiSpecNode::Branch(branch) => {
let mut child_exports = BTreeMap::new();
for (name, node) in &branch.children {
let exported_names = generate_tree_node(
node,
support,
model_hoists,
files,
warnings,
root_package_imports,
)?;
child_exports.insert(name.clone(), exported_names);
}
insert_branch_index_file(
files,
branch,
model_hoists,
&child_exports,
&RootPackageImports::new(),
)?;
Ok(child_exports.into_values().flatten().collect())
}
}
}
fn insert_branch_index_file(
files: &mut GeneratedFileMap,
branch: &ApiSpecBranch<PlannedFamily>,
model_hoists: &PythonModelHoists,
child_exports: &BTreeMap<String, BTreeSet<String>>,
root_package_imports: &RootPackageImports,
) -> Result<()> {
let mut path = branch.module_path.to_path_buf();
path.push("__init__.py");
let mut contents = String::from(GENERATED_HEADER);
contents.push_str("\n\n");
let mut wrote_import = false;
if !root_package_imports.is_empty() {
render_root_package_imports(&mut contents, root_package_imports);
wrote_import = true;
}
for name in branch.children.keys() {
let module_name = name.replace('-', "_");
let names = child_exports
.get(name)
.into_iter()
.flatten()
.cloned()
.into_iter()
.collect::<Vec<_>>();
if names.is_empty() {
continue;
}
render_named_python_import(&mut contents, &format!(".{module_name}"), &names);
wrote_import = true;
}
if branch.module_path.is_root() && !model_hoists.is_empty() {
if wrote_import {
contents.push('\n');
}
render_named_python_import(
&mut contents,
"._recursive",
&model_hoists
.exported_names()
.iter()
.cloned()
.collect::<Vec<_>>(),
);
}
let mut export_names = child_exports
.values()
.flat_map(|names| names.iter().cloned())
.collect::<BTreeSet<_>>();
if branch.module_path.is_root() {
export_names.extend(model_hoists.exported_names().iter().cloned());
}
export_names.extend(root_package_export_names(root_package_imports));
if !export_names.is_empty() {
contents.push_str("\n__all__ = [\n");
for name in export_names {
contents.push_str(" ");
contents.push_str(&python_string_literal(&name));
contents.push_str(",\n");
}
contents.push_str("]\n");
}
files.insert(
path,
contents,
GeneratedFileOrigin::fixed(format!(
"generated Python package initializer for module `{}`",
branch.module_path.as_module_key()
)),
)
}
fn planned_module_export_model_names(plan: &PlannedSpec) -> BTreeSet<String> {
plan.types
.values()
.filter(|entry| entry.is_module_export())
.filter_map(|entry| match &entry.declaration {
TypeDeclSpec::Record(record) => Some(record.name.clone()),
TypeDeclSpec::Enum(enumeration) => Some(enumeration.name.clone()),
TypeDeclSpec::Flags(flags) => Some(flags.name.clone()),
TypeDeclSpec::Variant(variant) => Some(variant.name.clone()),
TypeDeclSpec::External(_) => None,
})
.collect()
}
fn support_fragments_for_plan(
plan: &PlannedSpec,
support: &crate::SupportFiles,
) -> Vec<SupportFragmentSpec> {
if support.fragments.is_empty() {
plan.support
.fragments_for_language(Language::Python)
.to_vec()
} else {
support.fragments.clone()
}
}
struct ApiPlanner<'a> {
api_plan: &'a PlannedSpec,
inline_model_rebuilds: bool,
model_hoists: Option<&'a PythonModelHoists>,
external_models: PythonExternalModels,
language_imports: Vec<LanguageImportSpec>,
enums: IndexMap<String, RenderedEnum>,
flags: IndexMap<String, RenderedFlags>,
variants: IndexMap<String, RenderedVariant>,
models: IndexMap<String, RenderedModel>,
}
#[derive(Debug, Default)]
struct PythonExternalModels {
proto: python_proto::ModelBackend,
json: python_json::ModelBackend,
}
impl PythonExternalModels {
fn new(api_plan: &PlannedSpec) -> Result<Self> {
let mut this = Self::default();
this.prepare(api_plan)?;
Ok(this)
}
fn new_with_hoists(api_plan: &PlannedSpec, hoists: &PythonModelHoists) -> Result<Self> {
let mut this = Self::default();
this.proto.prepare(api_plan)?;
this.json.prepare_with_hoists(api_plan, hoists)?;
Ok(this)
}
fn render_model_fragments(
&self,
models: &[&RenderedModel],
variants: &[&RenderedVariant],
api_plan: &PlannedSpec,
) -> Result<RenderedModelFragments> {
let mut fragments = RenderedModelFragments::default();
fragments.extend(render_record_models(models, api_plan, self)?);
fragments.extend(self.json.render_models()?);
fragments.extend(self.proto.render_variant_models(
api_plan,
variants,
&fragments.declared_type_parameters,
));
Ok(fragments)
}
fn owns_variant(&self, full_name: &str) -> bool {
self.proto.owns_variant(full_name)
}
fn render_support_files(&self) -> Result<BTreeMap<PathBuf, String>> {
let mut files = BTreeMap::new();
files.extend(self.json.render_support_files()?);
Ok(files)
}
fn service_model_ref(
&self,
model_type: &PlannedType,
native_module: &str,
api_plan: &PlannedSpec,
) -> (String, Option<String>) {
match model_type {
PlannedType::External(ExternalTypeSpec::Proto(_)) => {
if let Some(conversion) = self.proto.wire_conversion(model_type, None) {
return (
conversion.annotation.clone(),
python_qualified_module_paths(&conversion.annotation)
.into_iter()
.next(),
);
}
let reference = self
.proto
.service_wire_model_ref(model_type)
.expect("proto wire model ref should exist");
(reference.type_ref, Some(reference.module_path))
}
PlannedType::External(ExternalTypeSpec::Json(json_type)) => {
let type_name = self
.json
.model_type_annotation(json_type)
.expect("json model annotation should exist");
if self.json.is_hoisted(json_type) {
(
type_name,
Some(root_python_model_hoist_module(&api_plan.module_path)),
)
} else {
(
format!("{native_module}.{type_name}"),
Some(format!(".{native_module}")),
)
}
}
PlannedType::Record(record) => (
format!("{native_module}.{}", record.model_name),
Some(format!(".{native_module}")),
),
_ => panic!("operation service ref should be model-shaped"),
}
}
fn resolved_external_field_type(&self, value_type: &PlannedType) -> Option<ResolvedFieldType> {
match value_type {
PlannedType::External(ExternalTypeSpec::Proto(PlannedProtoType::Enum(_))) => {
self.proto.enum_field_type(value_type)
}
PlannedType::External(ExternalTypeSpec::Json(_)) => {
let conversion = self.wire_conversion(value_type, None)?;
Some(ResolvedFieldType {
annotation: conversion.annotation.clone(),
imports: conversion.imports.clone(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: Some(conversion),
})
}
_ => None,
}
}
fn render_record_wire_block(
&self,
api_plan: &PlannedSpec,
model: &RenderedModel,
planned_model: &RecordSpec<PlannedFamily>,
) -> Result<Option<RenderedRecordWireBlock>> {
self.proto
.render_record_wire_block(api_plan, model, planned_model, &|value_type| {
resolve_python_value_type(api_plan, self, value_type)
})
}
}
impl ExternalModelBackend for PythonExternalModels {
type ModelFragments = RenderedModelFragments;
type WireConversion = WireValueConversion;
fn prepare(&mut self, api_plan: &PlannedSpec) -> Result<()> {
self.proto.prepare(api_plan)?;
self.json.prepare(api_plan)
}
fn render_models(&self) -> Result<RenderedModelFragments> {
self.json.render_models()
}
fn model_type_annotation(&self, model_type: &PlannedType) -> Option<String> {
match model_type {
PlannedType::External(ExternalTypeSpec::Proto(_)) | PlannedType::Record(_) => {
if let Some(conversion) = self.proto.wire_conversion(model_type, None) {
return Some(conversion.annotation);
}
self.proto.model_type_annotation(model_type)
}
PlannedType::External(ExternalTypeSpec::Json(json_type)) => {
self.json.model_type_annotation(json_type)
}
_ => None,
}
}
fn wire_type_identifier(&self, model_type: &PlannedType) -> Option<String> {
match model_type {
PlannedType::External(ExternalTypeSpec::Proto(_)) | PlannedType::Record(_) => {
self.proto.wire_type_identifier(model_type)
}
PlannedType::External(ExternalTypeSpec::Json(json_type)) => {
self.json.wire_type_identifier(json_type)
}
_ => None,
}
}
fn wire_conversion(
&self,
model_type: &PlannedType,
planned_record: Option<&RecordSpec<PlannedFamily>>,
) -> Option<WireValueConversion> {
match model_type {
PlannedType::External(ExternalTypeSpec::Proto(_)) | PlannedType::Record(_) => {
self.proto.wire_conversion(model_type, planned_record)
}
PlannedType::External(ExternalTypeSpec::Json(json_type)) => {
self.json.wire_conversion(json_type, planned_record)
}
_ => None,
}
}
}
impl<'a> ApiPlanner<'a> {
fn new(
api_plan: &'a PlannedSpec,
inline_model_rebuilds: bool,
model_hoists: Option<&'a PythonModelHoists>,
) -> Result<Self> {
let external_models = if let Some(model_hoists) = model_hoists {
PythonExternalModels::new_with_hoists(api_plan, model_hoists)?
} else {
PythonExternalModels::new(api_plan)?
};
Ok(Self {
api_plan,
inline_model_rebuilds,
model_hoists,
external_models,
language_imports: collect_python_language_imports(api_plan),
enums: IndexMap::new(),
flags: IndexMap::new(),
variants: IndexMap::new(),
models: IndexMap::new(),
})
}
fn build(
mut self,
support_fragments: &[SupportFragmentSpec],
module_origin: GeneratedFileOrigin,
) -> Result<PythonGenerationResult> {
let api_plan = self.api_plan;
let services = api_plan
.services
.iter()
.map(|service| {
let operations = service
.operations
.iter()
.map(|operation| self.resolve_operation(operation))
.collect::<Result<Vec<_>>>()?;
Ok(RenderedService {
name: service
.code_name
.for_language(Language::Python)
.unwrap_or(&service.name),
wire_name: &service.wire_name,
doc: service
.doc
.for_language(Language::Python)
.map(str::to_string),
endpoint: service.endpoint.clone(),
experimental: service.experimental,
deprecated: service.deprecated,
delay_load_temporalio_workflow: service.delay_load_temporalio_workflow,
operations,
resources: service
.resources
.iter()
.map(|resource| resource.data.clone())
.collect(),
})
})
.collect::<Result<Vec<_>>>()?;
for service in &services {
for resource in &service.resources {
self.ensure_resource_field_types(resource)?;
}
}
for record in api_plan.records().map(|(_, record)| record) {
let model_type = TypeSpec::Record(PlannedRecordType {
full_name: record.full_name.clone(),
model_name: record.name.clone(),
});
self.resolve_message_value_conversion(&model_type)?;
}
let model_refs = self.models.values().collect::<Vec<_>>();
let variant_refs = self.variants.values().collect::<Vec<_>>();
let model_fragments =
self.render_model_fragments(model_refs.as_slice(), variant_refs.as_slice())?;
validate_python_generated_names(self.api_plan, &model_fragments.generated_names)?;
let (files, exported_names) = self.render_package(
&model_fragments,
&services,
support_fragments,
module_origin,
)?;
Ok(PythonGenerationResult {
files,
warnings: Vec::new(),
root_package_imports: model_fragments.root_package_imports,
exported_names,
})
}
fn render_model_fragments(
&self,
models: &[&RenderedModel],
variants: &[&RenderedVariant],
) -> Result<RenderedModelFragments> {
self.external_models
.render_model_fragments(models, variants, self.api_plan)
}
fn render_package(
&self,
model_fragments: &RenderedModelFragments,
services: &[RenderedService<'_>],
support_fragments: &[SupportFragmentSpec],
module_origin: GeneratedFileOrigin,
) -> Result<(GeneratedFileMap, BTreeSet<String>)> {
let mode = crate::nexgen_config::current().mode;
let mut files = GeneratedFileMap::default();
render_support_package(&mut files, support_fragments)?;
files.insert_multi(
self.external_models.render_support_files()?,
GeneratedFileOrigin::fixed("generated Python external-model runtime"),
)?;
let variants = self
.variants
.iter()
.filter(|(full_name, _)| !self.external_models.owns_variant(full_name))
.map(|(_, variant)| variant)
.collect::<Vec<_>>();
let mut model_names = self
.enums
.values()
.map(|enumeration| enumeration.name.clone())
.chain(self.flags.values().map(|flag_set| flag_set.name.clone()))
.chain(variants.iter().map(|variant| variant.name.clone()))
.chain(model_fragments.exported_names.iter().cloned())
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
sort_python_model_names(&mut model_names, &self.variants, model_fragments);
let operation_model_names = model_names
.iter()
.cloned()
.chain(planned_module_import_names(self.api_plan))
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let resource_operation_owners = resource_operation_owners(services);
let has_standalone_operations = mode == GenerationMode::NativeApi
&& services.iter().any(|service| {
service.endpoint.is_some()
&& service.operations.iter().any(|operation| {
!resource_operation_owners
.contains_key(&operation_key(service.name, operation.name))
})
});
let resource_names = services
.iter()
.flat_map(|service| {
service
.resources
.iter()
.map(|resource| resource.type_name.clone())
})
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let support_names = support_export_names(support_fragments);
let rendered_model_names = model_names.iter().cloned().collect::<BTreeSet<_>>();
let mut package_model_names = if self.model_hoists.is_some() {
rendered_model_names.clone()
} else {
BTreeSet::new()
};
package_model_names.extend(
planned_module_export_model_names(self.api_plan)
.intersection(&rendered_model_names)
.cloned(),
);
package_model_names.extend(
model_fragments
.module_exported_names
.intersection(&rendered_model_names)
.cloned(),
);
let mut package_model_names = package_model_names.into_iter().collect::<Vec<_>>();
sort_python_model_names(&mut package_model_names, &self.variants, model_fragments);
let render_init = |root_package_imports: &RootPackageImports| {
if mode == GenerationMode::NativeApi {
render_package_init(
services,
&package_model_names,
&resource_names,
&resource_operation_owners,
&support_names,
self.api_plan,
self.model_hoists,
root_package_imports,
)
} else {
render_definitions_only_package_init(services, &model_names, root_package_imports)
}
};
let empty_root_package_imports = RootPackageImports::new();
let root_package_imports = if self.model_hoists.is_none() {
&model_fragments.root_package_imports
} else {
&empty_root_package_imports
};
let exported_names = package_export_names(
services,
if mode == GenerationMode::NativeApi {
&package_model_names
} else {
&model_names
},
root_package_imports,
);
files.insert(
"__init__.py",
render_init(root_package_imports),
module_origin.clone(),
)?;
// A module that declares nothing emits no models file. Emitting the file
// anyway leaves a header and unused imports behind.
if let Some(models_source) = render_models_module(
self.enums.values().collect::<Vec<_>>().as_slice(),
self.flags.values().collect::<Vec<_>>().as_slice(),
variants.as_slice(),
model_fragments,
&support_names,
&self.language_imports,
self.api_plan,
self.inline_model_rebuilds,
self.model_hoists,
)? {
files.insert("models.py", models_source, module_origin.clone())?;
}
if !resource_names.is_empty() {
files.insert(
"_resources/__init__.py",
render_resources_package_init(services),
GeneratedFileOrigin::fixed("generated Python resources package initializer"),
)?;
}
if !services.is_empty() {
files.insert(
"services.py",
render_service_module(
services,
self.api_plan,
self.model_hoists,
mode == GenerationMode::NativeApi,
),
module_origin.clone(),
)?;
}
if has_standalone_operations {
files.insert(
"operations/__init__.py",
render_operations_package_init(),
GeneratedFileOrigin::fixed("generated Python operations package initializer"),
)?;
}
if crate::nexgen_config::current().system_nexus && mode == GenerationMode::NativeApi {
files.insert(
"_system_nexus_interceptor.py",
render_system_nexus_interceptor(services),
GeneratedFileOrigin::fixed("generated Python system Nexus interceptor"),
)?;
}
for service in services {
for resource in &service.resources {
let bound_operations = if mode == GenerationMode::NativeApi {
resource_bound_operations(service, resource)
} else {
Vec::new()
};
files.insert(
format!("_resources/{}.py", resource_module_name(resource)),
self.render_resource_module_file(
service,
resource,
&bound_operations,
&operation_model_names,
&resource_names,
&support_names,
),
GeneratedFileOrigin::resource(Language::Python, service.name, &resource.name),
)?;
}
if mode == GenerationMode::NativeApi {
if service.endpoint.is_none() {
continue;
}
for operation in &service.operations {
if resource_operation_owners
.contains_key(&operation_key(service.name, operation.name))
{
continue;
}
files.insert(
format!("operations/{}.py", operation.attr_name),
render_operation_module(
service,
operation,
&operation_model_names,
&resource_names,
&support_names,
&self.language_imports,
self.api_plan,
self.model_hoists,
),
GeneratedFileOrigin::operation(
Language::Python,
service.name,
operation.name,
),
)?;
}
}
}
Ok((files, exported_names))
}
fn render_resource_module_file(
&self,
service: &RenderedService<'_>,
resource: &PlannedResource,
bound_operations: &[ResourceBoundOperation<'_>],
model_names: &[String],
resource_names: &[String],
support_names: &[String],
) -> String {
let mut module_imports = bound_operations
.iter()
.flat_map(|bound_operation| {
[
bound_operation
.operation
.input
.as_ref()
.and_then(|input| input.module_path.clone()),
bound_operation.operation.output_module_path.clone(),
]
})
.flatten()
.collect::<BTreeSet<_>>();
let mut body = String::new();
let function_fields = bound_operations
.iter()
.filter_map(|bound_operation| bound_operation.operation.unpacked_input.as_ref())
.flat_map(|unpacked_input| unpacked_input.functions.iter().cloned())
.collect::<Vec<_>>();
let output_type_parameters = bound_operations
.iter()
.flat_map(|bound_operation| bound_operation.operation.output_type_parameters.iter())
.cloned()
.collect::<BTreeSet<_>>();
if !function_fields.is_empty() {
render_function_type_parameter_definitions(
&mut body,
&function_fields,
&output_type_parameters,
);
body.push_str("\n\n");
}
self.render_resource(&mut body, service, resource, bound_operations);
if !bound_operations.is_empty() && service.endpoint.is_some() {
body.push_str("\n\n");
for (index, bound_operation) in bound_operations.iter().enumerate() {
render_operation_functions(&mut body, service, bound_operation.operation);
if index + 1 != bound_operations.len() {
body.push_str("\n\n");
}
}
}
if service.endpoint.is_none() {
body.push_str("\n\n");
render_resource_client_binder(&mut body, resource);
}
if !service.delay_load_temporalio_workflow {
module_imports.insert("temporalio.workflow".to_string());
}
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
let type_checking_names = temporalio_workflow_type_checking_names(&body);
let body_for_imports =
if service.delay_load_temporalio_workflow && !type_checking_names.is_empty() {
format!("{body}\ntyping.TYPE_CHECKING")
} else {
body.clone()
};
let skipped_language_modules =
delayed_temporalio_workflow_skipped_language_modules(service);
let wrote_imports = render_optional_python_imports_with_skipped_language_modules(
&mut output,
&body_for_imports,
&module_imports,
&self.language_imports,
&skipped_language_modules,
);
if wrote_imports
&& service.delay_load_temporalio_workflow
&& !type_checking_names.is_empty()
{
output.push('\n');
}
let wrote_type_checking_imports =
render_temporalio_workflow_type_checking_imports(&mut output, service, &body);
if wrote_imports || wrote_type_checking_imports {
output.push('\n');
}
let used_model_names = used_python_symbol_imports(&body, model_names);
if !used_model_names.is_empty() {
render_model_name_imports(
&mut output,
"..models",
&self.api_plan.module_path.child("_resources"),
self.api_plan,
&used_model_names,
self.model_hoists,
);
}
let used_resource_names = used_python_symbol_imports(&body, resource_names)
.into_iter()
.filter(|name| name != &resource.type_name)
.collect::<Vec<_>>();
if !used_resource_names.is_empty() {
render_named_python_import(&mut output, ".", &used_resource_names);
}
let used_support_names = used_python_symbol_imports(&body, support_names);
if !used_support_names.is_empty() {
render_named_python_import(&mut output, ".._support", &used_support_names);
}
if !body.is_empty() {
output.push('\n');
output.push('\n');
output.push_str(&body);
}
output
}
fn render_resource(
&self,
output: &mut String,
service: &RenderedService<'_>,
resource: &PlannedResource,
bound_operations: &[ResourceBoundOperation<'_>],
) {
output.push_str("@dataclasses.dataclass\n");
output.push_str("class ");
output.push_str(&resource.type_name);
output.push_str(":\n");
for field in &resource.fields {
output.push_str(" ");
output.push_str(&python_field_name(&field.name));
output.push_str(": ");
output.push_str(&self.python_resource_field_annotation(
&field.kind,
field.optional,
field.function.as_ref(),
));
output.push('\n');
}
if resource.methods.is_empty() {
output.push_str("\n pass\n");
return;
}
for method in &resource.methods {
output.push_str("\n");
self.render_resource_class_method(output, service, resource, method, bound_operations);
}
}
fn render_resource_class_method(
&self,
output: &mut String,
service: &RenderedService<'_>,
resource: &PlannedResource,
method: &PlannedResourceMethod,
bound_operations: &[ResourceBoundOperation<'_>],
) {
let result_annotation = self.python_resource_method_result_annotation(method);
output.push_str(" async def ");
output.push_str(&python_field_name(&method.name));
output.push_str("(\n");
output.push_str(" self,\n");
for param in &method.params {
output.push_str(" ");
output.push_str(&python_field_name(¶m.name));
output.push_str(": ");
output.push_str(&self.python_resource_field_annotation(
¶m.kind,
param.optional,
param.function.as_ref(),
));
if param.optional {
output.push_str(" = None");
}
output.push_str(",\n");
}
output.push_str(" ) -> ");
output.push_str(&result_annotation);
output.push_str(":\n");
match &method.binding {
PlannedResourceMethodBindingSpec::Operation { operation_name, .. } => {
let operation = bound_operations
.iter()
.find(|bound_operation| bound_operation.operation.name == operation_name)
.map(|bound_operation| bound_operation.operation)
.expect("bound resource operation should be rendered in the same module");
if service.endpoint.is_none() {
render_resource_method_inline_operation_body(
output,
service,
method,
operation,
&result_annotation,
);
} else {
render_resource_method_operation_body(
output,
service,
method,
operation,
&result_annotation,
);
}
}
PlannedResourceMethodBindingSpec::Stub => {
output.push_str(" raise NotImplementedError(");
output.push_str(&python_string_literal(&format!(
"{}.{} is not yet implemented",
resource.name,
python_field_name(&method.name)
)));
output.push_str(")\n");
}
}
}
fn python_resource_method_result_annotation(&self, method: &PlannedResourceMethod) -> String {
let Some(result) = &method.result else {
return "None".to_string();
};
let annotation = match &result.kind {
PlannedResourceMethodResultKind::Resource { type_name } => type_name.clone(),
PlannedResourceMethodResultKind::Value(kind) => {
self.python_resource_field_annotation(kind, result.optional, None)
}
};
if result.optional && !annotation.contains("| None") {
format!("{annotation} | None")
} else {
annotation
}
}
fn python_resource_field_annotation(
&self,
kind: &PlannedType,
optional: bool,
function: Option<&FunctionFieldSpec<PlannedFamily>>,
) -> String {
let base = if let Some(function) = function {
python_function_field_annotation(
function,
function
.alternate_type
.as_ref()
.map(python_authored_type_annotation),
)
} else {
match kind {
PlannedType::Option(value) | PlannedType::List(value) => {
format!(
"collections.abc.Sequence[{}]",
self.python_resource_value_annotation(value)
)
}
PlannedType::Map(key, value) => format!(
"collections.abc.Mapping[{}, {}]",
self.python_resource_value_annotation(key),
self.python_resource_value_annotation(value)
),
value => self.python_resource_value_annotation(value),
}
};
if optional {
format!("{base} | None")
} else {
base
}
}
fn python_resource_value_annotation(&self, value: &PlannedType) -> String {
match value {
PlannedType::Float => "float".to_string(),
PlannedType::Int(_) => "int".to_string(),
PlannedType::Bool => "bool".to_string(),
PlannedType::String => "str".to_string(),
PlannedType::Bytes => "bytes".to_string(),
PlannedType::TypeParameter(parameter) => parameter.name.clone(),
PlannedType::Enum(enum_type) => enum_type.name.clone(),
PlannedType::External(ExternalTypeSpec::Proto(_))
| PlannedType::External(ExternalTypeSpec::Json(_)) => self
.external_models
.model_type_annotation(value)
.expect("external resource value annotation should exist"),
PlannedType::Flags(flags_type) => flags_type.name.clone(),
PlannedType::Variant(variant_type) => variant_type.name.clone(),
PlannedType::Record(record) => record.model_name.clone(),
PlannedType::Resource(resource) => resource.type_name.clone(),
PlannedType::Option(inner) | PlannedType::List(inner) => {
format!(
"collections.abc.Sequence[{}]",
self.python_resource_value_annotation(inner)
)
}
PlannedType::Map(key, value) => format!(
"collections.abc.Mapping[{}, {}]",
self.python_resource_value_annotation(key),
self.python_resource_value_annotation(value)
),
PlannedType::Tuple(items) => format!(
"tuple[{}]",
items
.iter()
.map(|item| self.python_resource_value_annotation(item))
.collect::<Vec<_>>()
.join(", ")
),
PlannedType::Result { ok, err } => python_result_annotation(
ok.as_ref()
.map(|ok| self.python_resource_value_annotation(ok).to_string()),
err.as_ref()
.map(|err| self.python_resource_value_annotation(err).to_string()),
),
PlannedType::External(ExternalTypeSpec::Alias(AliasTypeSpec {
type_name,
target,
..
})) => type_name
.for_language(Language::Python)
.map(str::to_string)
.unwrap_or_else(|| self.python_resource_value_annotation(target)),
}
}
fn ensure_resource_field_types(&mut self, resource: &PlannedResource) -> Result<()> {
for field in &resource.fields {
self.ensure_resource_field_type(&field.kind)?;
}
for method in &resource.methods {
for param in &method.params {
self.ensure_resource_field_type(¶m.kind)?;
}
}
Ok(())
}
fn ensure_resource_field_type(&mut self, kind: &PlannedType) -> Result<()> {
match kind {
PlannedType::List(value) => {
self.resolve_planned_value_type(value)?;
}
PlannedType::Map(key, value) => {
self.resolve_planned_value_type(key)?;
self.resolve_planned_value_type(value)?;
}
value => {
self.resolve_planned_value_type(value)?;
}
}
Ok(())
}
fn resolve_operation<'operation>(
&mut self,
operation: &'operation OperationSpec<PlannedFamily>,
) -> Result<RenderedOperation<'operation>> {
let output_resource_return = operation.data.output_resource_return.clone();
let input = operation_input_model(operation);
let rendered_input = input
.map(|input| -> Result<_> {
let (type_ref, module_path) =
self.external_models
.service_model_ref(input, "models", self.api_plan);
let input_conversion = self.resolve_message_value_conversion(input)?;
let annotation = if let PlannedType::Record(record) = input {
let parameters = self
.api_plan
.record_type_parameters(&record.full_name, Language::Python);
python_generic_record_annotation(&input_conversion.annotation, ¶meters)
} else {
input_conversion.annotation.clone()
};
Ok(RenderedOperationInput {
descriptor_type_ref: if let PlannedType::Record(record) = input {
python_erased_generic_record_annotation(
&type_ref,
&self
.api_plan
.record_type_parameters(&record.full_name, Language::Python),
)
} else {
type_ref.clone()
},
type_ref,
module_path,
annotation,
supports_unpacked: input_conversion.supports_unpacked_input(),
})
})
.transpose()?;
let output_transform = operation.output_transform.as_ref();
let output_direct_result = operation_output_direct_result(operation);
let (output_ref, output_module_path, output_annotation_default) = match operation
.output_type()
{
Some(
output_model @ (PlannedType::External(ExternalTypeSpec::Proto(
PlannedProtoType::Message(_),
))
| PlannedType::External(ExternalTypeSpec::Json(_))
| PlannedType::Record(_)),
) => {
let (output_ref, output_module_path) =
self.external_models
.service_model_ref(output_model, "models", self.api_plan);
if output_transform.is_none()
&& output_resource_return.is_none()
&& !output_direct_result
&& matches!(output_model, PlannedType::Record(_))
{
self.resolve_message_value_conversion(output_model)?;
}
let annotation = self.resolve_output_annotation(output_model)?;
(output_ref, output_module_path, annotation)
}
Some(PlannedType::Resource(resource)) => {
if let Some(output) = &resource.wire_type {
let output = planned_record_for_external_source(self.api_plan, output)
.map(planned_record_type)
.unwrap_or_else(|| PlannedType::External(output.clone()));
let (output_ref, output_module_path) =
self.external_models
.service_model_ref(&output, "models", self.api_plan);
let conversion = self.resolve_message_value_conversion(&output)?;
(output_ref, output_module_path, conversion.annotation)
} else {
(
resource.type_name.clone(),
Some("._resources".to_string()),
resource.type_name.clone(),
)
}
}
None => ("None".to_string(), None, "None".to_string()),
Some(_) => {
panic!("planned operation output should be proto, record, resource, or none")
}
};
let model_input_parameters = input
.map(|input| self.api_plan.type_parameters(input, Language::Python))
.unwrap_or_default();
let model_input_identities = model_input_parameters
.iter()
.map(|usage| usage.parameter.full_name.as_str())
.collect::<BTreeSet<_>>();
let model_output_parameters = operation
.output_type()
.map(|output| self.api_plan.type_parameters(output, Language::Python))
.unwrap_or_default();
let generic_model_types =
!model_input_parameters.is_empty() || !model_output_parameters.is_empty();
let model_output_only_parameters = model_output_parameters
.into_iter()
.filter(|usage| !model_input_identities.contains(usage.parameter.full_name.as_str()))
.map(|usage| usage.parameter.name)
.collect::<BTreeSet<_>>();
let output_annotation_default =
erase_python_type_parameters(&output_annotation_default, &model_output_only_parameters);
let overload_output_annotation = output_transform
.and_then(|transform| {
transform
.type_name
.for_language(crate::language::Language::Python)
.map(str::to_string)
})
.or_else(|| {
output_resource_return
.as_ref()
.map(|resource| resource.resource_type_name.clone())
})
.unwrap_or_else(|| output_annotation_default.clone());
let unpacked_input = if let (Some(input), Some(rendered_input)) = (input, &rendered_input) {
if rendered_input.supports_unpacked {
Some(self.build_unpacked_input(match input {
PlannedType::Record(record) => &record.full_name,
_ => panic!("unpacked operation input should be a record"),
})?)
} else {
None
}
} else {
None
};
let output_type_parameters = unpacked_input
.iter()
.flat_map(|unpacked_input| &unpacked_input.functions)
.filter_map(|function| function.result_type_parameter.clone())
.filter(|name| python_annotation_uses_identifier(&overload_output_annotation, name))
.collect::<BTreeSet<_>>();
let output_annotation =
erase_python_type_parameters(&overload_output_annotation, &output_type_parameters);
let output_type_expr = erase_python_type_parameters(
&local_python_model_type_expr(&output_ref),
&output_type_parameters,
);
let descriptor_output_ref = match operation.output_type() {
Some(PlannedType::Record(record)) => python_erased_generic_record_annotation(
&output_ref,
&self
.api_plan
.record_type_parameters(&record.full_name, Language::Python),
),
_ => output_ref.clone(),
};
Ok(RenderedOperation {
name: operation.name.as_str(),
wire_name: operation.wire_name.as_str(),
attr_name: operation
.code_name
.for_language(Language::Python)
.map(str::to_string)
.unwrap_or_else(|| python_ident(&operation.name.to_snake_case())),
experimental: operation.experimental,
deprecated: operation.deprecated,
doc: operation
.doc
.for_language(Language::Python)
.map(str::to_string),
return_doc: operation
.return_doc
.for_language(Language::Python)
.map(str::to_string),
input: rendered_input,
output_ref,
descriptor_output_ref,
output_module_path,
output_annotation,
overload_output_annotation,
output_type_parameters,
output_type_expr,
output_transform_expr: output_transform.and_then(|transform| {
transform
.transform
.for_language(crate::language::Language::Python)
.map(str::to_string)
}),
serialization_context_expr: operation
.serialization_context
.for_language(crate::language::Language::Python)
.map(str::to_string),
output_resource_return,
output_direct_result,
output_none: operation.output_type().is_none(),
unpacked_input,
model_input_type_parameters: model_input_parameters
.into_iter()
.map(|usage| usage.parameter.name)
.collect(),
generic_model_types,
})
}
fn resolve_output_annotation(&mut self, model_type: &PlannedType) -> Result<String> {
match model_type {
PlannedType::External(ExternalTypeSpec::Proto(_))
| PlannedType::External(ExternalTypeSpec::Json(_)) => Ok(self
.external_models
.model_type_annotation(model_type)
.expect("external model annotation should exist")),
PlannedType::Record(record) => {
let conversion = self.resolve_message_value_conversion(model_type)?;
Ok(python_generic_record_annotation(
&conversion.annotation,
&self
.api_plan
.record_type_parameters(&record.full_name, Language::Python),
))
}
_ => panic!("operation output annotation should be model-shaped"),
}
}
fn resolve_message_value_conversion(
&mut self,
model_type: &PlannedType,
) -> Result<WireValueConversion> {
let planned_record = match model_type {
PlannedType::Record(record) => self.api_plan.record(&record.full_name),
_ => None,
};
if let Some(conversion) = self
.external_models
.wire_conversion(model_type, planned_record)
{
if matches!(model_type, PlannedType::Record(_)) {
self.ensure_rendered_model(model_type)?;
}
return Ok(conversion);
}
panic!("message conversion should be model-shaped")
}
fn build_unpacked_input(&self, input_full_name: &str) -> Result<RenderedUnpackedInput> {
let model = self
.models
.get(input_full_name)
.expect("input model should be rendered before building unpacked input");
let planned_model = planned_record(self.api_plan, input_full_name);
let mut parameters = Vec::new();
let mut request_fields = Vec::new();
let mut flattened_messages = Vec::new();
let mut parameter_sources = BTreeMap::<String, String>::new();
for ((_field_name, planned_field), rendered_field) in planned_model
.fields
.iter()
.filter(|(_, field)| field.visibility != RecordFieldVisibility::Omitted)
.zip(model.fields.iter())
{
if planned_field.visibility != RecordFieldVisibility::Public {
continue;
}
if let Some(flattened) = self.build_flattened_message(planned_field, rendered_field) {
request_fields.push(RenderedUnpackedRequestField {
attr_name: flattened.local_name.clone(),
value_expr: flattened.local_name.clone(),
default_kind: PythonFieldDefaultKind::Required,
});
for child in &flattened.fields {
register_unpacked_parameter_name(
&mut parameter_sources,
&child.value_expr,
&format!("{}.{}", flattened.local_name, child.attr_name),
&model.name,
)?;
parameters.push(RenderedUnpackedInputField {
attr_name: child.value_expr.clone(),
annotation: child.annotation.clone(),
default_kind: child.default_kind.clone(),
doc: child.doc.clone(),
});
}
flattened_messages.push(flattened);
continue;
}
register_unpacked_parameter_name(
&mut parameter_sources,
&rendered_field.attr_name,
&rendered_field.attr_name,
&model.name,
)?;
parameters.push(RenderedUnpackedInputField {
attr_name: rendered_field.attr_name.clone(),
annotation: rendered_field.annotation.clone(),
default_kind: rendered_field.default_kind.clone(),
doc: planned_field
.doc
.as_ref()
.and_then(|doc| doc.for_language(Language::Python))
.map(str::to_string),
});
request_fields.push(RenderedUnpackedRequestField {
attr_name: rendered_field.attr_name.clone(),
value_expr: rendered_field.attr_name.clone(),
default_kind: rendered_field.default_kind.clone(),
});
}
Ok(RenderedUnpackedInput {
model_name: model.name.clone(),
parameters,
request_fields,
flattened_messages,
functions: planned_model
.functions()
.map(|(field_name, function)| {
let result_annotation = python_function_result_annotation(&function.result);
let result_type_parameters = function
.result_type_parameter
.iter()
.cloned()
.collect::<BTreeSet<_>>();
let erased_result_annotation =
erase_python_type_parameters(&result_annotation, &result_type_parameters);
let args = python_function_args(&function.args);
let alternate_annotation = function
.alternate_type
.as_ref()
.map(python_authored_type_annotation);
RenderedFunctionField {
callable_field_name: planned_model
.field_name_override(field_name)
.map(python_field_name)
.unwrap_or_else(|| python_field_name(field_name)),
args_field_name: planned_model
.field_name_override(&function.args_field)
.map(python_field_name)
.unwrap_or_else(|| python_field_name(&function.args_field)),
args: args.clone(),
primary: function.primary,
callable_annotation: python_function_annotation_from_args(
&args,
alternate_annotation.as_deref(),
&erased_result_annotation,
),
result_annotation,
erased_result_annotation,
result_type_parameter: function.result_type_parameter.clone(),
alternate_annotation,
}
})
.collect(),
})
}
fn build_flattened_message(
&self,
planned_field: &RecordFieldSpec<PlannedFamily>,
rendered_field: &RenderedField,
) -> Option<RenderedFlattenedMessage> {
let model_type = &planned_field.field_type;
let PlannedType::Record(record) = model_type else {
return None;
};
let full_name = &record.full_name;
let nested_planned_model = self.api_plan.record(full_name)?;
if !nested_planned_model.flatten_in_api {
return None;
}
let nested_rendered_model = self.models.get(full_name)?;
Some(RenderedFlattenedMessage {
local_name: rendered_field.attr_name.clone(),
model_name: nested_rendered_model.name.clone(),
required: planned_field.required,
fields: nested_planned_model
.fields
.iter()
.filter(|(_, field)| field.visibility != RecordFieldVisibility::Omitted)
.zip(nested_rendered_model.fields.iter())
.filter(|((_, field), _)| field.visibility == RecordFieldVisibility::Public)
.map(
|((nested_field_name, nested_planned_field), nested_rendered_field)| {
RenderedFlattenedMessageField {
attr_name: nested_rendered_field.attr_name.clone(),
annotation: nested_planned_model
.field_flattened_annotation(nested_field_name)
.and_then(|annotation| annotation.for_language(Language::Python))
.map(|annotation| {
if nested_planned_field.required {
annotation.to_string()
} else {
format!("{annotation} | None")
}
})
.unwrap_or_else(|| nested_rendered_field.annotation.clone()),
value_expr: nested_rendered_field.attr_name.clone(),
default_kind: if nested_planned_field.required {
PythonFieldDefaultKind::Required
} else {
nested_rendered_field.default_kind.clone()
},
doc: nested_planned_field
.doc
.as_ref()
.and_then(|doc| doc.for_language(Language::Python))
.map(str::to_string),
}
},
)
.collect(),
})
}
fn ensure_rendered_model(&mut self, model_type: &PlannedType) -> Result<()> {
let PlannedType::Record(record) = model_type else {
panic!("rendered model should be a record");
};
let full_name = record.full_name.as_str();
let api_plan = self.api_plan;
let planned_model = planned_record(api_plan, full_name);
if self.models.contains_key(full_name) {
return Ok(());
}
self.models.insert(
full_name.to_string(),
RenderedModel {
full_name: full_name.to_string(),
name: planned_model.name.clone(),
type_parameters: api_plan
.record_type_parameters(full_name, Language::Python)
.into_iter()
.map(|usage| usage.parameter.name)
.collect(),
experimental: planned_model.experimental,
fields: Vec::new(),
},
);
let fields = planned_model
.fields
.iter()
.filter(|(_, field)| field.visibility != RecordFieldVisibility::Omitted)
.map(|(field_name, field)| {
if let RecordFieldVisibility::Sourced { source_expr } = &field.visibility {
self.build_public_sourced_field(planned_model, field_name, field, source_expr)
} else {
self.build_field(planned_model, field_name, field)
}
})
.collect::<Result<Vec<_>>>()?;
self.models
.get_mut(full_name)
.expect("model should be inserted before recursive field resolution")
.fields = fields;
Ok(())
}
fn ensure_rendered_enum(&mut self, enum_spec: &EnumSpec<PlannedFamily>) {
self.enums
.entry(enum_spec.full_name.clone())
.or_insert_with(|| RenderedEnum {
name: enum_spec.name.clone(),
values: enum_spec
.values
.iter()
.map(|value| RenderedEnumValue {
name: value.name.clone(),
number: value.number,
})
.collect(),
});
}
fn ensure_rendered_flags(&mut self, flags_spec: &FlagsSpec<PlannedFamily>) {
self.flags
.entry(flags_spec.full_name.clone())
.or_insert_with(|| RenderedFlags {
name: flags_spec.name.clone(),
flags: flags_spec
.flags
.iter()
.map(|flag| RenderedFlag {
name: flag.name.clone(),
bit: flag.bit,
})
.collect(),
});
}
fn ensure_rendered_variant(&mut self, variant_spec: &VariantSpec<PlannedFamily>) -> Result<()> {
if self.variants.contains_key(&variant_spec.full_name) {
return Ok(());
}
let cases = variant_spec
.cases
.iter()
.map(|case| -> Result<_> {
Ok(RenderedVariantCase {
name: case.name.clone(),
type_parameters: case
.payload
.as_ref()
.map(|payload| {
self.api_plan
.type_parameters(payload, Language::Python)
.into_iter()
.map(|usage| usage.parameter.name)
.collect()
})
.unwrap_or_default(),
payload_annotation: case
.payload
.as_ref()
.map(|payload| {
self.resolve_planned_value_type(payload)
.map(|resolved| resolved.annotation)
})
.transpose()?,
})
})
.collect::<Result<Vec<_>>>()?;
let rendered_variant = RenderedVariant {
full_name: variant_spec.full_name.clone(),
name: variant_spec.name.clone(),
type_parameters: self
.api_plan
.variant_type_parameters(&variant_spec.full_name, Language::Python)
.into_iter()
.map(|usage| usage.parameter.name)
.collect(),
cases,
};
self.variants
.insert(variant_spec.full_name.clone(), rendered_variant);
Ok(())
}
fn build_field(
&mut self,
record: &RecordSpec<PlannedFamily>,
field_name: &str,
field: &RecordFieldSpec<PlannedFamily>,
) -> Result<RenderedField> {
let attr_name = python_field_name(&field.name);
if let PlannedType::Map(key, value) = &field.field_type {
let key_type = self.resolve_planned_value_type(key)?;
let value_type = self.resolve_planned_value_type(value)?;
let mut imports = key_type.imports.clone();
imports.extend(&value_type.imports);
return Ok(RenderedField {
attr_name: attr_name.clone(),
annotation: python_field_annotation(
record,
field_name,
field,
format!("dict[{}, {}]", key_type.annotation, value_type.annotation),
false,
),
default_kind: PythonFieldDefaultKind::EmptyDict,
default_expr: Some("dataclasses.field(default_factory=dict)".to_string()),
wire_value_type: value_type,
imports,
});
}
let (resolved_type, repeated) = match &field.field_type {
PlannedType::List(value) => (self.resolve_planned_value_type(value)?, true),
PlannedType::Map(_, _) => unreachable!("handled above"),
value => (self.resolve_planned_value_type(value)?, false),
};
if repeated {
return Ok(RenderedField {
attr_name: attr_name.clone(),
annotation: python_field_annotation(
record,
field_name,
field,
format!("list[{}]", resolved_type.annotation),
false,
),
default_kind: PythonFieldDefaultKind::EmptyList,
default_expr: Some("dataclasses.field(default_factory=list)".to_string()),
wire_value_type: resolved_type.clone(),
imports: resolved_type.imports,
});
}
let imports = resolved_type.imports.clone();
if let Some(default_value) = &field.default_value {
let default_expr = enum_default_expr(&resolved_type, &default_value.enum_case);
return Ok(RenderedField {
attr_name: attr_name.clone(),
annotation: python_field_annotation(
record,
field_name,
field,
resolved_type.annotation.clone(),
false,
),
default_kind: PythonFieldDefaultKind::Expression(default_expr.clone()),
default_expr: Some(default_expr.clone()),
wire_value_type: resolved_type.clone(),
imports,
});
}
if field.required {
return Ok(RenderedField {
attr_name: attr_name.clone(),
annotation: python_field_annotation(
record,
field_name,
field,
resolved_type.annotation.clone(),
false,
),
default_kind: PythonFieldDefaultKind::Required,
default_expr: None,
wire_value_type: resolved_type.clone(),
imports,
});
}
Ok(RenderedField {
attr_name: attr_name.clone(),
annotation: python_field_annotation(
record,
field_name,
field,
resolved_type.annotation.clone(),
true,
),
default_kind: PythonFieldDefaultKind::None,
default_expr: Some("None".to_string()),
wire_value_type: resolved_type.clone(),
imports,
})
}
fn build_public_sourced_field(
&mut self,
record: &RecordSpec<PlannedFamily>,
field_name: &str,
field: &RecordFieldSpec<PlannedFamily>,
source_expr: &str,
) -> Result<RenderedField> {
let mut rendered = self.build_field(record, field_name, field)?;
rendered.default_kind = PythonFieldDefaultKind::Expression(source_expr.to_string());
rendered.default_expr = Some(python_dataclass_source_default_expr(source_expr));
Ok(rendered)
}
fn resolve_planned_value_type(
&mut self,
value_type: &PlannedType,
) -> Result<ResolvedFieldType> {
self.ensure_rendered_value_type(value_type)?;
resolve_python_value_type(self.api_plan, &self.external_models, value_type)
}
fn ensure_rendered_value_type(&mut self, value_type: &PlannedType) -> Result<()> {
match value_type {
PlannedType::Enum(enum_type) => {
if let Some(enum_spec) = self.api_plan.enum_decl(&enum_type.full_name) {
self.ensure_rendered_enum(enum_spec);
}
}
PlannedType::Flags(flags_type) => {
if let Some(flags_spec) = self.api_plan.flags_decl(&flags_type.full_name) {
self.ensure_rendered_flags(flags_spec);
}
}
PlannedType::Variant(variant_type) => {
if let Some(variant_spec) = self.api_plan.variant(&variant_type.full_name) {
self.ensure_rendered_variant(variant_spec)?;
}
}
PlannedType::Record(_) => self.ensure_rendered_model(value_type)?,
PlannedType::Option(inner) | PlannedType::List(inner) => {
self.ensure_rendered_value_type(inner)?;
}
PlannedType::Map(key, value) => {
self.ensure_rendered_value_type(key)?;
self.ensure_rendered_value_type(value)?;
}
PlannedType::Tuple(items) => {
for item in items {
self.ensure_rendered_value_type(item)?;
}
}
PlannedType::Result { ok, err } => {
if let Some(ok) = ok {
self.ensure_rendered_value_type(ok)?;
}
if let Some(err) = err {
self.ensure_rendered_value_type(err)?;
}
}
PlannedType::External(ExternalTypeSpec::Alias(AliasTypeSpec { target, .. })) => {
self.ensure_rendered_value_type(target)?;
}
PlannedType::TypeParameter(_)
| PlannedType::Float
| PlannedType::Int(_)
| PlannedType::Bool
| PlannedType::String
| PlannedType::Bytes
| PlannedType::External(_)
| PlannedType::Resource(_) => {}
}
Ok(())
}
}
fn resolve_python_value_type(
api_plan: &PlannedSpec,
external_models: &PythonExternalModels,
value_type: &PlannedType,
) -> Result<ResolvedFieldType> {
let resolved = match value_type {
PlannedType::TypeParameter(parameter) => ResolvedFieldType {
annotation: parameter.name.clone(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Float => ResolvedFieldType {
annotation: "float".to_string(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Int(_) => ResolvedFieldType {
annotation: "int".to_string(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Bool => ResolvedFieldType {
annotation: "bool".to_string(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::String => ResolvedFieldType {
annotation: "str".to_string(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Bytes => ResolvedFieldType {
annotation: "bytes".to_string(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Enum(enum_type) => ResolvedFieldType {
annotation: enum_type.name.clone(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Enum,
wire_conversion: None,
},
value_type @ PlannedType::External(ExternalTypeSpec::Proto(PlannedProtoType::Enum(_))) => {
external_models
.resolved_external_field_type(value_type)
.expect("proto enum field type should exist")
}
PlannedType::Flags(flags_type) => ResolvedFieldType {
annotation: flags_type.name.clone(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Variant(variant_type) => ResolvedFieldType {
annotation: python_generic_record_annotation(
&variant_type.name,
&api_plan.variant_type_parameters(&variant_type.full_name, Language::Python),
),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
message_type @ (PlannedType::External(ExternalTypeSpec::Proto(
PlannedProtoType::Message(_),
))
| PlannedType::Record(_)) => {
let planned_record = match message_type {
PlannedType::Record(record) => api_plan.record(&record.full_name),
_ => None,
};
let conversion = external_models
.wire_conversion(message_type, planned_record)
.expect("message conversion should be model-shaped");
let annotation = if let PlannedType::Record(record) = message_type {
python_generic_record_annotation(
&conversion.annotation,
&api_plan.record_type_parameters(&record.full_name, Language::Python),
)
} else {
conversion.annotation.clone()
};
ResolvedFieldType {
annotation,
imports: conversion.imports.clone(),
kind: ResolvedFieldKind::Message,
wire_conversion: Some(conversion),
}
}
value_type @ PlannedType::External(ExternalTypeSpec::Json(_)) => external_models
.resolved_external_field_type(value_type)
.expect("json model field type should exist"),
PlannedType::Resource(resource) => ResolvedFieldType {
annotation: resource.type_name.clone(),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
},
PlannedType::Option(inner) | PlannedType::List(inner) => {
let inner = resolve_python_value_type(api_plan, external_models, inner)?;
ResolvedFieldType {
annotation: format!("list[{}]", inner.annotation),
imports: inner.imports,
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
}
}
PlannedType::Map(key, value) => {
let key = resolve_python_value_type(api_plan, external_models, key)?;
let value = resolve_python_value_type(api_plan, external_models, value)?;
let mut imports = key.imports;
imports.extend(&value.imports);
ResolvedFieldType {
annotation: format!("dict[{}, {}]", key.annotation, value.annotation),
imports,
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
}
}
PlannedType::Tuple(items) => {
let items = items
.iter()
.map(|item| {
resolve_python_value_type(api_plan, external_models, item)
.map(|resolved| resolved.annotation)
})
.collect::<Result<Vec<_>>>()?;
ResolvedFieldType {
annotation: format!("tuple[{}]", items.join(", ")),
imports: PythonImports::default(),
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
}
}
PlannedType::Result { ok, err } => {
let ok = ok
.as_ref()
.map(|ok| resolve_python_value_type(api_plan, external_models, ok))
.transpose()?;
let err = err
.as_ref()
.map(|err| resolve_python_value_type(api_plan, external_models, err))
.transpose()?;
let mut imports = PythonImports::default();
if let Some(ok) = &ok {
imports.extend(&ok.imports);
}
if let Some(err) = &err {
imports.extend(&err.imports);
}
ResolvedFieldType {
annotation: python_result_annotation(
ok.as_ref().map(|ok| ok.annotation.clone()),
err.as_ref().map(|err| err.annotation.clone()),
),
imports,
kind: ResolvedFieldKind::Scalar,
wire_conversion: None,
}
}
PlannedType::External(ExternalTypeSpec::Alias(AliasTypeSpec {
type_name, target, ..
})) => {
let mut resolved = resolve_python_value_type(api_plan, external_models, target)?;
if let Some(annotation) = type_name.for_language(Language::Python) {
resolved.annotation = annotation.to_string();
}
resolved
}
};
Ok(resolved)
}
fn collect_python_language_imports(api_plan: &PlannedSpec) -> Vec<LanguageImportSpec> {
let mut imports = BTreeSet::new();
for variant in api_plan.variants().map(|(_, variant)| variant) {
for case in &variant.cases {
if let Some(payload) = &case.payload {
collect_value_type_imports(payload, &mut imports);
}
}
}
for record in api_plan.records().map(|(_, record)| record) {
if let Some(proto) = &record.data.proto {
collect_proto_type_imports(proto, &mut imports);
}
for field in record.fields.values() {
if let Some(annotation) = &field.annotation {
collect_python_import(annotation, &mut imports);
}
if let Some(annotation) = &field.flattened_annotation {
collect_python_import(annotation, &mut imports);
}
}
for (_, field) in record
.fields
.iter()
.filter(|(_, field)| field.visibility != RecordFieldVisibility::Omitted)
{
if let Some(function) = &field.function {
collect_function_imports(function, &mut imports);
}
collect_value_type_imports(&field.field_type, &mut imports);
}
}
for service in &api_plan.services {
for operation in &service.operations {
if let Some(input) = operation_input_model(operation) {
collect_model_type_imports(input, &mut imports);
}
match operation.output_type() {
Some(
model_type @ (PlannedType::External(ExternalTypeSpec::Proto(
PlannedProtoType::Message(_),
))
| PlannedType::External(ExternalTypeSpec::Json(_))
| PlannedType::Record(_)),
) => collect_model_type_imports(model_type, &mut imports),
Some(PlannedType::Resource(resource)) => {
if let Some(model_type) = &resource.wire_type {
collect_model_type_imports(
&PlannedType::External(model_type.clone()),
&mut imports,
);
}
}
_ => {}
}
if let Some(transform) = &operation.output_transform {
collect_python_import(&transform.type_name, &mut imports);
}
}
for resource in service.resources.iter().map(|resource| &resource.data) {
for field in &resource.fields {
collect_value_type_imports(&field.kind, &mut imports);
if let Some(function) = &field.function {
collect_function_imports(function, &mut imports);
}
}
for method in &resource.methods {
for param in &method.params {
collect_value_type_imports(¶m.kind, &mut imports);
if let Some(function) = ¶m.function {
collect_function_imports(function, &mut imports);
}
}
if let Some(result) = &method.result {
if let PlannedResourceMethodResultKind::Value(kind) = &result.kind {
collect_value_type_imports(kind, &mut imports);
}
}
}
}
}
imports.into_iter().collect()
}
fn render_record_models(
models: &[&RenderedModel],
api_plan: &PlannedSpec,
external_models: &PythonExternalModels,
) -> Result<RenderedModelFragments> {
let mut body = String::new();
let type_parameters = models
.iter()
.flat_map(|model| model.type_parameters.iter())
.cloned()
.collect::<BTreeSet<_>>();
for parameter in &type_parameters {
body.push_str(parameter);
body.push_str(" = typing.TypeVar(");
body.push_str(&python_string_literal(parameter));
body.push_str(")\n");
}
if !type_parameters.is_empty() && !models.is_empty() {
body.push_str("\n\n");
}
let mut wire_blocks = BTreeMap::new();
for (index, model) in models.iter().enumerate() {
let planned_model = api_plan
.record(&model.full_name)
.unwrap_or_else(|| panic!("planned model should exist for {}", model.full_name));
let wire_block =
external_models.render_record_wire_block(api_plan, model, planned_model)?;
render_record_model(&mut body, model, wire_block.as_ref());
if let Some(wire_block) = &wire_block {
for line in &wire_block.post_class_lines {
body.push_str(line);
body.push('\n');
}
}
if let Some(wire_block) = wire_block {
wire_blocks.insert(model.full_name.clone(), wire_block);
}
if index + 1 != models.len() {
body.push_str("\n\n");
}
}
let mut module_imports = BTreeSet::new();
for model in models {
if let Some(wire_block) = wire_blocks.get(&model.full_name) {
module_imports.extend(wire_block.imports.module_imports.iter().cloned());
}
for field in &model.fields {
module_imports.extend(field.imports.module_imports.iter().cloned());
}
}
Ok(RenderedModelFragments {
body,
post_model_statements: String::new(),
module_imports,
relative_imports: BTreeMap::new(),
root_package_imports: RootPackageImports::new(),
exported_names: models.iter().map(|model| model.name.clone()).collect(),
module_exported_names: BTreeSet::new(),
generated_names: Vec::new(),
export_sort_keys: BTreeMap::new(),
declared_type_parameters: type_parameters,
allows_private_wire_access: false,
})
}
fn render_record_model(
output: &mut String,
model: &RenderedModel,
wire_block: Option<&RenderedRecordWireBlock>,
) {
if model_needs_keyword_only_dataclass(model) {
output.push_str("@dataclasses.dataclass(slots=True, kw_only=True)\n");
} else {
output.push_str("@dataclasses.dataclass(slots=True)\n");
}
output.push_str("class ");
output.push_str(&model.name);
if !model.type_parameters.is_empty() {
output.push_str("(typing.Generic[");
output.push_str(&model.type_parameters.join(", "));
output.push_str("])");
}
output.push_str(":\n");
render_python_docstring(output, " ", None, &[], None, model.experimental);
let wire_lines = wire_block
.map(|block| block.class_body_lines.as_slice())
.unwrap_or(&[]);
if model.fields.is_empty() {
if wire_lines.is_empty() {
output.push_str(" pass\n");
} else {
render_record_wire_block_lines(output, wire_lines);
}
return;
}
for field in &model.fields {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(&python_parameter_annotation(
&field.annotation,
&field.default_kind,
));
if let Some(default_expr) = &field.default_expr {
render_python_default_expr(output, default_expr, " ");
}
output.push('\n');
}
render_record_wire_block_lines(output, wire_lines);
}
fn render_record_wire_block_lines(output: &mut String, lines: &[String]) {
for line in lines {
output.push_str(line);
output.push('\n');
}
}
fn model_needs_keyword_only_dataclass(model: &RenderedModel) -> bool {
let mut saw_defaulted_field = false;
for field in &model.fields {
if field.default_expr.is_some() {
saw_defaulted_field = true;
} else if saw_defaulted_field {
return true;
}
}
false
}
fn collect_proto_type_imports(
proto: &PlannedProtoTypeInfo,
imports: &mut BTreeSet<LanguageImportSpec>,
) {
collect_python_import(&proto.type_name, imports);
}
fn collect_model_type_imports(
model_type: &PlannedType,
imports: &mut BTreeSet<LanguageImportSpec>,
) {
let PlannedType::External(ExternalTypeSpec::Proto(PlannedProtoType::Message(proto))) =
model_type
else {
return;
};
collect_proto_type_imports(&proto.proto, imports);
if let Some(replacement) = &proto.replacement {
collect_type_replacement_imports(replacement, imports);
}
if let Some(authored_type) = &proto.authored_type {
collect_authored_type_imports(authored_type, imports);
}
}
fn collect_type_replacement_imports(
replacement: &TypeReplacementSpec,
imports: &mut BTreeSet<LanguageImportSpec>,
) {
collect_python_import(&replacement.type_name, imports);
collect_python_import(&replacement.from_proto, imports);
collect_python_import(&replacement.to_proto, imports);
}
fn collect_value_type_imports(value: &PlannedType, imports: &mut BTreeSet<LanguageImportSpec>) {
match value {
PlannedType::Enum(_) | PlannedType::TypeParameter(_) => {}
PlannedType::Flags(_) | PlannedType::Variant(_) => {}
PlannedType::External(ExternalTypeSpec::Proto(PlannedProtoType::Enum(enumeration))) => {
collect_proto_type_imports(&enumeration.proto, imports);
if let Some(replacement) = &enumeration.replacement {
collect_type_replacement_imports(replacement, imports);
}
}
PlannedType::External(ExternalTypeSpec::Proto(PlannedProtoType::Message(_)))
| PlannedType::External(ExternalTypeSpec::Json(_))
| PlannedType::Record(_) => {
collect_model_type_imports(value, imports);
}
PlannedType::Resource(resource) => {
if let Some(wire_type) = &resource.wire_type {
collect_model_type_imports(&PlannedType::External(wire_type.clone()), imports);
}
}
PlannedType::Option(inner) | PlannedType::List(inner) => {
collect_value_type_imports(inner, imports)
}
PlannedType::Map(key, value) => {
collect_value_type_imports(key, imports);
collect_value_type_imports(value, imports);
}
PlannedType::Tuple(items) => {
for item in items {
collect_value_type_imports(item, imports);
}
}
PlannedType::Result { ok, err } => {
if let Some(ok) = ok {
collect_value_type_imports(ok, imports);
}
if let Some(err) = err {
collect_value_type_imports(err, imports);
}
}
PlannedType::External(ExternalTypeSpec::Alias(AliasTypeSpec {
type_name, target, ..
})) => {
collect_python_import(type_name, imports);
collect_value_type_imports(target, imports);
}
PlannedType::Bool
| PlannedType::Int(_)
| PlannedType::Float
| PlannedType::String
| PlannedType::Bytes => {}
}
}
fn collect_function_imports(
function: &FunctionFieldSpec<PlannedFamily>,
imports: &mut BTreeSet<LanguageImportSpec>,
) {
if let FunctionResultSpec::Annotation(result) = &function.result {
collect_python_import(result, imports);
}
collect_function_args_imports(&function.args, imports);
if let Some(alternate_type) = &function.alternate_type {
collect_authored_type_imports(alternate_type, imports);
}
if let Some(descriptor) = &function.type_descriptor {
collect_python_import(&descriptor.value_type, imports);
collect_python_import(&descriptor.args_type, imports);
}
}
fn collect_function_args_imports(
args: &FunctionArgsSpec<PlannedFamily>,
imports: &mut BTreeSet<LanguageImportSpec>,
) {
let fields = match args {
FunctionArgsSpec::Varargs { prefix, .. } => prefix,
FunctionArgsSpec::Fixed(fields) => fields,
};
for field in fields {
collect_authored_type_imports(&field.field_type, imports);
}
}
fn collect_authored_type_imports(
field_type: &PlannedType,
imports: &mut BTreeSet<LanguageImportSpec>,
) {
match field_type {
TypeSpec::Option(inner) | TypeSpec::List(inner) => {
collect_authored_type_imports(inner, imports);
}
TypeSpec::Tuple(items) => {
for item in items {
collect_authored_type_imports(item, imports);
}
}
TypeSpec::Map(key, value) => {
collect_authored_type_imports(key, imports);
collect_authored_type_imports(value, imports);
}
TypeSpec::Result { ok, err } => {
if let Some(ok) = ok {
collect_authored_type_imports(ok, imports);
}
if let Some(err) = err {
collect_authored_type_imports(err, imports);
}
}
TypeSpec::External(ExternalTypeSpec::Alias(AliasTypeSpec {
target, type_name, ..
})) => {
collect_python_import(type_name, imports);
collect_authored_type_imports(target, imports);
}
TypeSpec::TypeParameter(_)
| TypeSpec::Bool
| TypeSpec::Int(_)
| TypeSpec::Float
| TypeSpec::String
| TypeSpec::Bytes
| TypeSpec::External(ExternalTypeSpec::Proto(_))
| TypeSpec::External(ExternalTypeSpec::Json(_))
| TypeSpec::Record(_)
| TypeSpec::Enum(_)
| TypeSpec::Flags(_)
| TypeSpec::Variant(_)
| TypeSpec::Resource(_) => {}
}
}
fn collect_python_import(spec: &LanguageStringSpec, imports: &mut BTreeSet<LanguageImportSpec>) {
if let Some(module) = spec.import_for_language(Language::Python) {
imports.insert(LanguageImportSpec {
language: Language::Python,
reference: module.to_string(),
module: module.to_string(),
name: None,
type_only: false,
import_style: LanguageImportStyle::Module,
});
return;
}
let Some(expression) = spec.for_language(Language::Python) else {
return;
};
for module_path in python_qualified_module_paths(expression) {
imports.insert(LanguageImportSpec {
language: Language::Python,
reference: module_path.clone(),
module: module_path,
name: None,
type_only: false,
import_style: LanguageImportStyle::Module,
});
}
}
fn python_qualified_module_paths(expression: &str) -> BTreeSet<String> {
let chars = expression.char_indices().collect::<Vec<_>>();
let mut module_paths = BTreeSet::new();
let mut index = 0;
while index < chars.len() {
let (start_byte, ch) = chars[index];
if !is_python_identifier_start(ch) {
index += 1;
continue;
}
let before = expression[..start_byte].chars().next_back();
if before.is_some_and(|before| is_python_identifier_char(before) || before == '.') {
index += 1;
continue;
}
let mut end = index + 1;
while end < chars.len() {
let ch = chars[end].1;
if is_python_identifier_char(ch) || ch == '.' {
end += 1;
} else {
break;
}
}
let end_byte = chars
.get(end)
.map(|(byte, _)| *byte)
.unwrap_or(expression.len());
let qualified_name = expression[start_byte..end_byte].trim_end_matches('.');
if let Some(module_path) = python_module_path_for_qualified_name(qualified_name) {
module_paths.insert(module_path.to_string());
}
index = end;
}
module_paths
}
fn python_module_path_for_qualified_name(qualified_name: &str) -> Option<&str> {
let (module_path, _) = qualified_name.rsplit_once('.')?;
if is_builtin_python_import(module_path) {
return None;
}
Some(module_path)
}
fn is_builtin_python_import(module_path: &str) -> bool {
matches!(
module_path,
"collections.abc" | "typing" | "typing_extensions"
)
}
fn is_python_identifier_start(ch: char) -> bool {
ch.is_ascii_alphabetic() || ch == '_'
}
fn reject_support_namespaces(
language: Language,
support_fragments: &[SupportFragmentSpec],
) -> Result<()> {
if let Some(namespace) = support_fragments
.iter()
.find_map(|fragment| fragment.namespace.as_deref())
{
return Err(Error::UnsupportedSupportNamespace {
language,
namespace: namespace.to_string(),
});
}
Ok(())
}
pub(crate) fn python_field_name(name: &str) -> String {
python_ident(&name.to_snake_case())
}
#[derive(Debug)]
struct RenderedService<'a> {
name: &'a str,
wire_name: &'a str,
doc: Option<String>,
endpoint: Option<String>,
experimental: bool,
deprecated: bool,
delay_load_temporalio_workflow: bool,
operations: Vec<RenderedOperation<'a>>,
resources: Vec<PlannedResource>,
}
#[derive(Debug)]
struct RenderedOperation<'a> {
name: &'a str,
wire_name: &'a str,
attr_name: String,
experimental: bool,
deprecated: bool,
doc: Option<String>,
return_doc: Option<String>,
input: Option<RenderedOperationInput>,
output_ref: String,
descriptor_output_ref: String,
output_module_path: Option<String>,
output_annotation: String,
overload_output_annotation: String,
output_type_parameters: BTreeSet<String>,
output_type_expr: String,
output_transform_expr: Option<String>,
serialization_context_expr: Option<String>,
output_resource_return: Option<PlannedOperationResourceReturn>,
output_direct_result: bool,
output_none: bool,
unpacked_input: Option<RenderedUnpackedInput>,
model_input_type_parameters: Vec<String>,
generic_model_types: bool,
}
#[derive(Debug)]
struct RenderedOperationInput {
type_ref: String,
descriptor_type_ref: String,
module_path: Option<String>,
annotation: String,
supports_unpacked: bool,
}
#[derive(Debug, Clone, Copy)]
struct ResourceBoundOperation<'a> {
operation: &'a RenderedOperation<'a>,
}
#[derive(Debug)]
struct RenderedEnum {
name: String,
values: Vec<RenderedEnumValue>,
}
#[derive(Debug)]
struct RenderedEnumValue {
name: String,
number: i32,
}
#[derive(Debug)]
struct RenderedFlags {
name: String,
flags: Vec<RenderedFlag>,
}
#[derive(Debug)]
struct RenderedFlag {
name: String,
bit: usize,
}
#[derive(Debug)]
pub(in crate::generator) struct RenderedVariant {
pub(in crate::generator) full_name: String,
pub(in crate::generator) name: String,
pub(in crate::generator) type_parameters: Vec<String>,
pub(in crate::generator) cases: Vec<RenderedVariantCase>,
}
#[derive(Debug)]
pub(in crate::generator) struct RenderedVariantCase {
pub(in crate::generator) name: String,
pub(in crate::generator) type_parameters: Vec<String>,
pub(in crate::generator) payload_annotation: Option<String>,
}
#[derive(Debug)]
pub(in crate::generator) struct PythonGeneratedName {
pub(in crate::generator) name: String,
pub(in crate::generator) generated_by: String,
}
fn sort_python_model_names(
names: &mut [String],
variants: &IndexMap<String, RenderedVariant>,
model_fragments: &RenderedModelFragments,
) {
names.sort_by_key(|name| {
model_fragments
.export_sort_keys
.get(name)
.cloned()
.or_else(|| {
variants
.values()
.find(|variant| name == &variant.name)
.map(|variant| (variant.name.clone(), 0))
})
.unwrap_or_else(|| (name.clone(), 0))
});
}
#[derive(Debug)]
pub(in crate::generator) struct RenderedModel {
pub(in crate::generator) full_name: String,
pub(in crate::generator) name: String,
pub(in crate::generator) type_parameters: Vec<String>,
pub(in crate::generator) experimental: bool,
pub(in crate::generator) fields: Vec<RenderedField>,
}
#[derive(Debug)]
pub(in crate::generator) struct RenderedField {
pub(in crate::generator) attr_name: String,
pub(in crate::generator) annotation: String,
pub(in crate::generator) default_kind: PythonFieldDefaultKind,
pub(in crate::generator) default_expr: Option<String>,
pub(in crate::generator) wire_value_type: ResolvedFieldType,
pub(in crate::generator) imports: PythonImports,
}
#[derive(Debug, Default)]
pub(in crate::generator) struct RenderedRecordWireBlock {
pub(in crate::generator) imports: PythonImports,
pub(in crate::generator) class_body_lines: Vec<String>,
pub(in crate::generator) post_class_lines: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(in crate::generator) enum PythonFieldDefaultKind {
Required,
None,
EmptyList,
EmptyDict,
Expression(String),
}
#[derive(Debug)]
struct RenderedUnpackedInput {
model_name: String,
parameters: Vec<RenderedUnpackedInputField>,
request_fields: Vec<RenderedUnpackedRequestField>,
flattened_messages: Vec<RenderedFlattenedMessage>,
functions: Vec<RenderedFunctionField>,
}
#[derive(Debug)]
struct RenderedUnpackedInputField {
attr_name: String,
annotation: String,
default_kind: PythonFieldDefaultKind,
doc: Option<String>,
}
#[derive(Debug)]
struct RenderedUnpackedRequestField {
attr_name: String,
value_expr: String,
default_kind: PythonFieldDefaultKind,
}
#[derive(Debug)]
struct RenderedFlattenedMessage {
local_name: String,
model_name: String,
required: bool,
fields: Vec<RenderedFlattenedMessageField>,
}
#[derive(Debug)]
struct RenderedFlattenedMessageField {
attr_name: String,
annotation: String,
value_expr: String,
default_kind: PythonFieldDefaultKind,
doc: Option<String>,
}
#[derive(Debug, Clone)]
struct RenderedFunctionField {
callable_field_name: String,
args_field_name: String,
args: RenderedFunctionArgs,
primary: bool,
callable_annotation: String,
result_annotation: String,
erased_result_annotation: String,
result_type_parameter: Option<String>,
alternate_annotation: Option<String>,
}
#[derive(Debug, Clone)]
enum RenderedFunctionArgs {
Varargs {
prefix: Vec<RenderedFunctionArg>,
},
Typed {
parameters: Vec<RenderedFunctionArg>,
},
}
#[derive(Debug, Clone)]
struct RenderedFunctionArg {
name: String,
annotation: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(in crate::generator) enum ResolvedFieldKind {
Scalar,
Message,
Enum,
}
#[derive(Debug, Clone)]
pub(in crate::generator) struct ResolvedFieldType {
pub(in crate::generator) annotation: String,
pub(in crate::generator) imports: PythonImports,
pub(in crate::generator) kind: ResolvedFieldKind,
pub(in crate::generator) wire_conversion: Option<WireValueConversion>,
}
#[derive(Debug, Clone)]
pub(in crate::generator) struct WireValueConversion {
pub(in crate::generator) annotation: String,
pub(in crate::generator) from_wire: String,
pub(in crate::generator) to_wire: String,
pub(in crate::generator) imports: PythonImports,
pub(in crate::generator) supports_unpacked_input: bool,
}
impl WireValueConversion {
pub(in crate::generator) fn from_wire_expr(&self, wire_expr: &str) -> String {
self.from_wire_expr_with_type_hint(wire_expr, &self.annotation)
}
pub(in crate::generator) fn from_wire_expr_with_type_hint(
&self,
wire_expr: &str,
type_hint_expr: &str,
) -> String {
self.from_wire
.replace("{wire}", wire_expr)
.replace("{type_hint}", type_hint_expr)
}
pub(in crate::generator) fn to_wire_expr(&self, value_expr: &str) -> String {
self.to_wire.replace("{value}", value_expr)
}
fn supports_unpacked_input(&self) -> bool {
self.supports_unpacked_input
}
}
#[derive(Debug, Clone, Default)]
pub(in crate::generator) struct PythonImports {
pub(in crate::generator) module_imports: BTreeSet<String>,
}
impl PythonImports {
pub(in crate::generator) fn extend(&mut self, other: &Self) {
self.module_imports
.extend(other.module_imports.iter().cloned());
}
}
#[derive(Debug, Default)]
pub(in crate::generator) struct RenderedModelFragments {
pub(in crate::generator) body: String,
pub(in crate::generator) post_model_statements: String,
pub(in crate::generator) module_imports: BTreeSet<String>,
pub(in crate::generator) relative_imports: BTreeMap<String, BTreeSet<String>>,
/// Backend-owned public names imported only by the generated package tree's root module.
pub(in crate::generator) root_package_imports: RootPackageImports,
pub(in crate::generator) exported_names: BTreeSet<String>,
pub(in crate::generator) module_exported_names: BTreeSet<String>,
pub(in crate::generator) generated_names: Vec<PythonGeneratedName>,
pub(in crate::generator) export_sort_keys: BTreeMap<String, (String, usize)>,
pub(in crate::generator) declared_type_parameters: BTreeSet<String>,
pub(in crate::generator) allows_private_wire_access: bool,
}
impl RenderedModelFragments {
pub(in crate::generator) fn extend(&mut self, other: Self) {
if !other.body.is_empty() {
if !self.body.is_empty() {
self.body.push_str("\n\n");
}
self.body.push_str(&other.body);
}
if !other.post_model_statements.is_empty() {
if !self.post_model_statements.is_empty() {
self.post_model_statements.push_str("\n\n");
}
self.post_model_statements
.push_str(&other.post_model_statements);
}
self.module_imports.extend(other.module_imports);
self.allows_private_wire_access |= other.allows_private_wire_access;
for (module, names) in other.relative_imports {
self.relative_imports
.entry(module)
.or_default()
.extend(names);
}
extend_root_package_imports(&mut self.root_package_imports, other.root_package_imports);
self.exported_names.extend(other.exported_names);
self.module_exported_names
.extend(other.module_exported_names);
self.generated_names.extend(other.generated_names);
self.export_sort_keys.extend(other.export_sort_keys);
self.declared_type_parameters
.extend(other.declared_type_parameters);
}
}
pub(crate) fn render_tree_support_files(
branch: &ApiSpecBranch<PlannedFamily>,
) -> BTreeMap<PathBuf, String> {
if !branch_has_json_models(branch) {
return BTreeMap::new();
}
BTreeMap::from([(
PathBuf::from("_definitions.py"),
python_json::render_support_file(),
)])
}
#[derive(Debug, Default)]
pub(crate) struct PythonModelHoists {
hoisted: BTreeMap<ModulePath, BTreeSet<String>>,
runtime_imports: BTreeMap<ModulePath, BTreeSet<String>>,
files: BTreeMap<PathBuf, String>,
source_paths: BTreeSet<PathBuf>,
exported_names: BTreeSet<String>,
}
impl PythonModelHoists {
pub(in crate::generator) fn add_module_hoists(
&mut self,
module_path: ModulePath,
names: BTreeSet<String>,
) {
if names.is_empty() {
return;
}
self.hoisted.entry(module_path).or_default().extend(names);
}
pub(in crate::generator) fn add_file(&mut self, path: PathBuf, contents: String) {
self.files.insert(path, contents);
}
pub(in crate::generator) fn add_source_paths(&mut self, paths: BTreeSet<PathBuf>) {
self.source_paths.extend(paths);
}
pub(in crate::generator) fn add_runtime_imports(
&mut self,
module_path: ModulePath,
names: BTreeSet<String>,
) {
if names.is_empty() {
return;
}
self.runtime_imports
.entry(module_path)
.or_default()
.extend(names);
}
pub(in crate::generator) fn add_exported_names(&mut self, names: BTreeSet<String>) {
self.exported_names.extend(names);
}
pub(crate) fn is_empty(&self) -> bool {
self.hoisted.is_empty() && self.files.is_empty()
}
pub(in crate::generator) fn is_hoisted(&self, module_path: &ModulePath, name: &str) -> bool {
self.hoisted
.get(module_path)
.is_some_and(|names| names.contains(name))
}
fn runtime_imports_for(&self, module_path: &ModulePath) -> impl Iterator<Item = &String> {
self.runtime_imports.get(module_path).into_iter().flatten()
}
pub(crate) fn files(&self) -> &BTreeMap<PathBuf, String> {
&self.files
}
pub(crate) fn source_paths(&self) -> &BTreeSet<PathBuf> {
&self.source_paths
}
pub(crate) fn exported_names(&self) -> &BTreeSet<String> {
&self.exported_names
}
}
pub(crate) fn tree_model_hoists(
branch: &ApiSpecBranch<PlannedFamily>,
) -> Result<PythonModelHoists> {
python_json::tree_model_hoists(branch)
}
fn branch_has_json_models(branch: &ApiSpecBranch<PlannedFamily>) -> bool {
branch.children.values().any(|node| match node {
ApiSpecNode::Leaf(leaf) => leaf
.spec
.external_types()
.map(|(_, binding)| binding)
.any(|binding| binding.json_model().is_some()),
ApiSpecNode::Branch(branch) => branch_has_json_models(branch),
})
}
fn tree_json_source_paths(branch: &ApiSpecBranch<PlannedFamily>) -> BTreeSet<PathBuf> {
let mut paths = BTreeSet::new();
for node in branch.children.values() {
match node {
ApiSpecNode::Leaf(leaf) => {
if leaf
.spec
.external_types()
.map(|(_, binding)| binding)
.any(|binding| binding.json_model().is_some())
{
paths.insert(leaf.source_path.clone());
}
}
ApiSpecNode::Branch(branch) => paths.extend(tree_json_source_paths(branch)),
}
}
paths
}
fn source_paths_label(paths: &BTreeSet<PathBuf>) -> String {
paths
.iter()
.map(|path| path.display().to_string())
.collect::<Vec<_>>()
.join(", ")
}
fn python_function_result_annotation(result: &FunctionResultSpec<PlannedFamily>) -> String {
match result {
FunctionResultSpec::Annotation(annotation) => annotation
.for_language(Language::Python)
.unwrap_or("typing.Any")
.to_string(),
FunctionResultSpec::Authored(authored_type) => {
python_authored_type_annotation(authored_type)
}
}
}
pub(in crate::generator) fn python_authored_type_annotation(authored_type: &PlannedType) -> String {
match authored_type {
TypeSpec::Bool => "bool".to_string(),
TypeSpec::Int(_) => "int".to_string(),
TypeSpec::Float => "float".to_string(),
TypeSpec::String => "str".to_string(),
TypeSpec::Bytes => "bytes".to_string(),
TypeSpec::TypeParameter(parameter) => parameter.name.clone(),
TypeSpec::Option(inner) => {
format!("{} | None", python_authored_type_annotation(inner))
}
TypeSpec::List(inner) => {
format!(
"collections.abc.Sequence[{}]",
python_authored_type_annotation(inner)
)
}
TypeSpec::Tuple(items) => format!(
"tuple[{}]",
items
.iter()
.map(python_authored_type_annotation)
.collect::<Vec<_>>()
.join(", ")
),
TypeSpec::Map(key, value) => format!(
"collections.abc.Mapping[{}, {}]",
python_authored_type_annotation(key),
python_authored_type_annotation(value)
),
TypeSpec::Result { ok, err } => python_result_annotation(
ok.as_ref()
.map(|ok| python_authored_type_annotation(ok).to_string()),
err.as_ref()
.map(|err| python_authored_type_annotation(err).to_string()),
),
TypeSpec::External(ExternalTypeSpec::Alias(AliasTypeSpec {
type_name, target, ..
})) => type_name
.for_language(Language::Python)
.map(str::to_string)
.unwrap_or_else(|| python_authored_type_annotation(target)),
TypeSpec::External(ExternalTypeSpec::Proto(proto_name)) => proto_name
.full_name()
.rsplit('.')
.next()
.unwrap_or(proto_name.full_name())
.to_string(),
TypeSpec::External(ExternalTypeSpec::Json(json_type)) => json_type.model_name.clone(),
TypeSpec::Record(name) => name
.full_name
.as_str()
.rsplit('.')
.next()
.unwrap_or(&name.full_name)
.to_upper_camel_case(),
TypeSpec::Enum(name) => name
.full_name
.as_str()
.rsplit('.')
.next()
.unwrap_or(&name.full_name)
.to_upper_camel_case(),
TypeSpec::Flags(name) => name
.full_name
.as_str()
.rsplit('.')
.next()
.unwrap_or(&name.full_name)
.to_upper_camel_case(),
TypeSpec::Variant(name) => name
.full_name
.as_str()
.rsplit('.')
.next()
.unwrap_or(&name.full_name)
.to_upper_camel_case(),
TypeSpec::Resource(name) => name
.type_name
.as_str()
.rsplit('.')
.next()
.unwrap_or(&name.type_name)
.to_upper_camel_case(),
}
}
fn python_generic_record_annotation(
model_name: &str,
parameters: &[crate::spec::TypeParameterUsage],
) -> String {
if parameters.is_empty() {
model_name.to_string()
} else {
format!(
"{}[{}]",
model_name,
parameters
.iter()
.map(|usage| usage.parameter.name.as_str())
.collect::<Vec<_>>()
.join(", ")
)
}
}
fn python_erased_generic_record_annotation(
model_name: &str,
parameters: &[crate::spec::TypeParameterUsage],
) -> String {
if parameters.is_empty() {
model_name.to_string()
} else {
format!(
"{}[{}]",
model_name,
std::iter::repeat_n("typing.Any", parameters.len())
.collect::<Vec<_>>()
.join(", ")
)
}
}
fn python_function_args(args: &FunctionArgsSpec<PlannedFamily>) -> RenderedFunctionArgs {
match args {
FunctionArgsSpec::Varargs { prefix, .. } => RenderedFunctionArgs::Varargs {
prefix: prefix
.iter()
.map(|arg| RenderedFunctionArg {
name: python_field_name(&arg.name),
annotation: python_authored_type_annotation(&arg.field_type),
})
.collect(),
},
FunctionArgsSpec::Fixed(args) => RenderedFunctionArgs::Typed {
parameters: args
.iter()
.map(|arg| RenderedFunctionArg {
name: python_field_name(&arg.name),
annotation: python_authored_type_annotation(&arg.field_type),
})
.collect(),
},
}
}
fn python_function_args_field_annotation(args: &FunctionArgsSpec<PlannedFamily>) -> String {
match python_function_args(args) {
RenderedFunctionArgs::Varargs { .. } => "list[typing.Any]".to_string(),
RenderedFunctionArgs::Typed { parameters } => python_function_args_list_annotation(
parameters
.iter()
.map(|parameter| parameter.annotation.as_str())
.collect::<Vec<_>>()
.as_slice(),
),
}
}
fn python_function_arg_field_annotation(
function: &FunctionFieldSpec<PlannedFamily>,
field_name: &str,
) -> String {
match &function.args {
FunctionArgsSpec::Varargs { .. } => "list[typing.Any]".to_string(),
FunctionArgsSpec::Fixed(args) => args
.iter()
.find(|arg| arg.name.to_snake_case() == field_name)
.map(|arg| python_authored_type_annotation(&arg.field_type))
.unwrap_or_else(|| python_function_args_field_annotation(&function.args)),
}
}
fn python_function_args_list_annotation(annotations: &[&str]) -> String {
let Some(first) = annotations.first() else {
return "list[typing.Any]".to_string();
};
if annotations.iter().all(|annotation| annotation == first) {
format!("list[{first}]")
} else {
"list[typing.Any]".to_string()
}
}
fn python_result_annotation(ok: Option<String>, err: Option<String>) -> String {
let ok_case = if let Some(ok) = ok {
format!("tuple[typing.Literal[\"ok\"], {ok}]")
} else {
"tuple[typing.Literal[\"ok\"]]".to_string()
};
let err_case = if let Some(err) = err {
format!("tuple[typing.Literal[\"err\"], {err}]")
} else {
"tuple[typing.Literal[\"err\"]]".to_string()
};
format!("{ok_case} | {err_case}")
}
fn planned_record<'a>(api_plan: &'a PlannedSpec, full_name: &str) -> &'a RecordSpec<PlannedFamily> {
api_plan
.record(full_name)
.unwrap_or_else(|| panic!("planned record should exist for {full_name}"))
}
fn register_unpacked_parameter_name(
parameter_sources: &mut BTreeMap<String, String>,
parameter_name: &str,
source: &str,
model_name: &str,
) -> Result<()> {
if let Some(conflicting_field) = parameter_sources.get(parameter_name) {
return Err(Error::FlattenedApiFieldConflict {
type_name: model_name.to_string(),
field: parameter_name.to_string(),
conflicting_field: conflicting_field.clone(),
});
}
parameter_sources.insert(parameter_name.to_string(), source.to_string());
Ok(())
}
fn python_field_annotation(
record: &RecordSpec<PlannedFamily>,
field_name: &str,
field: &RecordFieldSpec<PlannedFamily>,
default_base_annotation: String,
is_optional: bool,
) -> String {
if let Some(function) = record.function_for_args_field(field_name) {
let args_annotation = python_function_arg_field_annotation(function, field_name);
return if is_optional {
format!("{args_annotation} | None")
} else {
args_annotation
};
}
if let Some(annotation) = &field.annotation {
if let Some(annotation) = annotation.for_language(Language::Python) {
return annotation.to_string();
}
}
if let Some(function) = &field.function {
let result_annotation = erase_python_type_parameters(
&python_function_result_annotation(&function.result),
&function
.result_type_parameter
.iter()
.cloned()
.collect::<BTreeSet<_>>(),
);
return python_function_annotation_from_args(
&python_function_args(&function.args),
function
.alternate_type
.as_ref()
.map(python_authored_type_annotation)
.as_deref(),
&result_annotation,
);
}
if is_optional {
format!("{default_base_annotation} | None")
} else {
default_base_annotation
}
}
fn python_function_field_annotation(
function: &FunctionFieldSpec<PlannedFamily>,
alternate_annotation: Option<String>,
) -> String {
let result_annotation = erase_python_type_parameters(
&python_function_result_annotation(&function.result),
&function
.result_type_parameter
.iter()
.cloned()
.collect::<BTreeSet<_>>(),
);
python_function_annotation_from_args(
&python_function_args(&function.args),
alternate_annotation.as_deref(),
&result_annotation,
)
}
fn python_function_annotation_from_args(
args: &RenderedFunctionArgs,
alternate_annotation: Option<&str>,
result_annotation: &str,
) -> String {
let callable_annotation = match args {
RenderedFunctionArgs::Varargs { .. } => {
format!("collections.abc.Callable[..., {}]", result_annotation)
}
RenderedFunctionArgs::Typed { parameters } => format!(
"collections.abc.Callable[[{}], {}]",
parameters
.iter()
.map(|parameter| parameter.annotation.clone())
.collect::<Vec<_>>()
.join(", "),
result_annotation
),
};
if let Some(alternate_annotation) = alternate_annotation {
format!("{alternate_annotation} | {callable_annotation}")
} else {
callable_annotation
}
}
pub(in crate::generator) fn enum_default_expr(
resolved_type: &ResolvedFieldType,
enum_case: &str,
) -> String {
let value_name = if resolved_type.wire_conversion.is_some() {
enum_case.to_shouty_snake_case()
} else {
enum_case.to_upper_camel_case()
};
format!("{}.{}", resolved_type.annotation, value_name)
}
fn render_definitions_only_package_init(
services: &[RenderedService<'_>],
model_names: &[String],
root_package_imports: &RootPackageImports,
) -> String {
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
let service_names = services
.iter()
.map(|service| service.name.to_string())
.collect::<Vec<_>>();
render_root_package_imports(&mut output, root_package_imports);
if !model_names.is_empty() {
render_named_python_import(&mut output, ".models", model_names);
}
if !service_names.is_empty() {
render_named_python_import(&mut output, ".services", &service_names);
}
output.push_str("\n__all__ = [\n");
for name in root_package_export_names(root_package_imports)
.iter()
.map(String::as_str)
.chain(model_names.iter().map(String::as_str))
.chain(service_names.iter().map(String::as_str))
{
output.push_str(" ");
output.push_str(&python_string_literal(name));
output.push_str(",\n");
}
output.push_str("]\n");
output
}
fn package_export_names(
services: &[RenderedService<'_>],
model_names: &[String],
root_package_imports: &RootPackageImports,
) -> BTreeSet<String> {
let mut names = root_package_export_names(root_package_imports);
names.extend(model_names.iter().cloned());
match crate::nexgen_config::current().mode {
GenerationMode::DefinitionsOnly => {
names.extend(services.iter().map(|service| service.name.to_string()));
}
GenerationMode::NativeApi => {
names.extend(
services
.iter()
.filter(|service| service.endpoint.is_none())
.map(endpoint_service_object_name),
);
names.extend(
services
.iter()
.filter(|service| service.endpoint.is_some())
.flat_map(|service| {
service
.operations
.iter()
.map(|operation| operation.attr_name.clone())
}),
);
}
}
names
}
fn operation_key(service_name: &str, operation_name: &str) -> String {
format!("{service_name}::{operation_name}")
}
fn resource_module_name(resource: &PlannedResource) -> String {
python_ident(&resource.name.to_snake_case())
}
fn resource_client_bind_function_name(resource: &PlannedResource) -> String {
format!(
"bind_{}_client",
python_ident(&resource.name.to_snake_case())
)
}
fn resource_operation_owners(services: &[RenderedService<'_>]) -> BTreeMap<String, String> {
let mut owners = BTreeMap::new();
for service in services {
for resource in &service.resources {
let resource_module = resource_module_name(resource);
for method in &resource.methods {
if let PlannedResourceMethodBindingSpec::Operation { operation_name, .. } =
&method.binding
{
owners.insert(
operation_key(service.name, operation_name),
resource_module.clone(),
);
}
}
}
}
owners
}
fn render_support_package(
files: &mut GeneratedFileMap,
support_fragments: &[SupportFragmentSpec],
) -> Result<Vec<String>> {
let mut module_names = Vec::new();
for fragment in support_fragments {
let module_name = support_module_name(fragment)?;
files.insert(
format!("_support/{module_name}.py"),
fragment.contents.clone(),
GeneratedFileOrigin::support_fragment(Language::Python, &fragment.path),
)?;
module_names.push(module_name);
}
if !module_names.is_empty() {
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
for module_name in &module_names {
output.push_str("from .");
output.push_str(module_name);
output.push_str(" import * # noqa: F401,F403\n");
}
files.insert(
"_support/__init__.py",
output,
GeneratedFileOrigin::fixed("generated Python support package initializer"),
)?;
}
Ok(module_names)
}
fn support_module_name(fragment: &SupportFragmentSpec) -> Result<String> {
let path = Path::new(&fragment.path);
let Some(file_name) = path.file_name() else {
return Err(Error::InvalidGeneratedPath {
path: path.to_path_buf(),
reason: "support path must have a file name".to_string(),
});
};
let Some(file_name) = file_name.to_str() else {
return Err(Error::InvalidGeneratedPath {
path: path.to_path_buf(),
reason: "support path must be valid UTF-8".to_string(),
});
};
let Some(module_name) = Path::new(file_name).file_stem() else {
return Err(Error::InvalidGeneratedPath {
path: path.to_path_buf(),
reason: "support path must have a module stem".to_string(),
});
};
let Some(module_name) = module_name.to_str() else {
return Err(Error::InvalidGeneratedPath {
path: path.to_path_buf(),
reason: "support path must be valid UTF-8".to_string(),
});
};
if Path::new(file_name)
.extension()
.and_then(|extension| extension.to_str())
!= Some("py")
{
return Err(Error::InvalidGeneratedPath {
path: path.to_path_buf(),
reason: "Python support files must end with `.py`".to_string(),
});
}
if module_name.is_empty() {
return Err(Error::InvalidGeneratedPath {
path: path.to_path_buf(),
reason: "support module name cannot be empty".to_string(),
});
}
Ok(module_name.to_string())
}
pub(in crate::generator) fn render_generated_file_header(output: &mut String) {
output.push_str(GENERATED_HEADER);
output.push_str("\n\n");
output.push_str("from __future__ import annotations\n");
}
fn render_python_module_imports(output: &mut String, module_imports: &BTreeSet<String>) {
render_python_module_imports_with_indent(output, module_imports, "");
}
fn render_python_module_imports_with_indent(
output: &mut String,
module_imports: &BTreeSet<String>,
indent: &str,
) {
for import in module_imports {
output.push_str(indent);
if let Some(relative_import) = import.strip_prefix('.') {
output.push_str("from . import ");
output.push_str(relative_import.trim_start_matches('.'));
output.push('\n');
} else {
output.push_str("import ");
output.push_str(import);
output.push('\n');
}
}
}
fn body_uses_python_module_path(body: &str, module_path: &str) -> bool {
let dotted_module = format!("{module_path}.");
body.match_indices(&dotted_module).any(|(index, _)| {
let previous = body[..index].chars().next_back();
!previous.is_some_and(|character| is_python_identifier_char(character) || character == '.')
})
}
fn used_python_module_imports(body: &str, module_imports: &BTreeSet<String>) -> BTreeSet<String> {
module_imports
.iter()
.filter(|module_import| body_uses_python_module_path(body, module_import))
.cloned()
.collect()
}
fn is_python_identifier_char(character: char) -> bool {
character.is_ascii_alphanumeric() || character == '_'
}
fn body_uses_python_symbol(body: &str, symbol: &str) -> bool {
body.match_indices(symbol).any(|(index, _)| {
let previous = body[..index].chars().next_back();
let next = body[index + symbol.len()..].chars().next();
!previous.is_some_and(|character| is_python_identifier_char(character) || character == '.')
&& !next.is_some_and(is_python_identifier_char)
})
}
fn used_python_symbol_imports(body: &str, candidates: &[String]) -> Vec<String> {
candidates
.iter()
.filter(|candidate| body_uses_python_symbol(body, candidate))
.cloned()
.collect()
}
pub(in crate::generator) fn render_named_python_import(
output: &mut String,
module: &str,
names: &[String],
) {
render_named_python_import_with_indent(output, module, names, "");
}
fn render_root_package_imports(output: &mut String, imports: &RootPackageImports) {
for (module, names) in imports {
render_named_python_import(output, module, &names.iter().cloned().collect::<Vec<_>>());
}
}
fn root_package_export_names(imports: &RootPackageImports) -> BTreeSet<String> {
imports
.values()
.flat_map(|names| names.iter().cloned())
.collect()
}
fn extend_root_package_imports(target: &mut RootPackageImports, source: RootPackageImports) {
for (module, names) in source {
target.entry(module).or_default().extend(names);
}
}
fn render_named_python_import_with_indent(
output: &mut String,
module: &str,
names: &[String],
indent: &str,
) {
if names.is_empty() {
return;
}
if names.len() == 1 {
output.push_str(indent);
output.push_str("from ");
output.push_str(module);
output.push_str(" import ");
output.push_str(&names[0]);
output.push('\n');
return;
}
output.push_str(indent);
output.push_str("from ");
output.push_str(module);
output.push_str(" import (\n");
for name in names {
output.push_str(indent);
output.push_str(" ");
output.push_str(name);
output.push_str(",\n");
}
output.push_str(indent);
output.push_str(")\n");
}
fn top_level_python_symbols(source: &str) -> BTreeSet<String> {
let mut symbols = BTreeSet::new();
for line in source.lines() {
let line = line.trim_end();
if line.is_empty()
|| line.starts_with(' ')
|| line.starts_with('\t')
|| line.starts_with('#')
|| line.starts_with('@')
|| line.starts_with("import ")
|| line.starts_with("from ")
{
continue;
}
if let Some(name) = line
.strip_prefix("def ")
.and_then(|line| line.split('(').next())
{
symbols.insert(name.trim().to_string());
continue;
}
if let Some(name) = line
.strip_prefix("class ")
.and_then(|line| line.split(['(', ':']).next())
{
symbols.insert(name.trim().to_string());
continue;
}
if let Some((name, _)) = line.split_once('=') {
let name = name.trim();
if !name.is_empty() && name.chars().all(is_python_identifier_char) {
symbols.insert(name.to_string());
}
}
}
symbols
}
fn support_export_names(support_fragments: &[SupportFragmentSpec]) -> Vec<String> {
support_fragments
.iter()
.flat_map(|fragment| top_level_python_symbols(&fragment.contents))
.filter(|name| !name.starts_with('_'))
.collect::<BTreeSet<_>>()
.into_iter()
.collect()
}
pub(in crate::generator) fn render_optional_python_imports(
output: &mut String,
body: &str,
module_imports: &BTreeSet<String>,
language_imports: &[LanguageImportSpec],
) -> bool {
render_optional_python_imports_with_skipped_language_modules(
output,
body,
module_imports,
language_imports,
&BTreeSet::new(),
)
}
fn render_optional_python_imports_with_skipped_language_modules(
output: &mut String,
body: &str,
module_imports: &BTreeSet<String>,
language_imports: &[LanguageImportSpec],
skipped_language_modules: &BTreeSet<String>,
) -> bool {
let simple_imports = [
("import collections.abc\n", "collections.abc"),
("import dataclasses\n", "dataclasses"),
("import enum\n", "enum"),
("import typing\n", "typing"),
("import typing_extensions\n", "typing_extensions"),
];
let mut wrote_any = false;
for (import_line, module_path) in simple_imports {
if body_uses_python_module_path(body, module_path) {
output.push_str(import_line);
wrote_any = true;
}
}
let mut used_module_imports = used_python_module_imports(body, module_imports);
for (_, module_path) in simple_imports {
if body_uses_python_module_path(body, module_path) {
used_module_imports.remove(module_path);
}
}
wrote_any |= render_language_python_imports(
output,
body,
language_imports,
&used_module_imports,
skipped_language_modules,
);
if !used_module_imports.is_empty() {
render_python_module_imports(output, &used_module_imports);
wrote_any = true;
}
wrote_any
}
fn render_language_python_imports(
output: &mut String,
body: &str,
language_imports: &[LanguageImportSpec],
skipped_module_imports: &BTreeSet<String>,
skipped_language_modules: &BTreeSet<String>,
) -> bool {
let mut module_imports = BTreeSet::new();
let mut named_imports = BTreeMap::<String, BTreeSet<String>>::new();
for import in language_imports {
if skipped_language_modules.contains(&import.module) {
continue;
}
let used = match import.import_style {
LanguageImportStyle::Module | LanguageImportStyle::Namespace => {
body_uses_python_module_path(body, &import.reference)
}
LanguageImportStyle::Named => body_uses_python_module_path(body, &import.reference),
};
if !used {
continue;
}
match import.import_style {
LanguageImportStyle::Module | LanguageImportStyle::Namespace => {
if !skipped_module_imports.contains(&import.module) {
module_imports.insert(import.module.clone());
}
}
LanguageImportStyle::Named => {
named_imports
.entry(import.module.clone())
.or_default()
.insert(
import
.name
.clone()
.unwrap_or_else(|| import.reference.clone()),
);
}
}
}
if module_imports.is_empty() && named_imports.is_empty() {
return false;
}
render_python_module_imports(output, &module_imports);
for (module, names) in named_imports {
render_named_python_import(
&mut *output,
&module,
&names.into_iter().collect::<Vec<_>>(),
);
}
true
}
fn delayed_temporalio_workflow_skipped_language_modules(
service: &RenderedService<'_>,
) -> BTreeSet<String> {
if service.delay_load_temporalio_workflow {
BTreeSet::from(["temporalio.workflow".to_string()])
} else {
BTreeSet::new()
}
}
fn temporalio_workflow_type_annotation(service: &RenderedService<'_>, annotation: &str) -> String {
if !service.delay_load_temporalio_workflow {
return annotation.to_string();
}
annotation
.replace(
"temporalio.workflow.NexusOperationHandle",
"NexusOperationHandle",
)
.replace(
"temporalio.workflow.ExternalWorkflowHandle",
"ExternalWorkflowHandle",
)
}
fn temporalio_workflow_type_ref(service: &RenderedService<'_>, name: &str) -> String {
if service.delay_load_temporalio_workflow {
name.to_string()
} else {
format!("temporalio.workflow.{name}")
}
}
fn temporalio_workflow_value_expr(service: &RenderedService<'_>, expression: &str) -> String {
if service.delay_load_temporalio_workflow {
expression.replace("temporalio.workflow.", "")
} else {
expression.to_string()
}
}
fn temporalio_workflow_runtime_names(expression: &str) -> Vec<String> {
const PREFIX: &str = "temporalio.workflow.";
let mut names = BTreeSet::new();
let mut index = 0;
while let Some(relative_start) = expression[index..].find(PREFIX) {
let name_start = index + relative_start + PREFIX.len();
let mut name_end = name_start;
for character in expression[name_start..].chars() {
if is_python_identifier_char(character) {
name_end += character.len_utf8();
} else {
break;
}
}
if name_end > name_start {
names.insert(expression[name_start..name_end].to_string());
}
index = name_end;
}
names.into_iter().collect()
}
fn render_temporalio_workflow_runtime_imports(
output: &mut String,
service: &RenderedService<'_>,
names: &[String],
indent: &str,
) {
if !service.delay_load_temporalio_workflow || names.is_empty() {
return;
}
render_named_python_import_with_indent(output, "temporalio.workflow", names, indent);
}
fn temporalio_workflow_operation_runtime_names(
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) -> Vec<String> {
if !service.delay_load_temporalio_workflow {
return Vec::new();
}
let mut names = BTreeSet::from(["create_nexus_client".to_string()]);
if let Some(transform_expr) = &operation.output_transform_expr {
names.extend(temporalio_workflow_runtime_names(transform_expr));
}
names.into_iter().collect()
}
fn temporalio_workflow_type_checking_names(body: &str) -> Vec<String> {
["ExternalWorkflowHandle", "NexusOperationHandle"]
.into_iter()
.filter(|name| body_uses_python_symbol(body, name))
.map(ToString::to_string)
.collect()
}
fn render_temporalio_workflow_type_checking_imports(
output: &mut String,
service: &RenderedService<'_>,
body: &str,
) -> bool {
if !service.delay_load_temporalio_workflow {
return false;
}
let names = temporalio_workflow_type_checking_names(body);
if names.is_empty() {
return false;
}
output.push_str("if typing.TYPE_CHECKING:\n");
render_named_python_import_with_indent(output, "temporalio.workflow", &names, " ");
true
}
fn render_models_module(
enums: &[&RenderedEnum],
flags: &[&RenderedFlags],
variants: &[&RenderedVariant],
model_fragments: &RenderedModelFragments,
support_names: &[String],
language_imports: &[LanguageImportSpec],
api_plan: &PlannedSpec,
inline_model_rebuilds: bool,
model_hoists: Option<&PythonModelHoists>,
) -> Result<Option<String>> {
let mut module_imports = BTreeSet::new();
module_imports.extend(model_fragments.module_imports.iter().cloned());
let mut body = String::new();
if !enums.is_empty() {
if !body.is_empty() {
body.push_str("\n\n");
}
for (index, enumeration) in enums.iter().enumerate() {
render_enum(&mut body, enumeration);
if index + 1 != enums.len() {
body.push_str("\n\n");
}
}
}
if !flags.is_empty() {
if !body.is_empty() {
body.push_str("\n\n");
}
for (index, flag_set) in flags.iter().enumerate() {
render_flags(&mut body, flag_set);
if index + 1 != flags.len() {
body.push_str("\n\n");
}
}
}
if !model_fragments.body.is_empty() {
if !body.is_empty() {
body.push_str("\n\n");
}
body.push_str(&model_fragments.body);
}
if !variants.is_empty() {
if !body.is_empty() {
body.push_str("\n\n");
}
let variant_only_type_parameters = variants
.iter()
.flat_map(|variant| variant.type_parameters.iter())
.filter(|parameter| {
!model_fragments
.declared_type_parameters
.contains(parameter.as_str())
})
.cloned()
.collect::<BTreeSet<_>>();
for parameter in &variant_only_type_parameters {
body.push_str(parameter);
body.push_str(" = typing.TypeVar(");
body.push_str(&python_string_literal(parameter));
body.push_str(")\n");
}
if !variant_only_type_parameters.is_empty() {
body.push_str("\n\n");
}
for (index, variant) in variants.iter().enumerate() {
render_variant(&mut body, variant);
if index + 1 != variants.len() {
body.push_str("\n\n");
}
}
}
if inline_model_rebuilds && !model_fragments.post_model_statements.is_empty() {
if !body.is_empty() {
body.push_str("\n\n");
}
body.push_str(&model_fragments.post_model_statements);
}
// Nothing was declared, so there is no models module to emit.
if body.is_empty() {
return Ok(None);
}
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
let import_scan_body = if model_hoists.is_none() && api_plan.data.module_imports.is_empty() {
body.clone()
} else {
format!("{body}\ntyping.TYPE_CHECKING\n")
};
let wrote_imports = render_optional_python_imports(
&mut output,
&import_scan_body,
&module_imports,
language_imports,
);
let mut wrote_relative_imports = false;
for (module, names) in &model_fragments.relative_imports {
if names.is_empty() {
continue;
}
if wrote_imports || wrote_relative_imports {
output.push('\n');
}
render_named_python_import(
&mut output,
module,
&names.iter().cloned().collect::<Vec<_>>(),
);
wrote_relative_imports = true;
}
let wrote_module_model_imports = if let Some(model_hoists) = model_hoists {
render_python_module_model_runtime_imports(
&mut output,
api_plan,
model_hoists,
&body,
wrote_imports || wrote_relative_imports,
)
} else {
render_python_module_model_type_checking_imports(
&mut output,
api_plan,
wrote_imports || wrote_relative_imports,
)
};
let used_support_names = used_python_symbol_imports(&body, support_names);
if !used_support_names.is_empty() {
if wrote_imports || wrote_relative_imports || wrote_module_model_imports {
output.push('\n');
}
render_named_python_import(&mut output, "._support", &used_support_names);
}
output.push('\n');
output.push('\n');
output.push_str(&body);
Ok(Some(output))
}
fn render_resources_package_init(services: &[RenderedService<'_>]) -> String {
let resources = services
.iter()
.flat_map(|service| service.resources.iter())
.collect::<Vec<_>>();
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
for resource in &resources {
render_named_python_import(
&mut output,
&format!(".{}", resource_module_name(resource)),
std::slice::from_ref(&resource.type_name),
);
}
output.push_str("\n__all__ = [\n");
for resource in &resources {
output.push_str(" ");
output.push_str(&python_string_literal(&resource.type_name));
output.push_str(",\n");
}
output.push_str("]\n");
output
}
fn resource_bound_operations<'a>(
service: &'a RenderedService<'a>,
resource: &'a PlannedResource,
) -> Vec<ResourceBoundOperation<'a>> {
resource
.methods
.iter()
.filter_map(|method| match &method.binding {
PlannedResourceMethodBindingSpec::Operation { operation_name, .. } => {
Some(ResourceBoundOperation {
operation: service
.operations
.iter()
.find(|operation| operation.name == operation_name)
.expect("bound resource operation should exist on the service"),
})
}
PlannedResourceMethodBindingSpec::Stub => None,
})
.collect()
}
fn render_service_module(
services: &[RenderedService<'_>],
api_plan: &PlannedSpec,
model_hoists: Option<&PythonModelHoists>,
include_endpoint_clients: bool,
) -> String {
let mut module_imports = BTreeSet::new();
let mut model_type_imports = BTreeSet::new();
let resource_modules = services
.iter()
.flat_map(|service| service.resources.iter())
.map(|resource| (resource.type_name.clone(), resource_module_name(resource)))
.collect::<BTreeMap<_, _>>();
let mut resource_type_imports = BTreeMap::new();
for service in services {
for operation in &service.operations {
if let Some(input_module_path) = operation
.input
.as_ref()
.and_then(|input| input.module_path.as_ref())
{
if input_module_path == ".models" {
if let Some(type_name) =
operation_input_type_ref(operation).strip_prefix("models.")
{
model_type_imports.insert(type_name.to_string());
}
} else if model_hoists.is_some()
&& input_module_path == &root_python_model_hoist_module(&api_plan.module_path)
{
let type_ref = operation_input_type_ref(operation);
if !type_ref.contains('.') {
model_type_imports.insert(type_ref.to_string());
} else if let Some(type_name) = type_ref.strip_prefix("_recursive.") {
model_type_imports.insert(type_name.to_string());
}
} else {
module_imports.insert(input_module_path.clone());
}
}
if let Some(output_module_path) = &operation.output_module_path {
if output_module_path == "._resources" {
if let Some(module_name) = resource_modules.get(&operation.output_ref) {
resource_type_imports
.insert(operation.output_ref.clone(), module_name.clone());
}
} else if output_module_path == ".models" {
if let Some(type_name) = operation.output_ref.strip_prefix("models.") {
model_type_imports.insert(type_name.to_string());
}
} else if model_hoists.is_some()
&& output_module_path == &root_python_model_hoist_module(&api_plan.module_path)
{
if !operation.output_ref.contains('.') {
model_type_imports.insert(operation.output_ref.clone());
} else if let Some(type_name) = operation.output_ref.strip_prefix("_recursive.")
{
model_type_imports.insert(type_name.to_string());
}
} else {
module_imports.insert(output_module_path.clone());
}
}
}
}
let endpoint_services = services
.iter()
.filter(|service| include_endpoint_clients && service.endpoint.is_none())
.collect::<Vec<_>>();
let mut endpoint_client_body = String::new();
for (index, service) in endpoint_services.iter().enumerate() {
render_endpoint_service_object(&mut endpoint_client_body, service);
if index + 1 != endpoint_services.len() {
endpoint_client_body.push('\n');
}
}
let runtime_model_names = endpoint_services
.iter()
.flat_map(|service| {
service.operations.iter().filter_map(|operation| {
operation.unpacked_input.as_ref().map(|input| {
std::iter::once(input.model_name.clone())
.chain(
input
.flattened_messages
.iter()
.map(|message| message.model_name.clone()),
)
.collect::<Vec<_>>()
})
})
})
.flatten()
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let used_runtime_model_names =
used_python_symbol_imports(&endpoint_client_body, &runtime_model_names);
model_type_imports.extend(used_runtime_model_names.iter().cloned());
let type_checking_model_names = endpoint_services
.iter()
.flat_map(|service| {
service.operations.iter().flat_map(|operation| {
let input_model = operation
.input
.as_ref()
.and_then(|input| model_type_name(&input.type_ref));
let output_model = (!resource_modules.contains_key(&operation.output_ref))
.then(|| model_type_name(&operation.output_ref))
.flatten();
input_model.into_iter().chain(output_model)
})
})
.filter(|name| !model_type_imports.contains(name))
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let used_type_checking_model_names =
used_python_symbol_imports(&endpoint_client_body, &type_checking_model_names);
let resource_names = resource_modules.keys().cloned().collect::<Vec<_>>();
let used_resource_names = used_python_symbol_imports(&endpoint_client_body, &resource_names);
let needs_temporalio_workflow_runtime = endpoint_client_body.contains("temporalio.workflow.");
let needs_temporalio_workflow_type_checking = false;
let needs_type_checking_imports = needs_temporalio_workflow_type_checking
|| !used_type_checking_model_names.is_empty()
|| !used_resource_names.is_empty();
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
let uses_service = !services.is_empty();
let uses_operation = services
.iter()
.any(|service| !service.operations.is_empty());
if uses_operation {
output.push_str("from nexusrpc import Operation, service\n");
} else if uses_service {
output.push_str("from nexusrpc import service\n");
}
let uses_typing = services
.iter()
.flat_map(|service| service.operations.iter())
.any(|operation| {
operation
.input
.as_ref()
.is_some_and(|input| input.descriptor_type_ref.contains("typing."))
|| operation.descriptor_output_ref.contains("typing.")
// A deprecated operation is declared as
// `typing.Annotated[Operation[...], ...]`, and `nexusrpc`'s
// `@service` resolves the annotation with `eval_str=True`, so
// the name has to be bound at runtime, not just under
// `TYPE_CHECKING`.
|| operation.deprecated
});
if uses_typing || endpoint_client_body.contains("typing.") || needs_type_checking_imports {
output.push_str("import typing\n");
}
if services.iter().any(|service| {
service.deprecated
|| service
.operations
.iter()
.any(|operation| operation.deprecated)
}) {
output.push_str("import typing_extensions\n");
}
if needs_temporalio_workflow_runtime {
output.push_str("import temporalio.workflow\n");
}
let mut resource_imports_by_module = BTreeMap::<String, BTreeSet<String>>::new();
for (resource_type_name, module_name) in &resource_type_imports {
resource_imports_by_module
.entry(module_name.clone())
.or_default()
.insert(resource_type_name.clone());
}
for service in &endpoint_services {
for resource in &service.resources {
let binder_name = resource_client_bind_function_name(resource);
if endpoint_client_body.contains(&binder_name) {
resource_imports_by_module
.entry(resource_module_name(resource))
.or_default()
.insert(binder_name);
}
}
}
if !module_imports.is_empty()
|| !model_type_imports.is_empty()
|| !resource_imports_by_module.is_empty()
{
if uses_typing || endpoint_client_body.contains("typing.") || needs_type_checking_imports {
output.push('\n');
}
render_python_module_imports(&mut output, &module_imports);
render_model_name_imports(
&mut output,
".models",
&api_plan.module_path,
api_plan,
&model_type_imports.into_iter().collect::<Vec<_>>(),
model_hoists,
);
for (module_name, names) in &resource_imports_by_module {
render_named_python_import(
&mut output,
&format!("._resources.{module_name}"),
&names.iter().cloned().collect::<Vec<_>>(),
);
}
}
if needs_temporalio_workflow_type_checking {
output.push_str("if typing.TYPE_CHECKING:\n");
output.push_str(" import temporalio.workflow\n");
render_model_name_imports_with_indent(
&mut output,
".models",
&api_plan.module_path,
api_plan,
&used_type_checking_model_names,
model_hoists,
" ",
);
render_named_python_import_with_indent(
&mut output,
"._resources",
&used_resource_names,
" ",
);
} else if needs_type_checking_imports {
output.push_str("if typing.TYPE_CHECKING:\n");
render_model_name_imports_with_indent(
&mut output,
".models",
&api_plan.module_path,
api_plan,
&used_type_checking_model_names,
model_hoists,
" ",
);
render_named_python_import_with_indent(
&mut output,
"._resources",
&used_resource_names,
" ",
);
}
if !services.is_empty() {
output.push_str("\n\n");
for (index, service) in services.iter().enumerate() {
render_service_definition(&mut output, service);
if index + 1 != services.len() {
output.push_str("\n\n");
}
}
}
if !endpoint_client_body.is_empty() {
output.push('\n');
output.push_str(&endpoint_client_body);
}
output
}
fn render_operation_module(
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
model_names: &[String],
resource_names: &[String],
support_names: &[String],
language_imports: &[LanguageImportSpec],
api_plan: &PlannedSpec,
model_hoists: Option<&PythonModelHoists>,
) -> String {
let mut module_imports = BTreeSet::new();
if let Some(input_module_path) = operation
.input
.as_ref()
.and_then(|input| input.module_path.as_ref())
{
module_imports.insert(input_module_path.clone());
}
if let Some(output_module_path) = &operation.output_module_path {
module_imports.insert(output_module_path.clone());
}
let mut body = String::new();
for parameter in &operation.model_input_type_parameters {
body.push_str(parameter);
body.push_str(" = typing.TypeVar(");
body.push_str(&python_string_literal(parameter));
body.push_str(")\n");
}
if !operation.model_input_type_parameters.is_empty() {
body.push_str("\n\n");
}
if let Some(unpacked_input) = &operation.unpacked_input
&& !unpacked_input.functions.is_empty()
{
render_function_type_parameter_definitions(
&mut body,
&unpacked_input.functions,
&operation.output_type_parameters,
);
body.push_str("\n\n");
}
render_operation_functions(&mut body, service, operation);
if !service.delay_load_temporalio_workflow {
module_imports.insert("temporalio.workflow".to_string());
}
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
let type_checking_names = temporalio_workflow_type_checking_names(&body);
let body_for_imports =
if service.delay_load_temporalio_workflow && !type_checking_names.is_empty() {
format!("{body}\ntyping.TYPE_CHECKING")
} else {
body.clone()
};
let skipped_language_modules = delayed_temporalio_workflow_skipped_language_modules(service);
let wrote_imports = render_optional_python_imports_with_skipped_language_modules(
&mut output,
&body_for_imports,
&module_imports,
language_imports,
&skipped_language_modules,
);
if wrote_imports && service.delay_load_temporalio_workflow && !type_checking_names.is_empty() {
output.push('\n');
}
let wrote_type_checking_imports =
render_temporalio_workflow_type_checking_imports(&mut output, service, &body);
if wrote_imports || wrote_type_checking_imports {
output.push('\n');
}
let used_model_names = used_python_symbol_imports(&body, model_names);
render_model_name_imports(
&mut output,
"..models",
&api_plan.module_path.child("operations"),
api_plan,
&used_model_names,
model_hoists,
);
let used_resource_names = used_python_symbol_imports(&body, resource_names);
if !used_resource_names.is_empty() {
render_named_python_import(&mut output, ".._resources", &used_resource_names);
}
let used_support_names = used_python_symbol_imports(&body, support_names);
if !used_support_names.is_empty() {
render_named_python_import(&mut output, ".._support", &used_support_names);
}
output.push('\n');
output.push('\n');
output.push_str(&body);
output
}
fn planned_module_import_names(api_plan: &PlannedSpec) -> BTreeSet<String> {
api_plan
.data
.module_imports
.values()
.flat_map(|names| names.iter().cloned())
.collect()
}
fn render_model_name_imports(
output: &mut String,
local_module: &str,
current_module_path: &ModulePath,
api_plan: &PlannedSpec,
names: &[String],
model_hoists: Option<&PythonModelHoists>,
) {
render_model_name_imports_with_indent(
output,
local_module,
current_module_path,
api_plan,
names,
model_hoists,
"",
);
}
fn render_model_name_imports_with_indent(
output: &mut String,
local_module: &str,
current_module_path: &ModulePath,
api_plan: &PlannedSpec,
names: &[String],
model_hoists: Option<&PythonModelHoists>,
indent: &str,
) {
for (module, names) in model_name_import_groups(
local_module,
current_module_path,
api_plan,
names,
model_hoists,
) {
render_named_python_import_with_indent(
output,
&module,
&names.into_iter().collect::<Vec<_>>(),
indent,
);
}
}
fn model_name_import_groups(
local_module: &str,
current_module_path: &ModulePath,
api_plan: &PlannedSpec,
names: &[String],
model_hoists: Option<&PythonModelHoists>,
) -> BTreeMap<String, BTreeSet<String>> {
let foreign_modules = foreign_model_modules_by_name(api_plan);
let mut imports = BTreeMap::<String, BTreeSet<String>>::new();
for name in names {
if let Some(module_path) = foreign_modules.get(name) {
let module = if model_hoists.is_some_and(|hoists| hoists.is_hoisted(module_path, name))
{
root_python_model_hoist_module(current_module_path)
} else {
python_relative_models_module(current_module_path, module_path)
};
imports.entry(module).or_default().insert(name.clone());
} else if model_hoists.is_some_and(|hoists| hoists.is_hoisted(&api_plan.module_path, name))
{
imports
.entry(root_python_model_hoist_module(current_module_path))
.or_default()
.insert(name.clone());
} else {
imports
.entry(local_module.to_string())
.or_default()
.insert(name.clone());
}
}
imports
}
fn foreign_model_modules_by_name(api_plan: &PlannedSpec) -> BTreeMap<String, ModulePath> {
api_plan
.data
.module_imports
.iter()
.flat_map(|(module_path, names)| {
names
.iter()
.map(|name| (name.clone(), module_path.clone()))
.collect::<Vec<_>>()
})
.collect()
}
fn render_python_module_model_type_checking_imports(
output: &mut String,
api_plan: &PlannedSpec,
wrote_previous_imports: bool,
) -> bool {
if api_plan.data.module_imports.is_empty() {
return false;
}
if wrote_previous_imports {
output.push('\n');
}
output.push_str("if typing.TYPE_CHECKING:\n");
for (module_path, names) in &api_plan.data.module_imports {
let names = names.iter().cloned().collect::<Vec<_>>();
render_named_python_import_with_indent(
output,
&python_relative_models_module(&api_plan.module_path, module_path),
&names,
" ",
);
}
true
}
fn render_python_module_model_runtime_imports(
output: &mut String,
api_plan: &PlannedSpec,
model_hoists: &PythonModelHoists,
body: &str,
wrote_previous_imports: bool,
) -> bool {
let mut imports = BTreeMap::<String, BTreeSet<String>>::new();
// A local reference becomes cross-module after its target is hoisted.
// Seed these imports from the actual model-reference graph retained by the
// hoist plan. Rendered text is not an authority: a docstring can mention a
// type name without creating an edge, while a generated expression may
// require one regardless of how its source happens to be formatted.
for name in model_hoists.runtime_imports_for(&api_plan.module_path) {
imports
.entry(root_python_model_hoist_module(&api_plan.module_path))
.or_default()
.insert(name.clone());
}
for (module_path, names) in &api_plan.data.module_imports {
for name in names {
if !body_uses_python_symbol(body, name) {
continue;
}
let module = if model_hoists.is_hoisted(module_path, name) {
root_python_model_hoist_module(&api_plan.module_path)
} else {
python_relative_models_module(&api_plan.module_path, module_path)
};
imports.entry(module).or_default().insert(name.clone());
}
}
if imports.is_empty() {
return false;
}
if wrote_previous_imports {
output.push('\n');
}
for (module, names) in imports {
render_named_python_import(output, &module, &names.into_iter().collect::<Vec<_>>());
}
true
}
fn root_python_model_hoist_module(module_path: &ModulePath) -> String {
format!("{}{}", ".".repeat(module_path.0.len() + 1), "_recursive")
}
fn python_relative_models_module(from: &ModulePath, to: &ModulePath) -> String {
let common = module_common_prefix_len(&from.0, &to.0);
let dot_count = from.0.len().saturating_sub(common) + 1;
let mut module = ".".repeat(dot_count);
let rest = to.0[common..]
.iter()
.map(|segment| segment.replace('-', "_"))
.chain(std::iter::once("models".to_string()))
.collect::<Vec<_>>();
module.push_str(&rest.join("."));
module
}
pub(in crate::generator) fn module_common_prefix_len(left: &[String], right: &[String]) -> usize {
left.iter()
.zip(right.iter())
.take_while(|(left, right)| left == right)
.count()
}
fn render_operations_package_init() -> String {
let mut output = String::new();
render_generated_file_header(&mut output);
output
}
/// Renders the mixins that turn System Nexus operations into operation-specific
/// workflow-outbound interception points.
fn render_system_nexus_interceptor(services: &[RenderedService<'_>]) -> String {
let operations = services
.iter()
.flat_map(|service| {
service
.operations
.iter()
.map(move |operation| (service, operation))
})
.collect::<Vec<_>>();
let mut output = String::new();
render_generated_file_header(&mut output);
output.push_str("\n\nimport abc\nimport typing\n\nfrom . import models\n\nif typing.TYPE_CHECKING:\n import temporalio.workflow\n from temporalio.worker._interceptor import StartNexusOperationInput\n\n\n__all__ = [\n \"_start_system_nexus_operation\",\n \"_SystemNexusWorkflowOutboundInterceptorBase\",\n \"_SystemNexusWorkflowOutboundInterceptorTerminal\",\n]\n\n\n_InputT = typing.TypeVar(\"_InputT\")\n_OutputT = typing.TypeVar(\"_OutputT\")\n\n\n");
output.push_str(
"async def _start_system_nexus_operation(\n interceptor: _SystemNexusWorkflowOutboundInterceptorBase,\n input: StartNexusOperationInput[_InputT, _OutputT],\n) -> temporalio.workflow.NexusOperationHandle[_OutputT]:\n",
);
for (service, operation) in &operations {
let output_type = system_nexus_type_expr(&operation.output_type_expr);
output.push_str(" if input.service == ");
output.push_str(&python_string_literal(service.wire_name));
output.push_str(" and input.operation_name == ");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(
":\n typed_input = typing.cast(\n \"StartNexusOperationInput[",
);
output.push_str(operation_input_type_ref(operation));
output.push_str(", ");
output.push_str(&output_type);
output.push_str("]\",\n input,\n )\n # The dispatch check above establishes that this operation's response type is _OutputT.\n return typing.cast(\n \"temporalio.workflow.NexusOperationHandle[_OutputT]\",\n await interceptor.start_");
output.push_str(&operation.attr_name);
output.push_str("(typed_input.input),\n )\n");
}
output.push_str(" raise ValueError(f\"unsupported System Nexus operation: {input.service}/{input.operation_name}\")\n");
output.push_str("\n\nclass _SystemNexusWorkflowOutboundInterceptorBase(abc.ABC):\n");
output.push_str(" @abc.abstractmethod\n def _next_system_nexus_interceptor(\n self,\n ) -> _SystemNexusWorkflowOutboundInterceptorBase:\n ...\n");
for (_service, operation) in &operations {
let output_type = system_nexus_type_expr(&operation.output_type_expr);
output.push_str("\n async def start_");
output.push_str(&operation.attr_name);
output.push_str("(\n self, request: ");
output.push_str(operation_input_type_ref(operation));
output.push_str("\n ) -> temporalio.workflow.NexusOperationHandle[");
output.push_str(&output_type);
output.push_str("]:\n");
output.push_str(" \"\"\"Intercept the ");
output.push_str(operation.name);
output.push_str(" operation.\"\"\"\n");
output.push_str(" return await self._next_system_nexus_interceptor().start_");
output.push_str(&operation.attr_name);
output.push_str("(request)\n");
}
output.push_str("\n\nclass _SystemNexusWorkflowOutboundInterceptorTerminal(abc.ABC):\n");
output.push_str(" @abc.abstractmethod\n async def _intercept_system_nexus_operation(\n self,\n input: StartNexusOperationInput[_InputT, _OutputT],\n ) -> temporalio.workflow.NexusOperationHandle[_OutputT]:\n ...\n");
for (service, operation) in operations {
let output_type = system_nexus_type_expr(&operation.output_type_expr);
output.push_str("\n async def start_");
output.push_str(&operation.attr_name);
output.push_str("(\n self, request: ");
output.push_str(operation_input_type_ref(operation));
output.push_str("\n ) -> temporalio.workflow.NexusOperationHandle[");
output.push_str(&output_type);
output.push_str("]:\n");
output.push_str(
" from temporalio.worker._interceptor import StartNexusOperationInput\n",
);
output.push_str(" from temporalio.nexus.system import TEMPORAL_SYSTEM_ENDPOINT\n");
output
.push_str(" from temporalio.workflow import NexusOperationCancellationType\n\n");
output.push_str(" return await self._intercept_system_nexus_operation(\n");
output.push_str(" StartNexusOperationInput(\n endpoint=TEMPORAL_SYSTEM_ENDPOINT,\n service=");
output.push_str(&python_string_literal(service.wire_name));
output.push_str(",\n operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n input=request,\n output_type=");
output.push_str(&output_type);
output.push_str(",\n schedule_to_close_timeout=None,\n schedule_to_start_timeout=None,\n start_to_close_timeout=None,\n cancellation_type=NexusOperationCancellationType.WAIT_COMPLETED,\n headers=None,\n summary=None,\n )\n )\n");
}
output
}
fn system_nexus_type_expr(type_expr: &str) -> String {
if type_expr == "None" || type_expr.contains('.') || type_expr.contains('[') {
type_expr.to_string()
} else {
format!("models.{type_expr}")
}
}
fn function_type_parameters(
functions: &[RenderedFunctionField],
output_type_parameters: &BTreeSet<String>,
) -> Vec<PythonTypeParameter> {
let mut parameters = BTreeSet::new();
for function in functions {
for case in function_overload_cases(function, output_type_parameters) {
parameters.extend(case.type_parameters);
}
}
parameters.into_iter().collect()
}
fn render_function_type_parameter_definitions(
output: &mut String,
functions: &[RenderedFunctionField],
output_type_parameters: &BTreeSet<String>,
) {
let parameters = function_type_parameters(functions, output_type_parameters);
if parameters.is_empty() {
return;
}
for parameter in parameters {
match parameter {
PythonTypeParameter::TypeVar(name) => {
output.push_str(&name);
output.push_str(" = typing.TypeVar(");
output.push_str(&python_string_literal(&name));
output.push_str(")\n");
}
PythonTypeParameter::TypeVarTuple(name) => {
output.push_str(&name);
output.push_str(" = typing_extensions.TypeVarTuple(");
output.push_str(&python_string_literal(&name));
output.push_str(")\n");
}
}
}
}
fn render_package_init(
services: &[RenderedService<'_>],
model_names: &[String],
_resource_names: &[String],
resource_operation_owners: &BTreeMap<String, String>,
support_names: &[String],
_api_plan: &PlannedSpec,
_model_hoists: Option<&PythonModelHoists>,
root_package_imports: &RootPackageImports,
) -> String {
let operation_function_names = services
.iter()
.filter(|service| service.endpoint.is_some())
.flat_map(|service| {
service
.operations
.iter()
.map(|operation| operation.attr_name.clone())
})
.collect::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
let service_object_names = services
.iter()
.filter(|service| service.endpoint.is_none())
.map(endpoint_service_object_name)
.collect::<Vec<_>>();
let operation_registry_body = render_operation_registry_body(services);
let operation_registry_support_names =
used_python_symbol_imports(&operation_registry_body, support_names);
let mut output = String::new();
render_generated_file_header(&mut output);
output.push('\n');
render_root_package_imports(&mut output, root_package_imports);
if !model_names.is_empty() {
render_named_python_import(&mut output, ".models", model_names);
}
if services
.iter()
.any(|service| !service.operations.is_empty())
{
output.push_str("import collections.abc\n");
output.push_str("import typing\n");
output.push_str("\n");
output.push_str("import nexusrpc\n");
output.push_str("import temporalio.converter\n");
}
if !service_object_names.is_empty() {
render_named_python_import(&mut output, ".services", &service_object_names);
}
if services
.iter()
.any(|service| !service.operations.is_empty())
{
output.push_str("from . import services as _services\n");
}
if !services.is_empty() {
for service in services {
for operation in &service.operations {
let module = if let Some(resource_owner) =
resource_operation_owners.get(&operation_key(service.name, operation.name))
{
format!("._resources.{resource_owner}")
} else {
format!(".operations.{}", operation.attr_name)
};
if service.endpoint.is_some() {
render_named_python_import(
&mut output,
&module,
std::slice::from_ref(&operation.attr_name),
);
}
}
}
}
if !operation_registry_support_names.is_empty() {
render_named_python_import(&mut output, "._support", &operation_registry_support_names);
}
output.push_str("\n__all__ = [\n");
for name in root_package_export_names(root_package_imports) {
output.push_str(" ");
output.push_str(&python_string_literal(&name));
output.push_str(",\n");
}
for name in model_names {
output.push_str(" ");
output.push_str(&python_string_literal(name));
output.push_str(",\n");
}
for name in &service_object_names {
output.push_str(" ");
output.push_str(&python_string_literal(name));
output.push_str(",\n");
}
for name in &operation_function_names {
output.push_str(" ");
output.push_str(&python_string_literal(name));
output.push_str(",\n");
}
output.push_str("]\n");
output.push_str(&operation_registry_body);
output
}
fn render_operation_registry_body(services: &[RenderedService<'_>]) -> String {
if !services
.iter()
.any(|service| !service.operations.is_empty())
{
return String::new();
}
let mut output = String::new();
output.push_str("\n\n_InputT = typing.TypeVar(\"_InputT\")\n");
output.push_str("_OutputT = typing.TypeVar(\"_OutputT\")\n");
output.push_str(
"\n\n_SerializationContextFactory = collections.abc.Callable[\n [_InputT], temporalio.converter.SerializationContext\n]\n",
);
output.push_str("\n\nclass _NexusOperationInfo(typing.Generic[_InputT, _OutputT]):\n");
output.push_str(
" def __init__(\n self,\n *,\n operation: nexusrpc.Operation[_InputT, _OutputT],\n serialization_context: _SerializationContextFactory[_InputT] | None = None,\n ) -> None:\n",
);
output.push_str(" self.operation: nexusrpc.Operation[_InputT, _OutputT] = operation\n");
output.push_str(" self.serialization_context: _SerializationContextFactory[_InputT] | None = serialization_context\n");
output.push_str("\n\n__nexus_operation_registry__ = {\n");
for service in services {
for operation in &service.operations {
output.push_str(" (\n");
output.push_str(" \"");
output.push_str(service.wire_name);
output.push_str("\",\n");
output.push_str(" \"");
output.push_str(operation.wire_name);
output.push_str("\",\n");
output.push_str(" ): _NexusOperationInfo(\n");
output.push_str(" operation=_services.");
output.push_str(service.name);
output.push('.');
output.push_str(&operation.attr_name);
output.push_str(",\n");
if let Some(serialization_context) = &operation.serialization_context_expr {
output.push_str(" serialization_context=");
output.push_str(serialization_context);
output.push_str(",\n");
}
output.push_str(" ),\n");
}
}
output.push_str("}\n");
output
}
fn render_endpoint_service_object(output: &mut String, service: &RenderedService<'_>) {
if service.deprecated {
output.push_str(
"\n@typing_extensions.deprecated(\"This service is deprecated.\", category=None)\n",
);
} else {
output.push('\n');
}
output.push_str("class ");
output.push_str(&endpoint_service_object_name(service));
output.push_str(":\n");
output.push_str(" def __init__(self, endpoint: str) -> None:\n");
output.push_str(
" self._nexus_client: typing.Any = temporalio.workflow.create_nexus_client(\n",
);
output.push_str(" service=");
output.push_str(&python_string_literal(service.wire_name));
output.push_str(",\n");
output.push_str(" endpoint=endpoint,\n");
output.push_str(" )\n");
if service.operations.is_empty() {
output.push('\n');
output.push_str(" pass\n");
return;
}
for operation in &service.operations {
output.push('\n');
render_endpoint_service_object_operation_method(output, service, operation);
}
}
fn render_endpoint_service_object_operation_method(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) {
if let Some(unpacked_input) = &operation.unpacked_input {
render_endpoint_service_object_unpacked_method(output, service, operation, unpacked_input);
return;
}
if operation.deprecated {
output.push_str(
" @typing_extensions.deprecated(\"This operation is deprecated.\", category=None)\n",
);
}
output.push_str(" async def ");
output.push_str(&operation.attr_name);
output.push_str("(\n");
output.push_str(" self,\n");
if let Some(input) = &operation.input {
output.push_str(" request: ");
output.push_str(&input.annotation);
output.push_str(",\n");
}
output.push_str(" ) -> ");
output.push_str(&operation_public_return_annotation(service, operation));
output.push_str(":\n");
render_endpoint_service_object_operation_body(output, service, operation);
}
fn render_endpoint_service_object_unpacked_method(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
unpacked_input: &RenderedUnpackedInput,
) {
if operation.deprecated {
output.push_str(
" @typing_extensions.deprecated(\"This operation is deprecated.\", category=None)\n",
);
}
output.push_str(" async def ");
output.push_str(&operation.attr_name);
output.push_str("(\n");
output.push_str(" self,\n");
output.push_str(" *,\n");
for field in &unpacked_input.parameters {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(&python_parameter_annotation(
&field.annotation,
&field.default_kind,
));
if let Some(default_expr) = python_parameter_default_expr(&field.default_kind) {
render_python_default_expr(output, &default_expr, " ");
}
output.push_str(",\n");
}
output.push_str(" ) -> ");
output.push_str(&operation_public_return_annotation(service, operation));
output.push_str(":\n");
render_unpacked_operation_docstring(output, operation, unpacked_input, None);
for flattened_message in &unpacked_input.flattened_messages {
render_flattened_message_setup(output, flattened_message);
}
output.push_str(" request = ");
output.push_str(&unpacked_input.model_name);
output.push_str("(\n");
for field in &unpacked_input.request_fields {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push('=');
output.push_str(&unpacked_request_field_value_expr(field));
output.push_str(",\n");
}
output.push_str(" )\n");
render_endpoint_service_object_operation_body(output, service, operation);
}
fn render_endpoint_service_object_operation_body(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) {
if operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result
{
if operation.output_direct_result {
output.push_str(" handle = await self._nexus_client.start_operation(\n");
output.push_str(" operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" input=");
output.push_str(operation_input_wire_expr(operation));
output.push_str(",\n");
output.push_str(" output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
output.push_str(" )\n");
if operation.output_module_path.as_deref() == Some("._resources") {
output.push_str(" resource = await handle\n");
if let Some(resource) = service
.resources
.iter()
.find(|resource| resource.type_name == operation.output_annotation)
{
output.push_str(" return ");
output.push_str(&resource_client_bind_function_name(resource));
output.push_str("(resource, self._nexus_client)\n");
return;
}
output.push_str(" return resource\n");
} else {
output.push_str(" return await handle\n");
}
return;
}
output.push_str(" handle = await self._nexus_client.start_operation(\n");
output.push_str(" operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" input=");
output.push_str(operation_input_wire_expr(operation));
output.push_str(",\n");
if !operation.output_none {
output.push_str(" output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
if let Some(transform_expr) = &operation.output_transform_expr {
output.push_str(" result = await handle\n");
output.push_str(" return ");
output.push_str(&temporalio_workflow_value_expr(service, transform_expr));
output.push('\n');
} else if let Some(resource_return) = &operation.output_resource_return {
output.push_str(" result = await handle\n");
output.push_str(" resource = ");
output.push_str(&resource_return.resource_type_name);
output.push_str("(\n");
for binding in &resource_return.bindings {
output.push_str(" ");
output.push_str(&python_field_name(&binding.field_name));
output.push('=');
output.push_str(&resource_return_binding_expr_python(binding));
output.push_str(",\n");
}
output.push_str(" )\n");
if let Some(resource) = service
.resources
.iter()
.find(|resource| resource.type_name == resource_return.resource_type_name)
{
output.push_str(" return ");
output.push_str(&resource_client_bind_function_name(resource));
output.push_str("(resource, self._nexus_client)\n");
} else {
output.push_str(" return resource\n");
}
}
return;
}
output.push_str(" return await self._nexus_client.start_operation(\n");
output.push_str(" operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" input=");
output.push_str(operation_input_wire_expr(operation));
output.push_str(",\n");
if !operation.output_none {
output.push_str(" output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
}
fn endpoint_service_object_name(service: &RenderedService<'_>) -> String {
format!("{}Client", service.name)
}
fn render_service_definition(output: &mut String, service: &RenderedService<'_>) {
if service.deprecated {
output.push_str(
"@typing_extensions.deprecated(\"This service is deprecated.\", category=None)\n",
);
}
if service.wire_name == service.name {
output.push_str("@service\n");
} else {
output.push_str("@service(name=");
output.push_str(&python_string_literal(service.wire_name));
output.push_str(")\n");
}
output.push_str("class ");
output.push_str(service.name);
output.push_str(":\n");
render_python_docstring(
output,
" ",
service.doc.as_deref(),
&[],
None,
service.experimental,
);
if service.operations.is_empty() {
if !service.experimental {
output.push_str(" pass\n");
}
return;
}
for (operation_index, operation) in service.operations.iter().enumerate() {
if operation.experimental {
output.push_str(" # ");
output.push_str(".. warning:: ");
output.push_str(EXPERIMENTAL_WARNING);
output.push('\n');
}
output.push_str(" ");
output.push_str(&operation.attr_name);
if operation.deprecated {
output.push_str(": typing.Annotated[Operation[\n");
} else {
output.push_str(": Operation[\n");
}
output.push_str(" ");
output.push_str(&service_type_ref(
operation
.input
.as_ref()
.map(|input| input.descriptor_type_ref.as_str())
.unwrap_or("None"),
));
output.push_str(",\n");
output.push_str(" ");
output.push_str(&service_type_ref(&operation.descriptor_output_ref));
output.push_str(",\n");
if operation.deprecated {
output.push_str(
" ], typing_extensions.deprecated(\"This operation is deprecated.\", category=None)] = Operation(name=",
);
} else {
output.push_str(" ] = Operation(name=");
}
output.push_str(&python_string_literal(operation.wire_name));
if operation.deprecated {
// `nexusrpc.service` currently ignores `Annotated[Operation[...]]`
// while discovering the generic arguments. Supply the same types on
// the descriptor value so the PEP 702 metadata remains available to
// type checkers without weakening the runtime service definition.
let input_type = operation
.input
.as_ref()
.map(|input| input.descriptor_type_ref.as_str())
.unwrap_or("None");
output.push_str(", input_type=");
output.push_str(&service_operation_runtime_type_ref(input_type));
output.push_str(", output_type=");
output.push_str(&service_operation_runtime_type_ref(
&operation.descriptor_output_ref,
));
}
output.push_str(")\n");
if service.doc.is_some() {
render_python_docstring(output, " ", operation.doc.as_deref(), &[], None, false);
}
if operation_index + 1 != service.operations.len() {
output.push('\n');
}
}
}
fn service_type_ref(type_ref: &str) -> String {
type_ref
.strip_prefix("models.")
.unwrap_or(type_ref)
.to_string()
}
fn service_operation_runtime_type_ref(type_ref: &str) -> String {
let type_ref = service_type_ref(type_ref);
if type_ref == "None" {
"type(None)".to_string()
} else {
type_ref
}
}
fn local_python_model_type_expr(type_ref: &str) -> String {
type_ref
.strip_prefix("models.")
.or_else(|| type_ref.strip_prefix("_recursive."))
.unwrap_or(type_ref)
.to_string()
}
fn model_type_name(type_ref: &str) -> Option<String> {
let type_name = type_ref
.strip_prefix("models.")
.or_else(|| type_ref.strip_prefix("_recursive."))
.unwrap_or(type_ref);
(type_name != "None" && !type_name.contains('.')).then(|| type_name.to_string())
}
fn render_enum(output: &mut String, enumeration: &RenderedEnum) {
output.push_str("class ");
output.push_str(&enumeration.name);
output.push_str("(enum.IntEnum):\n");
if enumeration.values.is_empty() {
output.push_str(" pass\n");
return;
}
for value in &enumeration.values {
output.push_str(" ");
output.push_str(&value.name);
output.push_str(" = ");
output.push_str(&value.number.to_string());
output.push('\n');
}
}
fn render_flags(output: &mut String, flags: &RenderedFlags) {
// Keep native WIT flags default-converter friendly for now. We can return
// to enum.IntFlag once native WIT types have a consistent cross-language
// serialization story.
output.push_str(&flags.name);
output.push_str(": typing.TypeAlias = int\n");
if flags.flags.is_empty() {
return;
}
for flag in &flags.flags {
output.push_str(&flags.name);
output.push_str(&flag.name);
output.push_str(" = 1 << ");
output.push_str(&flag.bit.to_string());
output.push('\n');
}
}
fn render_variant(output: &mut String, variant: &RenderedVariant) {
output.push_str(variant.name.as_str());
output.push_str(" = ");
if variant.cases.is_empty() {
output.push_str("typing.Never\n");
return;
}
if variant.cases.len() > 1 {
output.push_str("(\n ");
}
for (index, case) in variant.cases.iter().enumerate() {
if index > 0 {
output.push_str("\n | ");
}
output.push_str("tuple[typing.Literal[");
output.push_str(&python_string_literal(&case.name));
output.push(']');
if let Some(payload_annotation) = &case.payload_annotation {
output.push_str(", ");
output.push_str(payload_annotation);
}
output.push(']');
}
if variant.cases.len() > 1 {
output.push_str("\n)");
}
output.push('\n');
}
fn render_resource_client_binder(output: &mut String, resource: &PlannedResource) {
output.push_str("def ");
output.push_str(&resource_client_bind_function_name(resource));
output.push_str("(\n");
output.push_str(" resource: ");
output.push_str(&resource.type_name);
output.push_str(",\n");
output.push_str(" nexus_client: typing.Any,\n");
output.push_str(") -> ");
output.push_str(&resource.type_name);
output.push_str(":\n");
output.push_str(" setattr(resource, \"_nexus_client\", nexus_client)\n");
output.push_str(" return resource\n");
}
fn render_resource_method_inline_operation_body(
output: &mut String,
service: &RenderedService<'_>,
method: &PlannedResourceMethod,
operation: &RenderedOperation<'_>,
result_annotation: &str,
) {
let PlannedResourceMethodBindingSpec::Operation {
request_plan,
direct_return,
..
} = &method.binding
else {
return;
};
render_resource_method_request(output, operation, request_plan);
output.push_str(" nexus_client = getattr(self, \"_nexus_client\")\n");
let returns_direct = operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result;
if !returns_direct && *direct_return {
if result_annotation == "None" {
output.push_str(" await ");
} else {
output.push_str(" return await ");
}
render_resource_nexus_start_operation(
output,
operation,
operation_input_wire_expr(operation),
);
return;
}
if returns_direct {
if operation.output_direct_result {
output.push_str(" handle = await ");
render_resource_nexus_start_operation(
output,
operation,
operation_input_wire_expr(operation),
);
if operation.output_module_path.as_deref() == Some("._resources") {
output.push_str(" resource = await handle\n");
if let Some(returned_resource) = service
.resources
.iter()
.find(|candidate| candidate.type_name == operation.output_annotation)
{
output.push_str(" return ");
output.push_str(&resource_client_bind_function_name(returned_resource));
output.push_str("(\n");
output.push_str(" resource,\n");
output.push_str(" nexus_client,\n");
output.push_str(" )\n");
} else {
output.push_str(" return resource\n");
}
} else {
output.push_str(" return await handle\n");
}
return;
}
output.push_str(" handle = await ");
render_resource_nexus_start_operation(
output,
operation,
operation_input_wire_expr(operation),
);
if let Some(transform_expr) = &operation.output_transform_expr {
output.push_str(" result = await handle\n");
output.push_str(" return ");
output.push_str(&temporalio_workflow_value_expr(service, transform_expr));
output.push('\n');
} else if let Some(resource_return) = &operation.output_resource_return {
output.push_str(" result = await handle\n");
output.push_str(" resource = ");
output.push_str(&resource_return.resource_type_name);
output.push_str("(\n");
for binding in &resource_return.bindings {
output.push_str(" ");
output.push_str(&python_field_name(&binding.field_name));
output.push('=');
output.push_str(&resource_return_binding_expr_python(binding));
output.push_str(",\n");
}
output.push_str(" )\n");
if let Some(returned_resource) = service
.resources
.iter()
.find(|candidate| candidate.type_name == resource_return.resource_type_name)
{
output.push_str(" return ");
output.push_str(&resource_client_bind_function_name(returned_resource));
output.push_str("(\n");
output.push_str(" resource,\n");
output.push_str(" nexus_client,\n");
output.push_str(" )\n");
} else {
output.push_str(" return resource\n");
}
}
return;
}
output.push_str(" handle = await ");
render_resource_nexus_start_operation(output, operation, operation_input_wire_expr(operation));
if result_annotation == "None" {
output.push_str(" await handle\n");
} else {
output.push_str(" return await handle\n");
}
}
fn render_resource_method_request(
output: &mut String,
operation: &RenderedOperation<'_>,
request_plan: &RequestPlan,
) {
let request_expr = render_request_plan_python(request_plan);
output.push_str(" request = ");
if let Some(unpacked_input) = &operation.unpacked_input {
output.push_str(&unpacked_input.model_name);
} else {
output.push_str(operation_input_annotation(operation));
}
output.push_str("(\n");
match request_plan {
RequestPlan::Construct { fields, .. } => {
for field in fields {
output.push_str(" ");
output.push_str(&python_field_name(&field.field_name));
output.push('=');
output.push_str(&render_request_plan_python(&field.value));
output.push_str(",\n");
}
}
RequestPlan::Source(_) => {
output.push_str(" value=");
output.push_str(&request_expr);
output.push_str(",\n");
}
}
output.push_str(" )\n");
}
fn render_resource_nexus_start_operation(
output: &mut String,
operation: &RenderedOperation<'_>,
input_expr: &str,
) {
output.push_str("nexus_client.start_operation(\n");
output.push_str(" operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" input=");
output.push_str(input_expr);
output.push_str(",\n");
if !operation.output_none {
output.push_str(" output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
}
fn render_resource_method_operation_body(
output: &mut String,
service: &RenderedService<'_>,
method: &PlannedResourceMethod,
operation: &RenderedOperation<'_>,
result_annotation: &str,
) {
let PlannedResourceMethodBindingSpec::Operation {
request_plan,
direct_return,
..
} = &method.binding
else {
return;
};
let operation_attr_name = request_operation_function_name(service, operation);
render_resource_method_request(output, operation, request_plan);
if *direct_return {
if result_annotation == "None" {
output.push_str(" await ");
output.push_str(&operation_attr_name);
output.push('(');
render_resource_operation_call_args(output, service);
output.push_str("request)\n");
} else {
output.push_str(" return await ");
output.push_str(&operation_attr_name);
output.push('(');
render_resource_operation_call_args(output, service);
output.push_str("request)\n");
}
} else {
output.push_str(" handle = await ");
output.push_str(&operation_attr_name);
output.push('(');
render_resource_operation_call_args(output, service);
output.push_str("request)\n");
if result_annotation == "None" {
output.push_str(" await handle\n");
} else {
output.push_str(" return await handle\n");
}
}
}
fn render_resource_operation_call_args(output: &mut String, service: &RenderedService<'_>) {
if service.endpoint.is_none() {
output.push_str("typing.cast(str, getattr(self, \"_endpoint\")), ");
}
}
fn render_request_plan_python(plan: &RequestPlan) -> String {
render_request_plan(
plan,
python_field_name,
|name, value| format!("{name}={value}"),
|message_name, fields| {
let model_name = message_model_name(message_name);
if fields.is_empty() {
format!("{model_name}()")
} else {
format!("{model_name}({})", fields.join(", "))
}
},
|name| format!("self.{}", python_field_name(name)),
python_field_name,
)
}
fn resource_return_binding_expr_python(binding: &PlannedOperationResourceFieldBinding) -> String {
match &binding.source {
ResolvedResourceBindingSource::RequestField {
field_name,
proto_field_name: _,
hidden,
} => {
if *hidden {
let expr = format!("request.{}", python_field_name(field_name));
if binding.optional {
format!("{expr} or None")
} else {
expr
}
} else {
format!("request.{}", python_field_name(field_name))
}
}
ResolvedResourceBindingSource::ResultField {
proto_field_name, ..
} => {
let expr = format!("result.{proto_field_name}");
if binding.optional {
format!("{expr} or None")
} else {
expr
}
}
}
}
#[derive(Debug, Clone)]
struct RenderedFunctionOverloadCase {
callable_field_name: String,
args_field_name: String,
callable_annotation: String,
positional_args: Vec<RenderedFunctionOverloadPositionalArg>,
args_annotation: Option<String>,
args_optional: bool,
result_type_parameter: Option<String>,
self_type_parameter: Option<String>,
type_parameters: Vec<PythonTypeParameter>,
}
#[derive(Debug, Clone)]
struct RenderedFunctionOverloadPositionalArg {
name: String,
annotation: String,
variadic: bool,
}
#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord)]
enum PythonTypeParameter {
TypeVar(String),
TypeVarTuple(String),
}
fn render_function_unpacked_overloads(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
unpacked_input: &RenderedUnpackedInput,
) {
let overload_combinations = collect_function_overload_combinations(
&unpacked_input.functions,
&operation.output_type_parameters,
);
for (index, overload_cases) in overload_combinations.iter().enumerate() {
render_function_unpacked_overload(
output,
service,
operation,
unpacked_input,
overload_cases,
);
if index + 1 != overload_combinations.len() {
output.push('\n');
}
}
}
fn collect_function_overload_combinations(
functions: &[RenderedFunctionField],
output_type_parameters: &BTreeSet<String>,
) -> Vec<Vec<RenderedFunctionOverloadCase>> {
let per_function_cases = functions
.iter()
.map(|function| function_overload_cases(function, output_type_parameters))
.collect::<Vec<_>>();
let mut combinations = Vec::new();
let mut current = Vec::new();
collect_function_overload_combinations_recursive(
&per_function_cases,
0,
&mut current,
&mut combinations,
);
combinations
}
fn collect_function_overload_combinations_recursive(
per_function_cases: &[Vec<RenderedFunctionOverloadCase>],
index: usize,
current: &mut Vec<RenderedFunctionOverloadCase>,
combinations: &mut Vec<Vec<RenderedFunctionOverloadCase>>,
) {
if index == per_function_cases.len() {
combinations.push(current.clone());
return;
}
for case in &per_function_cases[index] {
current.push(case.clone());
collect_function_overload_combinations_recursive(
per_function_cases,
index + 1,
current,
combinations,
);
current.pop();
}
}
fn function_overload_cases(
function: &RenderedFunctionField,
output_type_parameters: &BTreeSet<String>,
) -> Vec<RenderedFunctionOverloadCase> {
let result_type_parameter = function
.result_type_parameter
.clone()
.filter(|name| output_type_parameters.contains(name));
let result_annotation = if result_type_parameter.is_some() {
&function.result_annotation
} else {
&function.erased_result_annotation
};
let self_type_parameter = python_function_self_type_parameter(function);
let mut cases = Vec::new();
if let Some(alternate_annotation) = &function.alternate_annotation {
if function.primary && matches!(function.args, RenderedFunctionArgs::Varargs { .. }) {
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: alternate_annotation.clone(),
positional_args: python_alternate_function_positional_args(function),
args_annotation: None,
args_optional: false,
result_type_parameter: None,
self_type_parameter: None,
type_parameters: Vec::new(),
});
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: alternate_annotation.clone(),
positional_args: Vec::new(),
args_annotation: python_alternate_function_args_annotation(function),
args_optional: true,
result_type_parameter: None,
self_type_parameter: None,
type_parameters: Vec::new(),
});
} else {
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: alternate_annotation.clone(),
positional_args: python_alternate_function_positional_args(function),
args_annotation: python_alternate_function_args_annotation(function),
args_optional: true,
result_type_parameter: None,
self_type_parameter: None,
type_parameters: Vec::new(),
});
}
}
match &function.args {
RenderedFunctionArgs::Varargs { prefix } if function.primary => {
let type_parameter_prefix = function_overload_parameter_prefix(function);
let function_args = format!("{type_parameter_prefix}Args");
let callable_args = python_varargs_callable_arg_list(
prefix,
&function_args,
self_type_parameter.as_ref(),
);
let mut type_parameters = Vec::new();
if let Some(result_type_parameter) = &result_type_parameter {
type_parameters.push(PythonTypeParameter::TypeVar(result_type_parameter.clone()));
}
if let Some(self_type_parameter) = &self_type_parameter {
type_parameters.push(PythonTypeParameter::TypeVar(self_type_parameter.clone()));
}
type_parameters.push(PythonTypeParameter::TypeVarTuple(function_args.clone()));
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: format!(
"collections.abc.Callable[[{}], {}]",
callable_args, result_annotation
),
positional_args: vec![RenderedFunctionOverloadPositionalArg {
name: "positional_args".to_string(),
annotation: format!("typing_extensions.Unpack[{function_args}]"),
variadic: true,
}],
args_annotation: None,
args_optional: false,
result_type_parameter: result_type_parameter.clone(),
self_type_parameter: self_type_parameter.clone(),
type_parameters,
});
let type_parameters = result_type_parameter
.iter()
.map(|name| PythonTypeParameter::TypeVar(name.clone()))
.chain(self_type_parameter.iter().filter_map(|name| {
if output_type_parameters.is_empty() {
None
} else {
Some(PythonTypeParameter::TypeVar(name.clone()))
}
}))
.collect();
let (callable_annotation, case_self_type_parameter) =
if self_type_parameter.is_some() && !output_type_parameters.is_empty() {
(
format!(
"collections.abc.Callable[[{}], {}]",
callable_args, result_annotation
),
self_type_parameter.clone(),
)
} else {
(
format!("collections.abc.Callable[..., {}]", result_annotation),
None,
)
};
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation,
positional_args: Vec::new(),
args_annotation: Some("list[typing.Any]".to_string()),
args_optional: false,
result_type_parameter: result_type_parameter.clone(),
self_type_parameter: case_self_type_parameter,
type_parameters,
});
}
RenderedFunctionArgs::Varargs { prefix } => {
let type_parameter_prefix = function_overload_parameter_prefix(function);
let first_arg = format!("{type_parameter_prefix}Arg");
let callable_prefix = prefix
.iter()
.enumerate()
.map(|(index, parameter)| {
self_type_parameter
.as_ref()
.filter(|_| index == 0)
.cloned()
.unwrap_or_else(|| parameter.annotation.clone())
})
.collect::<Vec<_>>();
let no_arg_callable_args = python_callable_arg_list(&callable_prefix);
let mut one_arg_callable_prefix = callable_prefix;
one_arg_callable_prefix.push(first_arg.clone());
let one_arg_callable_args = python_callable_arg_list(&one_arg_callable_prefix);
let type_parameters = result_type_parameter
.iter()
.map(|name| PythonTypeParameter::TypeVar(name.clone()))
.chain(
self_type_parameter
.iter()
.map(|name| PythonTypeParameter::TypeVar(name.clone())),
)
.collect();
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: format!(
"collections.abc.Callable[[{}], {}]",
no_arg_callable_args, result_annotation
),
positional_args: Vec::new(),
args_annotation: None,
args_optional: false,
result_type_parameter: result_type_parameter.clone(),
self_type_parameter: self_type_parameter.clone(),
type_parameters,
});
let mut type_parameters = Vec::new();
if let Some(result_type_parameter) = &result_type_parameter {
type_parameters.push(PythonTypeParameter::TypeVar(result_type_parameter.clone()));
}
if let Some(self_type_parameter) = &self_type_parameter {
type_parameters.push(PythonTypeParameter::TypeVar(self_type_parameter.clone()));
}
type_parameters.push(PythonTypeParameter::TypeVar(first_arg.clone()));
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: format!(
"collections.abc.Callable[[{}], {}]",
one_arg_callable_args, result_annotation
),
positional_args: Vec::new(),
args_annotation: Some(first_arg.clone()),
args_optional: false,
result_type_parameter: result_type_parameter.clone(),
self_type_parameter: self_type_parameter.clone(),
type_parameters,
});
let type_parameters = result_type_parameter
.iter()
.map(|name| PythonTypeParameter::TypeVar(name.clone()))
.collect();
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: format!(
"collections.abc.Callable[..., {}]",
result_annotation
),
positional_args: Vec::new(),
args_annotation: Some("list[typing.Any]".to_string()),
args_optional: false,
result_type_parameter: result_type_parameter.clone(),
self_type_parameter: None,
type_parameters,
});
}
RenderedFunctionArgs::Typed { parameters } => {
let args_annotation = python_function_args_list_annotation(
parameters
.iter()
.map(|parameter| parameter.annotation.as_str())
.collect::<Vec<_>>()
.as_slice(),
);
let type_parameters = result_type_parameter
.iter()
.map(|name| PythonTypeParameter::TypeVar(name.clone()))
.collect();
cases.push(RenderedFunctionOverloadCase {
callable_field_name: function.callable_field_name.clone(),
args_field_name: function.args_field_name.clone(),
callable_annotation: python_rendered_function_callable_annotation_with_result(
function,
result_annotation,
),
positional_args: if function.primary {
parameters
.iter()
.map(|parameter| RenderedFunctionOverloadPositionalArg {
name: parameter.name.clone(),
annotation: parameter.annotation.clone(),
variadic: false,
})
.collect()
} else {
Vec::new()
},
args_annotation: if function.primary {
None
} else if parameters.len() == 1 {
Some(parameters[0].annotation.clone())
} else {
Some(args_annotation.clone())
},
args_optional: false,
result_type_parameter: result_type_parameter.clone(),
self_type_parameter: None,
type_parameters,
});
}
}
cases
}
fn python_alternate_function_positional_args(
function: &RenderedFunctionField,
) -> Vec<RenderedFunctionOverloadPositionalArg> {
if !function.primary {
return Vec::new();
}
match &function.args {
RenderedFunctionArgs::Typed { parameters } => parameters
.iter()
.map(|parameter| RenderedFunctionOverloadPositionalArg {
name: parameter.name.clone(),
annotation: parameter.annotation.clone(),
variadic: false,
})
.collect(),
RenderedFunctionArgs::Varargs { .. } => {
vec![RenderedFunctionOverloadPositionalArg {
name: "positional_args".to_string(),
annotation: "object".to_string(),
variadic: true,
}]
}
}
}
fn python_alternate_function_args_annotation(function: &RenderedFunctionField) -> Option<String> {
match (&function.args, function.primary) {
(RenderedFunctionArgs::Typed { .. }, true) => None,
_ => Some("list[typing.Any] | None".to_string()),
}
}
fn python_rendered_function_callable_annotation_with_result(
function: &RenderedFunctionField,
result_annotation: &str,
) -> String {
match &function.args {
RenderedFunctionArgs::Varargs { .. } => {
format!("collections.abc.Callable[..., {}]", result_annotation)
}
RenderedFunctionArgs::Typed { parameters } => format!(
"collections.abc.Callable[[{}], {}]",
parameters
.iter()
.map(|parameter| parameter.annotation.clone())
.collect::<Vec<_>>()
.join(", "),
result_annotation
),
}
}
fn python_function_self_type_parameter(function: &RenderedFunctionField) -> Option<String> {
match &function.args {
RenderedFunctionArgs::Varargs { prefix } if !prefix.is_empty() => {
Some("SelfType".to_string())
}
_ => None,
}
}
fn python_varargs_callable_arg_list(
prefix: &[RenderedFunctionArg],
varargs_name: &str,
self_type_parameter: Option<&String>,
) -> String {
let mut args = prefix
.iter()
.enumerate()
.map(|(index, parameter)| {
self_type_parameter
.filter(|_| index == 0)
.cloned()
.unwrap_or_else(|| parameter.annotation.clone())
})
.collect::<Vec<_>>();
args.push(format!("typing_extensions.Unpack[{varargs_name}]"));
args.join(", ")
}
fn python_callable_arg_list(args: &[String]) -> String {
if args.is_empty() {
String::new()
} else {
args.join(", ")
}
}
fn function_overload_parameter_prefix(function: &RenderedFunctionField) -> String {
function.callable_field_name.to_upper_camel_case()
}
fn primary_function(functions: &[RenderedFunctionField]) -> Option<&RenderedFunctionField> {
functions.iter().find(|function| function.primary)
}
fn python_function_positional_args_annotation(function: &RenderedFunctionField) -> String {
match &function.args {
RenderedFunctionArgs::Typed { parameters } => python_function_args_list_annotation(
parameters
.iter()
.map(|parameter| parameter.annotation.as_str())
.collect::<Vec<_>>()
.as_slice(),
)
.trim_start_matches("list[")
.trim_end_matches(']')
.to_string(),
RenderedFunctionArgs::Varargs { .. } => "object".to_string(),
}
}
fn python_function_fixed_positional_args(
function: &RenderedFunctionField,
) -> Option<&[RenderedFunctionArg]> {
match &function.args {
RenderedFunctionArgs::Typed { parameters } if function.primary => Some(parameters),
_ => None,
}
}
fn python_function_args_parameter_annotation(
function: &RenderedFunctionField,
primary_args: bool,
) -> String {
match (&function.args, primary_args) {
(RenderedFunctionArgs::Typed { parameters }, true) => {
let annotations = parameters
.iter()
.map(|parameter| parameter.annotation.as_str())
.collect::<Vec<_>>();
format!(
"{} | None",
python_function_args_list_annotation(annotations.as_slice())
)
}
(RenderedFunctionArgs::Typed { parameters }, false) => {
if parameters.len() == 1 {
format!("{} | None", parameters[0].annotation)
} else {
let annotations = parameters
.iter()
.map(|parameter| parameter.annotation.as_str())
.collect::<Vec<_>>();
format!(
"{} | None",
python_function_args_list_annotation(annotations.as_slice())
)
}
}
(RenderedFunctionArgs::Varargs { .. }, true) => "list[typing.Any] | None".to_string(),
(RenderedFunctionArgs::Varargs { .. }, false) => {
"object | list[typing.Any] | None".to_string()
}
}
}
fn python_function_normalized_args_annotation(function: &RenderedFunctionField) -> String {
match &function.args {
RenderedFunctionArgs::Typed { parameters } => {
let annotations = parameters
.iter()
.map(|parameter| parameter.annotation.as_str())
.collect::<Vec<_>>();
python_function_args_list_annotation(annotations.as_slice())
}
RenderedFunctionArgs::Varargs { .. } => "list[typing.Any]".to_string(),
}
}
fn overload_case_for_callable<'a>(
overload_cases: &'a [RenderedFunctionOverloadCase],
callable_field_name: &str,
) -> Option<&'a RenderedFunctionOverloadCase> {
overload_cases
.iter()
.find(|function_case| function_case.callable_field_name == callable_field_name)
}
fn render_function_unpacked_overload(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
unpacked_input: &RenderedUnpackedInput,
overload_cases: &[RenderedFunctionOverloadCase],
) {
render_function_unpacked_overload_case_doc(output, unpacked_input, overload_cases);
output.push_str("@typing.overload\n");
if operation.deprecated {
output.push_str(
"@typing_extensions.deprecated(\"This operation is deprecated.\", category=None)\n",
);
}
output.push_str("async def ");
output.push_str(&operation.attr_name);
output.push_str("(\n");
render_endpoint_parameter(output, service);
let primary_function = primary_function(&unpacked_input.functions);
let primary_case = primary_function.and_then(|function| {
overload_case_for_callable(overload_cases, &function.callable_field_name)
});
if let (Some(function), Some(function_case)) = (primary_function, primary_case) {
output.push_str(" ");
output.push_str(&function.callable_field_name);
output.push_str(": ");
output.push_str(&function_case.callable_annotation);
output.push_str(",\n");
for positional_arg in &function_case.positional_args {
output.push_str(" ");
if positional_arg.variadic {
output.push('*');
}
output.push_str(&function_overload_positional_arg_render_name(
positional_arg,
unpacked_input,
primary_function,
primary_case,
overload_cases,
));
output.push_str(": ");
output.push_str(&positional_arg.annotation);
output.push_str(",\n");
}
}
if !primary_case.is_some_and(|function_case| !function_case.positional_args.is_empty()) {
output.push_str(" *,\n");
}
for field in &unpacked_input.parameters {
if primary_function.is_some_and(|function| function.callable_field_name == field.attr_name)
{
continue;
}
if primary_function
.and_then(python_function_fixed_positional_args)
.is_some_and(|parameters| {
parameters
.iter()
.any(|parameter| parameter.name == field.attr_name)
})
{
continue;
}
if primary_function.is_some_and(|function| function.args_field_name == field.attr_name) {
if let Some(function_case) = primary_case {
if let Some(args_annotation) = &function_case.args_annotation {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(args_annotation);
if function_case.args_optional {
output.push_str(" = ...");
}
output.push_str(",\n");
}
}
continue;
}
if let Some(function_case) = overload_cases
.iter()
.find(|function_case| function_case.callable_field_name == field.attr_name)
{
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(&function_case.callable_annotation);
if field.default_kind != PythonFieldDefaultKind::Required {
output.push_str(" = ...");
}
output.push_str(",\n");
continue;
}
if let Some(function_case) = overload_cases
.iter()
.find(|function_case| function_case.args_field_name == field.attr_name)
{
if let Some(args_annotation) = &function_case.args_annotation {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(args_annotation);
if function_case.args_optional {
output.push_str(" = ...");
}
output.push_str(",\n");
}
continue;
}
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(&python_parameter_annotation(
&field.annotation,
&field.default_kind,
));
if field.default_kind != PythonFieldDefaultKind::Required {
output.push_str(" = ...");
}
output.push_str(",\n");
}
output.push_str(") -> ");
output.push_str(&function_unpacked_overload_return_annotation(
service,
operation,
overload_cases,
));
output.push_str(": ...\n");
}
fn function_overload_positional_arg_render_name(
positional_arg: &RenderedFunctionOverloadPositionalArg,
unpacked_input: &RenderedUnpackedInput,
primary_function: Option<&RenderedFunctionField>,
primary_case: Option<&RenderedFunctionOverloadCase>,
overload_cases: &[RenderedFunctionOverloadCase],
) -> String {
if positional_arg.variadic
&& positional_arg.name == "positional_args"
&& !function_unpacked_overload_renders_parameter_name(
unpacked_input,
primary_function,
primary_case,
overload_cases,
"args",
)
{
"args".to_string()
} else {
positional_arg.name.clone()
}
}
fn function_unpacked_overload_renders_parameter_name(
unpacked_input: &RenderedUnpackedInput,
primary_function: Option<&RenderedFunctionField>,
primary_case: Option<&RenderedFunctionOverloadCase>,
overload_cases: &[RenderedFunctionOverloadCase],
name: &str,
) -> bool {
if let Some(function) = primary_function {
if function.callable_field_name == name {
return true;
}
if primary_case.is_some_and(|case| {
case.positional_args.iter().any(|arg| arg.name == name)
|| (case.args_annotation.is_some() && case.args_field_name == name)
}) {
return true;
}
}
for field in &unpacked_input.parameters {
if primary_function.is_some_and(|function| function.callable_field_name == field.attr_name)
{
continue;
}
if primary_function
.and_then(python_function_fixed_positional_args)
.is_some_and(|parameters| {
parameters
.iter()
.any(|parameter| parameter.name == field.attr_name)
})
{
continue;
}
if primary_function.is_some_and(|function| function.args_field_name == field.attr_name) {
if primary_case.is_some_and(|case| case.args_annotation.is_some())
&& field.attr_name == name
{
return true;
}
continue;
}
if overload_cases
.iter()
.any(|case| case.callable_field_name == field.attr_name)
{
if field.attr_name == name {
return true;
}
continue;
}
if let Some(function_case) = overload_cases
.iter()
.find(|function_case| function_case.args_field_name == field.attr_name)
{
if function_case.args_annotation.is_some() && field.attr_name == name {
return true;
}
continue;
}
if field.attr_name == name {
return true;
}
}
false
}
fn render_function_unpacked_overload_case_doc(
output: &mut String,
unpacked_input: &RenderedUnpackedInput,
overload_cases: &[RenderedFunctionOverloadCase],
) {
output.push_str("# Overload case:\n");
for overload_case in ordered_function_overload_cases(unpacked_input, overload_cases) {
output.push_str("# - ");
output.push_str(&function_overload_case_doc(&overload_case));
output.push('\n');
}
}
fn ordered_function_overload_cases<'a>(
unpacked_input: &RenderedUnpackedInput,
overload_cases: &'a [RenderedFunctionOverloadCase],
) -> Vec<&'a RenderedFunctionOverloadCase> {
let mut ordered = Vec::new();
let mut rendered = BTreeSet::new();
for field in &unpacked_input.parameters {
if let Some(overload_case) = overload_cases
.iter()
.find(|overload_case| overload_case.callable_field_name == field.attr_name)
{
rendered.insert(overload_case.callable_field_name.clone());
ordered.push(overload_case);
}
}
for overload_case in overload_cases {
if !rendered.contains(&overload_case.callable_field_name) {
ordered.push(overload_case);
}
}
ordered
}
fn function_overload_case_doc(overload_case: &RenderedFunctionOverloadCase) -> String {
let callable_label = python_readable_name(&overload_case.callable_field_name);
let callable_kind = if overload_case
.callable_annotation
.starts_with("collections.abc.Callable")
{
if overload_case.self_type_parameter.is_some() {
"method callable"
} else {
"callable"
}
} else {
"name"
};
let args_label = python_function_args_doc_label(
&overload_case.callable_field_name,
&overload_case.args_field_name,
);
format!(
"{callable_label} {callable_kind} with {}",
function_overload_args_doc(overload_case, &args_label)
)
}
fn python_readable_name(name: &str) -> String {
name.replace('_', " ")
}
fn python_function_args_doc_label(callable_field_name: &str, args_field_name: &str) -> String {
if args_field_name == "args" {
return format!("{} arguments", python_readable_name(callable_field_name));
}
if let Some(prefix) = args_field_name.strip_suffix("_args") {
return format!("{} arguments", python_readable_name(prefix));
}
python_readable_name(args_field_name)
}
fn function_overload_args_doc(
overload_case: &RenderedFunctionOverloadCase,
args_label: &str,
) -> String {
if !overload_case.positional_args.is_empty() {
if overload_case
.callable_annotation
.starts_with("collections.abc.Callable")
&& overload_case
.positional_args
.iter()
.any(|arg| arg.annotation != "object")
{
return format!("typed positional {args_label}");
}
return format!("positional {args_label}");
}
let Some(args_annotation) = &overload_case.args_annotation else {
return format!("no {args_label}");
};
if args_annotation.contains("list[") {
if overload_case.args_optional {
format!("optional list-form {args_label}")
} else {
format!("list-form {args_label}")
}
} else {
format!("a typed single {args_label}")
}
}
fn function_unpacked_overload_return_annotation(
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
overload_cases: &[RenderedFunctionOverloadCase],
) -> String {
let mut active_type_parameters = overload_cases
.iter()
.filter_map(|function_case| function_case.result_type_parameter.as_ref())
.cloned()
.collect::<BTreeSet<_>>();
let self_type_parameter = overload_cases
.iter()
.find_map(|function_case| function_case.self_type_parameter.as_ref());
if self_type_parameter.is_some() {
active_type_parameters.extend(operation.output_type_parameters.iter().cloned());
}
let mut output_annotation = erase_inactive_python_type_parameters(
&operation.overload_output_annotation,
&operation.output_type_parameters,
&active_type_parameters,
);
if let Some(self_type_parameter) = self_type_parameter {
for output_type_parameter in &operation.output_type_parameters {
output_annotation = replace_python_identifier(
&output_annotation,
output_type_parameter,
self_type_parameter,
);
}
}
let output_annotation = temporalio_workflow_type_annotation(service, &output_annotation);
if operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result
{
output_annotation
} else {
format!(
"{}[{output_annotation}]",
temporalio_workflow_type_ref(service, "NexusOperationHandle")
)
}
}
fn erase_inactive_python_type_parameters(
annotation: &str,
type_parameters: &BTreeSet<String>,
active_type_parameters: &BTreeSet<String>,
) -> String {
let inactive_type_parameters = type_parameters
.difference(active_type_parameters)
.cloned()
.collect::<BTreeSet<_>>();
erase_python_type_parameters(annotation, &inactive_type_parameters)
}
fn erase_python_type_parameters(annotation: &str, type_parameters: &BTreeSet<String>) -> String {
erase_python_type_parameters_with_replacement(annotation, type_parameters, "object")
}
fn erase_python_type_parameters_with_replacement(
annotation: &str,
type_parameters: &BTreeSet<String>,
replacement: &str,
) -> String {
type_parameters
.iter()
.fold(annotation.to_string(), |current, name| {
replace_python_identifier(¤t, name, replacement)
})
}
fn python_annotation_uses_identifier(annotation: &str, identifier: &str) -> bool {
annotation != replace_python_identifier(annotation, identifier, "typing.Any")
}
fn replace_python_identifier(source: &str, identifier: &str, replacement: &str) -> String {
if identifier.is_empty() {
return source.to_string();
}
let mut output = String::with_capacity(source.len());
let mut index = 0;
while let Some(relative_match_start) = source[index..].find(identifier) {
let match_start = index + relative_match_start;
let match_end = match_start + identifier.len();
let before = source[..match_start].chars().next_back();
let after = source[match_end..].chars().next();
if before.is_some_and(|character| is_python_identifier_char(character) || character == '.')
|| after.is_some_and(is_python_identifier_char)
{
output.push_str(&source[index..match_end]);
} else {
output.push_str(&source[index..match_start]);
output.push_str(replacement);
}
index = match_end;
}
output.push_str(&source[index..]);
output
}
fn render_operation_functions(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) {
render_request_only_operation_function(output, service, operation);
if service.endpoint.is_none() {
return;
}
if let Some(unpacked_input) = &operation.unpacked_input {
output.push('\n');
if function_unpacked_needs_overloads(&unpacked_input.functions) {
render_function_unpacked_overloads(output, service, operation, unpacked_input);
output.push('\n');
}
render_unpacked_operation_function(output, service, operation, unpacked_input);
}
}
fn function_unpacked_needs_overloads(functions: &[RenderedFunctionField]) -> bool {
functions.iter().any(|function| {
function.alternate_annotation.is_some()
|| matches!(function.args, RenderedFunctionArgs::Varargs { .. })
})
}
fn request_operation_function_name(
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) -> String {
if service.endpoint.is_none() {
format!("{}_request", operation.attr_name)
} else if operation.unpacked_input.is_some() {
format!("_{}", operation.attr_name)
} else {
operation.attr_name.clone()
}
}
fn operation_input_type_ref<'a>(operation: &'a RenderedOperation<'_>) -> &'a str {
operation
.input
.as_ref()
.map(|input| input.type_ref.as_str())
.unwrap_or("None")
}
fn operation_input_annotation<'a>(operation: &'a RenderedOperation<'_>) -> &'a str {
operation
.input
.as_ref()
.map(|input| input.annotation.as_str())
.unwrap_or("None")
}
fn operation_input_wire_expr<'a>(operation: &'a RenderedOperation<'_>) -> &'a str {
operation
.input
.as_ref()
.map(|_| "request")
.unwrap_or("None")
}
fn planned_record_for_external_source<'a>(
api_plan: &'a PlannedSpec,
external: &ExternalTypeSpec<PlannedFamily>,
) -> Option<&'a RecordSpec<PlannedFamily>> {
let ExternalTypeSpec::Proto(external_proto) = external else {
return None;
};
api_plan.records().map(|(_, record)| record).find(|record| {
record
.source
.as_ref()
.and_then(ExternalTypeSourceSpec::proto_type)
.is_some_and(|source_proto| source_proto == external_proto)
})
}
fn planned_record_type(record: &RecordSpec<PlannedFamily>) -> PlannedType {
PlannedType::Record(PlannedRecordType {
full_name: record.full_name.clone(),
model_name: record.name.clone(),
})
}
fn render_endpoint_parameter(output: &mut String, service: &RenderedService<'_>) {
if service.endpoint.is_none() {
output.push_str(" endpoint: str,\n");
}
}
fn endpoint_forward_args(service: &RenderedService<'_>) -> &'static str {
if service.endpoint.is_none() {
"endpoint, request"
} else {
"request"
}
}
fn render_request_only_operation_function(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) {
let function_name = request_operation_function_name(service, operation);
let render_request_doc = operation.unpacked_input.is_none();
if operation.deprecated {
output.push_str(
"@typing_extensions.deprecated(\"This operation is deprecated.\", category=None)\n",
);
}
output.push_str("async def ");
output.push_str(&function_name);
output.push_str("(\n");
render_endpoint_parameter(output, service);
if let Some(input) = &operation.input {
output.push_str(" request: ");
output.push_str(&input.annotation);
output.push_str(",\n");
}
if operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result
{
output.push_str(") -> ");
output.push_str(&function_unpacked_implementation_return_annotation(
service, operation,
));
output.push_str(":\n");
} else {
output.push_str(") -> ");
output.push_str(&temporalio_workflow_type_ref(
service,
"NexusOperationHandle",
));
output.push_str("[\n");
output.push_str(" ");
output.push_str(&temporalio_workflow_type_annotation(
service,
&operation.output_annotation,
));
output.push_str(",\n");
output.push_str("]:\n");
if render_request_doc {
render_operation_docstring(output, operation, &[], None);
}
render_temporalio_workflow_runtime_imports(
output,
service,
&temporalio_workflow_operation_runtime_names(service, operation),
" ",
);
render_inline_nexus_client(output, service, " ");
if operation.generic_model_types {
output.push_str(" return typing.cast(\n");
output.push_str(" ");
output.push_str(&temporalio_workflow_type_ref(
service,
"NexusOperationHandle",
));
output.push_str("[\n");
output.push_str(" ");
output.push_str(&temporalio_workflow_type_annotation(
service,
&operation.output_annotation,
));
output.push_str(",\n");
output.push_str(" ],\n");
output.push_str(" await nexus_client.start_operation(\n");
} else {
output.push_str(" return await nexus_client.start_operation(\n");
}
output.push_str(" operation=");
if operation.generic_model_types {
output.push_str(" ");
}
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" ");
if operation.generic_model_types {
output.push_str(" ");
}
output.push_str("input=");
output.push_str(operation_input_wire_expr(operation));
output.push_str(",\n");
if !operation.output_none {
output.push_str(" ");
if operation.generic_model_types {
output.push_str(" ");
}
output.push_str("output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
}
if operation.generic_model_types {
output.push_str(" ),\n");
output.push_str(" )\n");
} else {
output.push_str(" )\n");
}
return;
}
if render_request_doc {
render_operation_docstring(output, operation, &[], None);
}
render_temporalio_workflow_runtime_imports(
output,
service,
&temporalio_workflow_operation_runtime_names(service, operation),
" ",
);
if operation.output_direct_result {
render_inline_nexus_client(output, service, " ");
output.push_str(" handle = await nexus_client.start_operation(\n");
output.push_str(" operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" ");
output.push_str("input=");
output.push_str(operation_input_wire_expr(operation));
output.push_str(",\n");
output.push_str(" output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
output.push_str(" )\n");
if service.endpoint.is_none()
&& operation.output_module_path.as_deref() == Some("._resources")
{
output.push_str(" resource = await handle\n");
output.push_str(" setattr(resource, \"_endpoint\", endpoint)\n");
output.push_str(" return resource\n");
} else {
output.push_str(" return await handle\n");
}
return;
}
render_inline_nexus_client(output, service, " ");
output.push_str(" handle = await nexus_client.start_operation(\n");
output.push_str(" operation=");
output.push_str(&python_string_literal(operation.wire_name));
output.push_str(",\n");
output.push_str(" input=");
output.push_str(operation_input_wire_expr(operation));
output.push_str(",\n");
if !operation.output_none {
output.push_str(" output_type=");
output.push_str(&operation.output_type_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
if let Some(transform_expr) = &operation.output_transform_expr {
output.push_str(" result = await handle\n");
output.push_str(" return ");
output.push_str(&temporalio_workflow_value_expr(service, transform_expr));
output.push('\n');
} else if let Some(resource_return) = &operation.output_resource_return {
output.push_str(" result = await handle\n");
if service.endpoint.is_none() {
output.push_str(" resource = ");
} else {
output.push_str(" return ");
}
output.push_str(&resource_return.resource_type_name);
output.push_str("(\n");
for binding in &resource_return.bindings {
output.push_str(" ");
output.push_str(&python_field_name(&binding.field_name));
output.push('=');
output.push_str(&resource_return_binding_expr_python(binding));
output.push_str(",\n");
}
output.push_str(" )\n");
if service.endpoint.is_none() {
output.push_str(" setattr(resource, \"_endpoint\", endpoint)\n");
output.push_str(" return resource\n");
}
}
}
fn render_unpacked_operation_docstring(
output: &mut String,
operation: &RenderedOperation<'_>,
unpacked_input: &RenderedUnpackedInput,
primary_function: Option<&RenderedFunctionField>,
) {
let mut args = Vec::<(String, String)>::new();
for field in &unpacked_input.parameters {
if let Some(function) = primary_function
&& field.attr_name == function.callable_field_name
{
if let Some(doc) = &field.doc {
args.push((field.attr_name.clone(), doc.clone()));
}
if let Some(parameters) = python_function_fixed_positional_args(function) {
for parameter in parameters {
args.push((
parameter.name.clone(),
format!("Argument for {}.", function.callable_field_name),
));
}
} else {
args.push((
"positional_args".to_string(),
format!(
"Positional arguments for {}. Cannot be set if args is set.",
function.callable_field_name
),
));
}
continue;
}
if let Some(function) = unpacked_input
.functions
.iter()
.find(|function| function.args_field_name == field.attr_name)
{
let generated_doc = if function.primary {
match &function.args {
RenderedFunctionArgs::Varargs { .. } => format!(
"List-form arguments for {}. Cannot be set if positional_args are set. \
For typed {} callables, list contents are not statically typechecked; \
pass {} arguments positionally for precise typechecking.",
function.callable_field_name,
function.callable_field_name,
function.callable_field_name
),
_ => format!(
"List-form arguments for {}. Cannot be set if positional_args are set.",
function.callable_field_name
),
}
} else {
match &function.args {
RenderedFunctionArgs::Varargs { .. } => format!(
"Argument value, or list of argument values, for {}. For typed \
single-argument signals, scalar signal_args values are statically \
typechecked. List-form signal_args values are not precisely typechecked. \
To pass a single signal argument that is itself a list, wrap it in \
another list; otherwise the list is interpreted as multiple signal \
arguments.",
function.callable_field_name
),
_ => format!(
"Argument value, or list of argument values, for {}. To pass a single \
argument that is itself a list, wrap it in another list; otherwise the \
list is interpreted as multiple arguments.",
function.callable_field_name
),
}
};
args.push((field.attr_name.clone(), generated_doc));
continue;
}
if let Some(doc) = &field.doc
&& !python_field_doc_rendered_as_function_arg(unpacked_input, &field.attr_name)
{
args.push((field.attr_name.clone(), doc.clone()));
}
}
render_operation_docstring(output, operation, &args, operation.return_doc.as_deref());
}
fn python_field_doc_rendered_as_function_arg(
unpacked_input: &RenderedUnpackedInput,
field_name: &str,
) -> bool {
unpacked_input.functions.iter().any(|function| {
python_function_fixed_positional_args(function).is_some_and(|parameters| {
parameters
.iter()
.any(|parameter| parameter.name == field_name)
})
})
}
fn render_operation_docstring(
output: &mut String,
operation: &RenderedOperation<'_>,
args: &[(String, String)],
returns: Option<&str>,
) {
render_python_docstring(
output,
" ",
operation.doc.as_deref(),
args,
returns,
operation.experimental,
);
}
pub(in crate::generator) fn render_python_docstring(
output: &mut String,
indent: &str,
summary: Option<&str>,
args: &[(String, String)],
returns: Option<&str>,
experimental: bool,
) {
let mut lines = Vec::<String>::new();
let docstring_width = PYTHON_FORMAT_LINE_LENGTH.saturating_sub(indent.chars().count());
let has_summary = summary.is_some_and(|summary| !summary.trim().is_empty());
if let Some(summary) = summary.map(str::trim).filter(|summary| !summary.is_empty()) {
for line in summary.lines() {
push_wrapped_python_docstring_line(&mut lines, "", "", line.trim(), docstring_width);
}
}
if experimental {
if !lines.is_empty() {
lines.push(String::new());
}
lines.push(".. warning::".to_string());
push_wrapped_python_docstring_line(
&mut lines,
" ",
" ",
EXPERIMENTAL_WARNING,
docstring_width,
);
}
if !args.is_empty() {
if !lines.is_empty() {
lines.push(String::new());
}
lines.push("Args:".to_string());
for (name, doc) in args {
let mut doc_lines = doc.trim().lines();
let first_prefix = format!(" {name}: ");
let continuation_prefix = " ";
let first = doc_lines.next().unwrap_or_default().trim();
push_wrapped_python_docstring_line(
&mut lines,
&first_prefix,
continuation_prefix,
first,
docstring_width,
);
for line in doc_lines {
push_wrapped_python_docstring_line(
&mut lines,
continuation_prefix,
continuation_prefix,
line.trim(),
docstring_width,
);
}
}
}
if let Some(returns) = returns.map(str::trim).filter(|returns| !returns.is_empty()) {
if !lines.is_empty() {
lines.push(String::new());
}
lines.push("Returns:".to_string());
let mut return_lines = returns.lines();
if let Some(first) = return_lines.next() {
push_wrapped_python_docstring_line(
&mut lines,
" ",
" ",
first.trim(),
docstring_width,
);
}
for line in return_lines {
push_wrapped_python_docstring_line(
&mut lines,
" ",
" ",
line.trim(),
docstring_width,
);
}
}
if lines.is_empty() {
return;
}
output.push_str(indent);
output.push_str("\"\"\"");
if !has_summary {
output.push('\n');
for line in &lines {
if !line.is_empty() {
output.push_str(indent);
output.push_str(&python_docstring_literal_text(line));
}
output.push('\n');
}
output.push_str(indent);
output.push_str("\"\"\"\n");
return;
}
if lines.len() == 1 {
output.push_str(&python_docstring_literal_text(&lines[0]));
output.push_str("\"\"\"\n");
return;
}
output.push_str(&python_docstring_literal_text(&lines[0]));
output.push('\n');
for line in lines.iter().skip(1) {
if !line.is_empty() {
output.push_str(indent);
output.push_str(&python_docstring_literal_text(line));
}
output.push('\n');
}
output.push_str(indent);
output.push_str("\"\"\"\n");
}
fn push_wrapped_python_docstring_line(
lines: &mut Vec<String>,
first_prefix: &str,
continuation_prefix: &str,
text: &str,
max_width: usize,
) {
if text.is_empty() {
lines.push(first_prefix.trim_end().to_string());
return;
}
let mut prefix = first_prefix;
let mut current = String::new();
for word in text.split_whitespace() {
let prefix_width = prefix.chars().count();
let current_width = current.chars().count();
let word_width = word.chars().count();
let separator_width = usize::from(!current.is_empty());
if current_width > 0
&& prefix_width + current_width + separator_width + word_width > max_width
{
lines.push(format!("{prefix}{current}"));
prefix = continuation_prefix;
current.clear();
}
if !current.is_empty() {
current.push(' ');
}
current.push_str(word);
}
lines.push(format!("{prefix}{current}"));
}
fn python_docstring_literal_text(value: &str) -> String {
let mut escaped = value
.replace('\\', "\\\\")
.replace("\"\"\"", "\\\"\\\"\\\"");
// A body that ends in an unescaped `"` would run into the closing `"""`
// delimiter and produce four consecutive quotes, which does not parse.
// Escaping the final quote keeps the rendered text identical while making
// the delimiter unambiguous. A quote that the passes above already escaped
// is preceded by an odd number of backslashes and must be left alone.
if escaped.ends_with('"') {
let preceding_backslashes = escaped[..escaped.len() - 1]
.chars()
.rev()
.take_while(|character| *character == '\\')
.count();
if preceding_backslashes % 2 == 0 {
escaped.insert(escaped.len() - 1, '\\');
}
}
escaped
}
fn render_unpacked_operation_function(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
unpacked_input: &RenderedUnpackedInput,
) {
if !unpacked_input.functions.is_empty() {
render_function_unpacked_implementation(output, service, operation, unpacked_input);
return;
}
if operation.deprecated {
output.push_str(
"@typing_extensions.deprecated(\"This operation is deprecated.\", category=None)\n",
);
}
output.push_str("async def ");
output.push_str(&operation.attr_name);
output.push_str("(\n");
render_endpoint_parameter(output, service);
output.push_str(" *,\n");
for field in &unpacked_input.parameters {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
output.push_str(&python_parameter_annotation(
&field.annotation,
&field.default_kind,
));
if let Some(default_expr) = python_parameter_default_expr(&field.default_kind) {
render_python_default_expr(output, &default_expr, " ");
}
output.push_str(",\n");
}
if operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result
{
output.push_str(") -> ");
output.push_str(&function_unpacked_implementation_return_annotation(
service, operation,
));
output.push_str(":\n");
} else {
output.push_str(") -> ");
output.push_str(&temporalio_workflow_type_ref(
service,
"NexusOperationHandle",
));
output.push_str("[\n");
output.push_str(" ");
output.push_str(&function_unpacked_implementation_return_annotation(
service, operation,
));
output.push_str(",\n");
output.push_str("]:\n");
}
render_unpacked_operation_docstring(output, operation, unpacked_input, None);
for flattened_message in &unpacked_input.flattened_messages {
render_flattened_message_setup(output, flattened_message);
}
output.push_str(" request = ");
output.push_str(&unpacked_input.model_name);
output.push_str("(\n");
for field in &unpacked_input.request_fields {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push('=');
output.push_str(&unpacked_request_field_value_expr(field));
output.push_str(",\n");
}
output.push_str(" )\n");
output.push_str(" return await ");
output.push_str(&request_operation_function_name(service, operation));
output.push('(');
output.push_str(endpoint_forward_args(service));
output.push_str(")\n");
}
fn function_unpacked_implementation_return_annotation(
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) -> String {
let output_annotation = erase_python_type_parameters_with_replacement(
&operation.overload_output_annotation,
&operation.output_type_parameters,
"typing.Any",
);
temporalio_workflow_type_annotation(service, &output_annotation)
}
fn operation_public_return_annotation(
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
) -> String {
let output_annotation = function_unpacked_implementation_return_annotation(service, operation);
if operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result
{
output_annotation
} else {
format!(
"{}[{output_annotation}]",
temporalio_workflow_type_ref(service, "NexusOperationHandle")
)
}
}
fn render_function_unpacked_implementation(
output: &mut String,
service: &RenderedService<'_>,
operation: &RenderedOperation<'_>,
unpacked_input: &RenderedUnpackedInput,
) {
if operation.deprecated {
output.push_str(
"@typing_extensions.deprecated(\"This operation is deprecated.\", category=None)\n",
);
}
output.push_str("async def ");
output.push_str(&operation.attr_name);
output.push_str("(\n");
render_endpoint_parameter(output, service);
let primary_function = primary_function(&unpacked_input.functions);
if let Some(function) = primary_function {
output.push_str(" ");
output.push_str(&function.callable_field_name);
output.push_str(": ");
output.push_str(&function.callable_annotation);
output.push_str(",\n");
if let Some(parameters) = python_function_fixed_positional_args(function) {
for parameter in parameters {
output.push_str(" ");
output.push_str(¶meter.name);
output.push_str(": ");
output.push_str(¶meter.annotation);
output.push_str(",\n");
}
} else {
output.push_str(" *positional_args: ");
output.push_str(&python_function_positional_args_annotation(function));
output.push_str(",\n");
}
}
let has_keyword_only_parameters = unpacked_input
.parameters
.iter()
.any(|field| rendered_function_implementation_renders_field(field, primary_function));
if primary_function.is_none()
|| (primary_function
.and_then(python_function_fixed_positional_args)
.is_some()
&& has_keyword_only_parameters)
{
output.push_str(" *,\n");
}
for field in &unpacked_input.parameters {
if !rendered_function_implementation_renders_field(field, primary_function) {
continue;
}
output.push_str(" ");
output.push_str(&field.attr_name);
output.push_str(": ");
if let Some(function) = unpacked_input
.functions
.iter()
.find(|function| function.callable_field_name == field.attr_name)
{
output.push_str(&function.callable_annotation);
} else if unpacked_input
.functions
.iter()
.any(|function| function.args_field_name == field.attr_name)
{
let function = unpacked_input
.functions
.iter()
.find(|function| function.args_field_name == field.attr_name)
.expect("function args field should have a function");
let primary_args = primary_function
.is_some_and(|function| function.args_field_name == field.attr_name);
output.push_str(&python_function_args_parameter_annotation(
function,
primary_args,
));
} else {
output.push_str(&python_parameter_annotation(
&field.annotation,
&field.default_kind,
));
}
if let Some(default_expr) = python_parameter_default_expr(&field.default_kind) {
render_python_default_expr(output, &default_expr, " ");
}
output.push_str(",\n");
}
if operation.output_transform_expr.is_some()
|| operation.output_resource_return.is_some()
|| operation.output_direct_result
{
output.push_str(") -> ");
output.push_str(&function_unpacked_implementation_return_annotation(
service, operation,
));
output.push_str(":\n");
} else {
output.push_str(") -> ");
output.push_str(&temporalio_workflow_type_ref(
service,
"NexusOperationHandle",
));
output.push_str("[\n");
output.push_str(" ");
output.push_str(&function_unpacked_implementation_return_annotation(
service, operation,
));
output.push_str(",\n");
output.push_str("]:\n");
}
render_unpacked_operation_docstring(output, operation, unpacked_input, primary_function);
for function in &unpacked_input.functions {
if function.primary && matches!(function.args, RenderedFunctionArgs::Typed { .. }) {
continue;
}
let normalized_args_name = format!("normalized_{}", function.args_field_name);
if function.primary {
output.push_str(" if positional_args and ");
output.push_str(&function.args_field_name);
output.push_str(" is not None:\n");
output.push_str(
" raise TypeError(\"cannot specify both positional arguments and ",
);
output.push_str(&function.args_field_name);
output.push_str("\")\n");
output.push_str(" ");
output.push_str(&normalized_args_name);
output.push_str(": ");
output.push_str(&python_function_normalized_args_annotation(function));
output.push_str(" | None = (\n");
output.push_str(" list(positional_args)\n");
output.push_str(" if positional_args\n");
output.push_str(" else ");
output.push_str(&function.args_field_name);
output.push('\n');
output.push_str(" )\n");
} else {
output.push_str(" ");
output.push_str(&normalized_args_name);
output.push_str(": ");
output.push_str(&python_function_normalized_args_annotation(function));
output.push_str(" | None\n");
output.push_str(" if ");
output.push_str(&function.args_field_name);
output.push_str(" is None:\n");
output.push_str(" ");
output.push_str(&normalized_args_name);
output.push_str(" = None\n");
output.push_str(" elif isinstance(");
output.push_str(&function.args_field_name);
output.push_str(", list):\n");
output.push_str(" ");
output.push_str(&normalized_args_name);
match &function.args {
RenderedFunctionArgs::Varargs { .. } => {
output.push_str(" = typing.cast(list[typing.Any], ");
output.push_str(&function.args_field_name);
output.push_str(")\n");
}
_ => {
output.push_str(" = ");
output.push_str(&function.args_field_name);
output.push('\n');
}
}
output.push_str(" else:\n");
output.push_str(" ");
output.push_str(&normalized_args_name);
output.push_str(" = [");
output.push_str(&function.args_field_name);
output.push_str("]\n");
}
}
for flattened_message in &unpacked_input.flattened_messages {
render_flattened_message_setup(output, flattened_message);
}
output.push_str(" request = ");
output.push_str(&unpacked_input.model_name);
output.push_str("(\n");
for field in &unpacked_input.request_fields {
let value_expr = unpacked_input
.functions
.iter()
.find(|function| function.args_field_name == field.value_expr)
.map(|function| format!("normalized_{}", function.args_field_name))
.unwrap_or_else(|| unpacked_request_field_value_expr(field));
output.push_str(" ");
output.push_str(&field.attr_name);
output.push('=');
output.push_str(&value_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
output.push_str(" return await ");
output.push_str(&request_operation_function_name(service, operation));
output.push('(');
output.push_str(endpoint_forward_args(service));
output.push_str(")\n");
}
fn rendered_function_implementation_renders_field(
field: &RenderedUnpackedInputField,
primary_function: Option<&RenderedFunctionField>,
) -> bool {
if primary_function.is_some_and(|function| function.callable_field_name == field.attr_name) {
return false;
}
if primary_function
.and_then(python_function_fixed_positional_args)
.is_some_and(|parameters| {
parameters
.iter()
.any(|parameter| parameter.name == field.attr_name)
})
{
return false;
}
true
}
fn render_flattened_message_setup(
output: &mut String,
flattened_message: &RenderedFlattenedMessage,
) {
let always_construct = flattened_message.required
|| flattened_message
.fields
.iter()
.any(|field| field.default_kind == PythonFieldDefaultKind::Required);
output.push_str(" ");
output.push_str(&flattened_message.local_name);
output.push_str(" = ");
if always_construct {
output.push_str(&flattened_message.model_name);
output.push_str("(\n");
for field in &flattened_message.fields {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push('=');
output.push_str(&field.value_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
return;
}
output.push_str("(\n");
output.push_str(" None\n");
output.push_str(" if ");
for (index, field) in flattened_message.fields.iter().enumerate() {
if index > 0 {
output.push_str(" and ");
}
output.push_str(&field.value_expr);
output.push_str(" is None");
}
output.push_str("\n");
output.push_str(" else ");
output.push_str(&flattened_message.model_name);
output.push_str("(\n");
for field in &flattened_message.fields {
output.push_str(" ");
output.push_str(&field.attr_name);
output.push('=');
output.push_str(&field.value_expr);
output.push_str(",\n");
}
output.push_str(" )\n");
output.push_str(" )\n");
}
fn render_inline_nexus_client(output: &mut String, service: &RenderedService<'_>, indent: &str) {
output.push_str(indent);
if service.delay_load_temporalio_workflow {
output.push_str("nexus_client = create_nexus_client(\n");
} else {
output.push_str("nexus_client = temporalio.workflow.create_nexus_client(\n");
}
output.push_str(indent);
output.push_str(" service=");
output.push_str(&python_string_literal(service.wire_name));
output.push_str(",\n");
output.push_str(indent);
output.push_str(" endpoint=");
output.push_str(&python_service_endpoint_expr(service));
output.push_str(",\n");
output.push_str(indent);
output.push_str(")\n");
}
fn python_service_endpoint_expr(service: &RenderedService<'_>) -> String {
service
.endpoint
.as_ref()
.map(|endpoint| python_string_literal(endpoint))
.unwrap_or_else(|| "endpoint".to_string())
}
fn python_parameter_default_expr(default_kind: &PythonFieldDefaultKind) -> Option<String> {
match default_kind {
PythonFieldDefaultKind::Required => None,
PythonFieldDefaultKind::None => Some("None".to_string()),
PythonFieldDefaultKind::EmptyList | PythonFieldDefaultKind::EmptyDict => {
Some("None".to_string())
}
PythonFieldDefaultKind::Expression(expr) => Some(expr.clone()),
}
}
fn python_dataclass_source_default_expr(source_expr: &str) -> String {
let source_provider = source_expr.strip_suffix("()").unwrap_or(source_expr);
if source_provider
.chars()
.all(|character| is_python_identifier_char(character) || character == '.')
{
return format!("dataclasses.field(default_factory={source_provider})");
}
source_expr.to_string()
}
pub(in crate::generator) fn render_python_default_expr(
output: &mut String,
default_expr: &str,
indent: &str,
) {
if default_expr.chars().count() <= 48 {
output.push_str(" = ");
output.push_str(default_expr);
return;
}
output.push_str(" = (\n");
output.push_str(indent);
output.push_str(" ");
output.push_str(default_expr);
output.push('\n');
output.push_str(indent);
output.push(')');
}
pub(in crate::generator) fn python_parameter_annotation(
annotation: &str,
default_kind: &PythonFieldDefaultKind,
) -> String {
if matches!(
default_kind,
PythonFieldDefaultKind::EmptyList | PythonFieldDefaultKind::EmptyDict
) && !annotation.contains("| None")
{
format!("{annotation} | None")
} else {
annotation.to_string()
}
}
fn unpacked_request_field_value_expr(field: &RenderedUnpackedRequestField) -> String {
match &field.default_kind {
PythonFieldDefaultKind::EmptyList => format!("{} or []", field.value_expr),
PythonFieldDefaultKind::EmptyDict => format!("{} or {{}}", field.value_expr),
PythonFieldDefaultKind::Required
| PythonFieldDefaultKind::None
| PythonFieldDefaultKind::Expression(_) => field.value_expr.clone(),
}
}
fn python_ident(name: &str) -> String {
if is_python_keyword(name) {
format!("{name}_")
} else {
name.to_string()
}
}
fn validate_python_generated_names<'a>(
api_plan: &PlannedSpec,
generated_names: impl IntoIterator<Item = &'a PythonGeneratedName>,
) -> Result<()> {
let mut declarations = BTreeMap::<String, String>::new();
for entry in api_plan.types.values() {
let (name, description) = match &entry.declaration {
TypeDeclSpec::Record(record) => {
(&record.name, format!("record `{}`", record.full_name))
}
TypeDeclSpec::Enum(enumeration) => (
&enumeration.name,
format!("enum `{}`", enumeration.full_name),
),
TypeDeclSpec::Flags(flags) => (&flags.name, format!("flags `{}`", flags.full_name)),
TypeDeclSpec::Variant(variant) => {
(&variant.name, format!("variant `{}`", variant.full_name))
}
TypeDeclSpec::External(binding) => {
let Some(json) = binding.json_model() else {
continue;
};
(
&json.model_name,
format!("JSON Schema model `{}`", json.full_name),
)
}
};
declarations.entry(name.clone()).or_insert(description);
}
for (_, record) in api_plan.records() {
for usage in api_plan.record_type_parameters(&record.full_name, Language::Python) {
declarations
.entry(usage.parameter.name.clone())
.or_insert_with(|| format!("type parameter `{}`", usage.parameter.full_name));
}
}
for (_, variant) in api_plan.variants() {
for usage in api_plan.variant_type_parameters(&variant.full_name, Language::Python) {
declarations
.entry(usage.parameter.name.clone())
.or_insert_with(|| format!("type parameter `{}`", usage.parameter.full_name));
}
}
for generated_name in generated_names {
if let Some(conflicting_declaration) = declarations.get(&generated_name.name) {
return Err(Error::PythonGeneratedNameConflict {
name: generated_name.name.clone(),
generated_by: generated_name.generated_by.clone(),
conflicting_declaration: conflicting_declaration.clone(),
});
}
declarations.insert(
generated_name.name.clone(),
generated_name.generated_by.clone(),
);
}
Ok(())
}
pub(in crate::generator) fn python_string_literal(value: &str) -> String {
format!("{value:?}")
}
fn is_python_keyword(name: &str) -> bool {
matches!(
name,
"False"
| "None"
| "True"
| "and"
| "as"
| "assert"
| "async"
| "await"
| "break"
| "class"
| "continue"
| "def"
| "del"
| "elif"
| "else"
| "except"
| "finally"
| "for"
| "from"
| "global"
| "if"
| "import"
| "in"
| "is"
| "lambda"
| "nonlocal"
| "not"
| "or"
| "pass"
| "raise"
| "return"
| "try"
| "while"
| "with"
| "yield"
| "match"
| "case"
)
}
#[cfg(test)]
mod tests {
use std::collections::{BTreeMap, BTreeSet};
use std::fs;
use std::path::{Path, PathBuf};
use std::process::Command;
use std::time::{SystemTime, UNIX_EPOCH};
use crate::SupportFiles;
use crate::descriptors::DescriptorIndex;
use crate::error::Error;
use crate::generator::{
GenerateFilesOptions, GeneratedOutputLayout, GenerationMode,
generate_files_for_tree_with_mode_and_options, generate_source,
};
use crate::language::Language;
use crate::nexgen_config::{NexgenConfig, current, scope};
use crate::spec::ApiSpecTree;
use crate::spec::{LanguageImportSpec, LanguageImportStyle};
fn sample_input_path(root: &std::path::Path) -> PathBuf {
root.join("advanced/samples/inputs/workflow-service.wit")
}
fn start_workflow_input_path(root: &std::path::Path) -> PathBuf {
root.join("advanced/samples/inputs/start-workflow.wit")
}
fn type_roundtrip_input_path(root: &std::path::Path) -> PathBuf {
root.join("advanced/samples/inputs/type-roundtrip.wit")
}
fn linked_inputs_path(root: &std::path::Path) -> PathBuf {
root.join("advanced/samples/inputs/deps")
}
fn example_input_paths(root: &std::path::Path, input_path: PathBuf) -> Vec<PathBuf> {
vec![input_path, linked_inputs_path(root)]
}
#[test]
fn renders_source_provider_defaults_as_dataclass_factories() {
assert_eq!(
super::python_dataclass_source_default_expr("workflow_namespace"),
"dataclasses.field(default_factory=workflow_namespace)"
);
assert_eq!(
super::python_dataclass_source_default_expr("workflow_namespace()"),
"dataclasses.field(default_factory=workflow_namespace)"
);
}
#[test]
fn rendered_model_fragments_merge_root_package_imports() {
let mut fragments = super::RenderedModelFragments::default();
fragments.root_package_imports.insert(
".runtime".to_string(),
BTreeSet::from(["First".to_string(), "Shared".to_string()]),
);
let mut other = super::RenderedModelFragments::default();
other.root_package_imports.insert(
".runtime".to_string(),
BTreeSet::from(["Second".to_string(), "Shared".to_string()]),
);
other.root_package_imports.insert(
".support".to_string(),
BTreeSet::from(["Support".to_string()]),
);
fragments.extend(other);
assert_eq!(
fragments.root_package_imports,
BTreeMap::from([
(
".runtime".to_string(),
BTreeSet::from([
"First".to_string(),
"Second".to_string(),
"Shared".to_string(),
]),
),
(
".support".to_string(),
BTreeSet::from(["Support".to_string()]),
),
])
);
}
#[test]
fn generated_tree_source_label_names_every_input() {
let paths = BTreeSet::from([
PathBuf::from("schemas/a.yaml"),
PathBuf::from("schemas/nested/b.yaml"),
]);
assert_eq!(
super::source_paths_label(&paths),
"schemas/a.yaml, schemas/nested/b.yaml"
);
}
fn unique_temp_dir(label: &str) -> PathBuf {
let unique = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
std::env::temp_dir().join(format!("nex-gen-python-{label}-{unique}"))
}
fn write_package(dir: &Path, files: &BTreeMap<PathBuf, String>) {
for (relative_path, contents) in files {
let path = dir.join(relative_path);
if let Some(parent) = path.parent() {
fs::create_dir_all(parent).unwrap();
}
fs::write(path, contents).unwrap();
}
}
#[test]
fn optional_imports_are_selected_from_generic_references() {
let mut module_imports = std::collections::BTreeSet::new();
module_imports.insert("temporalio.api.common.v1.message_pb2".to_string());
module_imports.insert("temporalio.api.unused.v1.message_pb2".to_string());
let body = r#"
@dataclasses.dataclass
class Example(enum.Enum):
value: collections.abc.Mapping[str, typing.Any]
retry_policy: temporalio.common.RetryPolicy
request: temporalio.api.common.v1.message_pb2.WorkflowExecution
deadline: datetime.timedelta
def run(self) -> typing_extensions.Self:
temporalio.workflow.create_nexus_client(service="svc", endpoint="endpoint")
return self
"#;
let language_imports = [
LanguageImportSpec {
language: Language::Python,
reference: "temporalio.common".to_string(),
module: "temporalio.common".to_string(),
name: None,
type_only: false,
import_style: LanguageImportStyle::Module,
},
LanguageImportSpec {
language: Language::Python,
reference: "datetime".to_string(),
module: "datetime".to_string(),
name: None,
type_only: false,
import_style: LanguageImportStyle::Module,
},
LanguageImportSpec {
language: Language::Python,
reference: "temporalio.workflow".to_string(),
module: "temporalio.workflow".to_string(),
name: None,
type_only: false,
import_style: LanguageImportStyle::Module,
},
];
let mut output = String::new();
let wrote_imports = super::render_optional_python_imports(
&mut output,
body,
&module_imports,
&language_imports,
);
assert!(wrote_imports);
assert!(output.contains("import collections.abc\n"));
assert!(output.contains("import dataclasses\n"));
assert!(output.contains("import enum\n"));
assert!(output.contains("import typing\n"));
assert!(output.contains("import typing_extensions\n"));
assert!(output.contains("import temporalio.common\n"));
assert!(output.contains("import datetime\n"));
assert!(output.contains("import temporalio.workflow\n"));
assert!(output.contains("import temporalio.api.common.v1.message_pb2\n"));
assert!(!output.contains("temporalio.api.unused.v1.message_pb2"));
let mut output = String::new();
super::render_optional_python_imports(
&mut output,
"async def call(workflow: str) -> None:\n pass\n",
&std::collections::BTreeSet::new(),
&language_imports,
);
assert!(!output.contains("import temporalio.workflow\n"));
}
fn read_python_package_files(dir: &Path) -> BTreeMap<PathBuf, String> {
fn visit(root: &Path, dir: &Path, files: &mut BTreeMap<PathBuf, String>) {
let mut entries = fs::read_dir(dir)
.unwrap()
.map(|entry| entry.unwrap().path())
.collect::<Vec<_>>();
entries.sort();
for path in entries {
if path.is_dir() {
visit(root, &path, files);
} else if path.extension().and_then(|extension| extension.to_str()) == Some("py") {
if path
.file_name()
.and_then(|file_name| file_name.to_str())
.is_some_and(|file_name| file_name.starts_with("test_"))
{
continue;
}
files.insert(
path.strip_prefix(root).unwrap().to_path_buf(),
fs::read_to_string(&path).unwrap(),
);
}
}
}
let mut files = BTreeMap::new();
visit(dir, dir, &mut files);
files
}
fn format_python_output(
root: &std::path::Path,
files: &BTreeMap<PathBuf, String>,
) -> BTreeMap<PathBuf, String> {
let temp_dir = unique_temp_dir("format");
fs::create_dir_all(&temp_dir).unwrap();
write_package(&temp_dir, files);
let status = Command::new("uv")
.current_dir(root.join("advanced/samples/python"))
.args([
"run",
"ruff",
"format",
"--line-length",
"88",
temp_dir.to_str().unwrap(),
])
.status()
.unwrap();
assert!(status.success());
let formatted = read_python_package_files(&temp_dir);
fs::remove_dir_all(&temp_dir).unwrap();
formatted
}
#[test]
fn renders_sample_with_versioned_header() {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
let spec = crate::parser::load_api_spec_from_wit_for_language_with_inputs(
Language::Python,
&example_input_paths(&root, sample_input_path(&root)),
)
.unwrap();
let descriptors =
DescriptorIndex::load(&root.join("advanced/samples/descriptors/temporal_api.bin"))
.unwrap();
let _scope = scope(NexgenConfig {
system_nexus: true,
..current()
});
let generated = generate_files_for_tree_with_mode_and_options(
Language::Python,
ApiSpecTree::single(spec.clone()),
&descriptors,
&crate::SupportFiles::default(),
GenerationMode::NativeApi,
GenerateFilesOptions::default(),
)
.unwrap();
assert_eq!(generated.layout, GeneratedOutputLayout::Directory);
let output = format_python_output(&root, &generated.files);
assert!(
output.values().any(|contents| {
contents.contains(concat!("nexgen v", env!("CARGO_PKG_VERSION")))
})
);
assert!(!output.is_empty());
}
#[test]
fn renders_required_fields_and_custom_message_types() {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
let spec = crate::parser::load_api_spec_from_wit_for_language_with_inputs(
Language::Python,
&example_input_paths(&root, sample_input_path(&root)),
)
.unwrap();
let descriptors =
DescriptorIndex::load(&root.join("advanced/samples/descriptors/temporal_api.bin"))
.unwrap();
let output = generate_source(
Language::Python,
spec.clone(),
&descriptors,
&crate::SupportFiles::default(),
)
.unwrap();
assert!(output.contains("from ._support import ("));
assert!(output.contains("retry_policy_to_proto,"));
assert!(output.contains("def retry_policy_from_proto("));
assert!(output.contains(
"workflow: str | collections.abc.Callable[..., collections.abc.Awaitable[object]]"
));
assert!(output.contains("id: str"));
assert!(output.contains("task_queue: str"));
assert!(output.contains(
"signal: str | collections.abc.Callable[..., None | collections.abc.Awaitable[None]]"
));
assert!(output.contains("class SignalWithStartWorkflowRequest:"));
assert!(output.contains("args: list[typing.Any] | None = None"));
assert!(!output.contains("namespace: str | None = None"));
assert!(output.contains("dataclasses.field(default_factory=workflow_namespace)"));
assert!(output.contains("message.namespace = value.namespace"));
assert!(output.contains("result = await handle"));
assert!(output.contains(
"from temporalio.workflow import (\n create_nexus_client,\n get_external_workflow_handle,\n )"
));
assert!(output.contains("return get_external_workflow_handle("));
assert!(output.contains("run_id=result.run_id"));
assert!(output.contains("signal_args: list[typing.Any] | None = None"));
assert!(output.contains("temporalio.common.WorkflowIDReusePolicy.ALLOW_DUPLICATE"));
assert!(output.contains(
"id_conflict_policy: temporalio.common.WorkflowIDConflictPolicy | None = None"
));
assert!(!output.contains("identity: str | None"));
assert!(output.contains("memo: collections.abc.Mapping[str, typing.Any] | None = None"));
assert!(output.contains("user_metadata: UserMetadata | None = None"));
assert!(output.contains("static_summary: typing.Any | None = None"));
assert!(output.contains("static_details: typing.Any | None = None"));
assert!(output.contains("retry_policy: temporalio.common.RetryPolicy"));
assert!(
output.contains(
"search_attributes: temporalio.common.TypedSearchAttributes | None = None"
)
);
assert!(output.contains("priority: temporalio.common.Priority | None = None"));
assert!(
output.contains(
"versioning_override: temporalio.common.VersioningOverride | None = None"
)
);
assert!(output.contains("workflow_id_reuse_policy_to_proto(value.id_reuse_policy)"));
assert!(output.contains("if value.id_conflict_policy is not None:"));
assert!(output.contains("workflow_id_conflict_policy_to_proto(value.id_conflict_policy)"));
assert!(output.contains("message.input.CopyFrom(payloads_to_proto(value.args))"));
assert!(output.contains("headers: collections.abc.Mapping[str, typing.Any] | None = None"));
assert!(output.contains("message.header.CopyFrom(header_to_proto(value.headers))"));
assert!(!output.contains("links:"));
assert!(!output.contains("link_to_proto("));
assert!(output.contains("async def _signal_with_start_workflow("));
assert!(output.contains("request: SignalWithStartWorkflowRequest"));
assert!(output.contains(
"if typing.TYPE_CHECKING:\n from temporalio.workflow import ExternalWorkflowHandle"
));
assert!(output.contains(") -> ExternalWorkflowHandle[object]:"));
assert!(output.contains(" nexus_client = create_nexus_client("));
assert!(!output.contains("(typing.TypedDict, total=False):"));
assert!(!output.contains("typing.Unpack["));
assert!(!output.contains("args: tuple[typing.Any, ...] | None = None"));
assert!(!output.contains("signal_args: tuple[typing.Any, ...] | None = None"));
assert!(output.contains("@typing.overload"));
assert!(output.contains("workflow: str,"));
assert!(output.contains("*args: object,"));
assert!(output.contains("args: list[typing.Any] | None = ...,"));
assert!(output.contains(
"workflow: collections.abc.Callable[[SelfType, typing_extensions.Unpack[WorkflowArgs]], collections.abc.Awaitable[WorkflowResult]],"
));
assert!(output.contains("WorkflowArgs = typing_extensions.TypeVarTuple(\"WorkflowArgs\")"));
assert!(output.contains("WorkflowResult = typing.TypeVar(\"WorkflowResult\")"));
assert!(output.contains("SelfType = typing.TypeVar(\"SelfType\")"));
assert!(output.contains(") -> ExternalWorkflowHandle[SelfType]:"));
assert!(output.contains("async def signal_with_start_workflow("));
assert!(output.contains("*args: typing_extensions.Unpack[WorkflowArgs],"));
assert!(output.contains("args: list[typing.Any],"));
assert!(!output.contains("tuple[FirstWorkflowArg"));
assert!(output.contains(
"signal: collections.abc.Callable[[SelfType, SignalArg], None | collections.abc.Awaitable[None]],"
));
assert!(output.contains("SignalArg = typing.TypeVar(\"SignalArg\")"));
assert!(output.contains("signal_args: SignalArg,"));
assert!(output.contains("signal_args: list[typing.Any],"));
assert!(output.contains("async def signal_with_start_workflow("));
assert!(output.contains(
"async def signal_with_start_workflow(\n workflow: str | collections.abc.Callable[..., collections.abc.Awaitable[object]],\n *positional_args: object,\n args: list[typing.Any] | None = None,"
));
assert!(output.contains(
"signal: str | collections.abc.Callable[..., None | collections.abc.Awaitable[None]],"
));
assert!(output.contains("signal_args: object | list[typing.Any] | None = None,"));
assert_eq!(
output
.matches("Signal a workflow, starting it first if needed.")
.count(),
1
);
assert!(output.contains(
"\"\"\"Signal a workflow, starting it first if needed.\n\n .. warning::\n This API is experimental and subject to change.\n\n Args:\n workflow: Workflow type name or callable identifying the workflow to start.\n positional_args: Positional arguments for workflow. Cannot be set if args is\n set.\n args: List-form arguments for workflow. Cannot be set if positional_args are\n set. For typed workflow callables, list contents are not statically\n typechecked; pass workflow arguments positionally for precise typechecking.\n id: Unique identifier for the workflow execution. Must be nonempty.\n task_queue: Task queue to run the workflow on.\n signal: Signal name or callable to send with the start request.\n signal_args: Argument value, or list of argument values, for signal. For typed\n single-argument signals, scalar signal_args values are statically\n typechecked. List-form signal_args values are not precisely typechecked. To\n pass a single signal argument that is itself a list, wrap it in another\n list; otherwise the list is interpreted as multiple signal arguments."
));
assert!(output.contains(
"cron_schedule: Cron schedule for recurring workflow executions. See\n https://docs.temporal.io/cron-job."
));
assert!(output.contains(
"static_summary: Single-line fixed summary for the workflow execution that may\n appear in UI and CLI. This can be in single-line Temporal Markdown format."
));
assert!(output.contains(
"\n\n Returns:\n A workflow handle to the started workflow.\n \"\"\""
));
assert!(output.contains("static_summary: str | None = None,"));
assert!(output.contains("static_details: str | None = None,"));
assert!(!output.contains("user_metadata_static_summary:"));
assert!(!output.contains("user_metadata_static_details:"));
assert!(!output.contains("def _nexus_is_function_args_list("));
assert!(!output.contains("def _nexus_normalize_function_args("));
assert!(!output.contains("_nexus_arg_unset = object()"));
assert!(output.contains(
"normalized_signal_args: list[typing.Any] | None\n if signal_args is None:\n normalized_signal_args = None\n elif isinstance(signal_args, list):\n normalized_signal_args = typing.cast(list[typing.Any], signal_args)\n else:\n normalized_signal_args = [signal_args]"
));
assert!(output.contains("if positional_args and args is not None:"));
assert!(
output
.contains("raise TypeError(\"cannot specify both positional arguments and args\")")
);
assert!(output.contains("normalized_args: list[typing.Any] | None = ("));
assert!(output.contains("list(positional_args)"));
assert!(output.contains("else args"));
assert!(output.contains("user_metadata = ("));
assert!(output.contains("if static_summary is None and static_details is None"));
assert!(output.contains("static_summary=static_summary,"));
assert!(output.contains("static_details=static_details,"));
assert!(output.contains("args=normalized_args,"));
assert!(output.contains("signal_args=normalized_signal_args,"));
assert!(output.contains("user_metadata=user_metadata,"));
assert!(output.contains("request = SignalWithStartWorkflowRequest("));
assert!(output.contains("workflow=workflow,"));
assert!(output.contains("args=normalized_args,"));
assert!(output.contains("return await _signal_with_start_workflow(request)"));
assert!(!output.contains("SignalWithStartWorkflowRequest.from_proto"));
assert!(!output.contains(
"proto: temporalio.api.workflowservice.v1.SignalWithStartWorkflowExecutionRequest,\n ) -> SignalWithStartWorkflowRequest:"
));
assert!(!output.contains("class SignalWithStartWorkflowRequestTyped["));
assert!(!output.contains("async def signal_with_start_workflow_typed("));
assert!(!output.contains("class RetryPolicy:"));
assert!(!output.contains("class WorkflowType:"));
assert!(!output.contains("class TaskQueue:"));
assert!(!output.contains("class WorkflowIdReusePolicy(enum.IntEnum):"));
assert!(!output.contains("class WorkflowIdConflictPolicy(enum.IntEnum):"));
assert!(!output.contains("class Payload:"));
assert!(!output.contains("class ExternalPayloadDetails:"));
assert!(!output.contains("class Payloads:"));
assert!(!output.contains("class Memo:"));
assert!(!output.contains("class Header:"));
assert!(!output.contains("class SearchAttributes:"));
assert!(output.contains("class UserMetadata:"));
assert!(!output.contains("class Link:"));
assert!(!output.contains("class Priority:"));
assert!(!output.contains("class VersioningOverride:"));
assert!(!output.contains("build_signal_with_start_workflow_request"));
assert!(!output.contains("from temporal_model_converters import"));
let type_roundtrip_spec = crate::parser::load_api_spec_from_wit_for_language_with_inputs(
Language::Python,
&example_input_paths(&root, type_roundtrip_input_path(&root)),
)
.unwrap();
let type_roundtrip_output = generate_source(
Language::Python,
type_roundtrip_spec,
&descriptors,
&crate::SupportFiles::default(),
)
.unwrap();
assert!(
type_roundtrip_output.contains(
"raise ValueError(\"missing required field ActivityOptions.retry_policy\")"
)
);
assert!(type_roundtrip_output.contains("if not value.HasField(\"retry_policy\"):\n raise ValueError(\"missing required field ActivityOptions.retry_policy\")"));
assert!(type_roundtrip_output.contains("retry_policy_from_proto("));
assert!(type_roundtrip_output.contains("value.retry_policy"));
assert!(!type_roundtrip_output.contains("async def retry_policy_operation("));
assert!(type_roundtrip_output.contains("async def activity_options_operation("));
assert!(type_roundtrip_output.contains("task_queue: str | None = None,"));
assert!(type_roundtrip_output.contains("retry_policy: temporalio.common.RetryPolicy,"));
assert!(type_roundtrip_output.contains("request = ActivityOptions("));
}
#[test]
fn renders_experimental_annotations() {
let wit = r#"
package temporal:nexus@1.0.0;
world system {
export example-service;
}
/// @nexus.endpoint "example"
/// @nexus.experimental
interface example-service {
/// @nexus.experimental
record request {
id: string,
}
/// @nexus.experimental
record response {
ok: bool,
}
/// @nexus.experimental
/// @nexus.doc "Runs example."
request-op: func(request: request) -> response;
}
"#;
let spec = crate::parser::parse_api_spec_from_wit_for_language(
Language::Python,
wit,
PathBuf::from("inline.wit"),
)
.unwrap();
let descriptors = DescriptorIndex::load(
&PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("advanced/samples/descriptors/temporal_api.bin"),
)
.unwrap();
let output = generate_source(
Language::Python,
spec.clone(),
&descriptors,
&crate::SupportFiles::default(),
)
.unwrap();
assert!(output.contains(
"class Request:\n \"\"\"\n .. warning::\n This API is experimental and subject to change.\n \"\"\""
));
assert!(output.contains(
"class ExampleService:\n \"\"\"\n .. warning::\n This API is experimental and subject to change.\n \"\"\""
));
assert!(output.contains(
" # .. warning:: This API is experimental and subject to change.\n request_op: Operation"
));
assert!(output.contains(
" \"\"\"Runs example.\n\n .. warning::\n This API is experimental and subject to change."
));
}
#[test]
fn delays_temporalio_workflow_imports_when_enabled() {
let wit = r#"
package temporal:nexus@1.0.0;
world system {
export example-service;
}
/// @nexus.endpoint "example"
/// @nexus.delay-load-temporalio-workflow
interface example-service {
record request {
id: string,
}
record response {
ok: bool,
}
start: func(request: request) -> response;
/// @nexus.output-transform
/// python-type="temporalio.workflow.ExternalWorkflowHandle[str]"
/// python="temporalio.workflow.get_external_workflow_handle(request.id)"
get-handle: func(request: request) -> response;
}
"#;
let spec = crate::parser::parse_api_spec_from_wit_for_language(
Language::Python,
wit,
PathBuf::from("inline.wit"),
)
.unwrap();
let descriptors =
DescriptorIndex::from_descriptor_set(prost_types::FileDescriptorSet::default())
.unwrap();
let generated = generate_files_for_tree_with_mode_and_options(
Language::Python,
ApiSpecTree::single(spec.clone()),
&descriptors,
&crate::SupportFiles::default(),
GenerationMode::NativeApi,
GenerateFilesOptions::default(),
)
.unwrap();
let start = generated
.files
.get(&PathBuf::from("operations/start.py"))
.expect("start operation should be rendered");
assert!(!start.contains("\nimport temporalio.workflow\n\n"));
assert!(start.contains(
"if typing.TYPE_CHECKING:\n from temporalio.workflow import NexusOperationHandle\n"
));
assert!(start.contains(
") -> NexusOperationHandle[\n Response,\n]:\n from temporalio.workflow import create_nexus_client\n nexus_client = create_nexus_client("
));
let get_handle = generated
.files
.get(&PathBuf::from("operations/get_handle.py"))
.expect("get-handle operation should be rendered");
assert!(!get_handle.contains("\nimport temporalio.workflow\n\n"));
assert!(get_handle.contains(
"if typing.TYPE_CHECKING:\n from temporalio.workflow import ExternalWorkflowHandle\n"
));
assert!(get_handle.contains(
" from temporalio.workflow import (\n create_nexus_client,\n get_external_workflow_handle,\n )\n"
));
assert!(get_handle.contains(" nexus_client = create_nexus_client("));
assert!(get_handle.contains("input=request,"));
assert!(get_handle.contains("return get_external_workflow_handle(request.id)\n"));
}
#[test]
fn renders_resource_bound_operations_inside_resource_modules() {
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
let spec = crate::parser::load_api_spec_from_wit_for_language_with_inputs(
Language::Python,
&example_input_paths(&root, start_workflow_input_path(&root)),
)
.unwrap();
let descriptors =
DescriptorIndex::load(&root.join("advanced/samples/descriptors/temporal_api.bin"))
.unwrap();
let output = generate_source(
Language::Python,
spec.clone(),
&descriptors,
&crate::SupportFiles::default(),
)
.unwrap();
assert!(output.contains("### _resources/started_workflow.py"));
assert!(output.contains("### operations/start_workflow.py"));
assert!(output.contains("async def cancel_workflow("));
assert!(output.contains("async def restart_workflow("));
assert!(output.contains("async def cancel("));
assert!(output.contains("async def restart_workflow("));
assert!(!output.contains("should be replaced during package initialization"));
}
#[test]
fn rejects_flattened_api_field_name_conflicts() {
let wit = r#"
package temporal:nexus@1.0.0;
world system {
export workflow-service;
}
/// @nexus.endpoint "__temporal_system"
interface workflow-service {
use nexus:temporal-types/model@1.0.0.{
duration,
placeholder,
signal-function,
task-queue,
user-metadata,
workflow-function,
};
/// @nexus.proto "temporal.api.workflowservice.v1.SignalWithStartWorkflowExecutionRequest"
record signal-with-start-workflow-request {
/// @nexus.proto-field "workflow_type"
workflow: workflow-function,
workflow-id: string,
task-queue: task-queue,
/// @nexus.proto-field "signal_name"
signal: signal-function,
/// @nexus.proto-field "request_id"
static-summary: string,
user-metadata: option<user-metadata>,
/// @nexus.source "workflow_namespace()"
namespace: option<string>,
/// @nexus.omit
workflow-execution-timeout: placeholder,
/// @nexus.omit
workflow-run-timeout: placeholder,
/// @nexus.omit
workflow-task-timeout: placeholder,
/// @nexus.omit
identity: placeholder,
/// @nexus.omit
workflow-id-reuse-policy: placeholder,
/// @nexus.omit
workflow-id-conflict-policy: placeholder,
/// @nexus.omit
control: placeholder,
/// @nexus.omit
retry-policy: placeholder,
/// @nexus.omit
cron-schedule: placeholder,
/// @nexus.omit
memo: placeholder,
/// @nexus.omit
search-attributes: placeholder,
/// @nexus.omit
header: placeholder,
workflow-start-delay: option<duration>,
/// @nexus.omit
links: placeholder,
/// @nexus.omit
versioning-override: placeholder,
/// @nexus.omit
priority: placeholder,
/// @nexus.omit
time-skipping-config: placeholder,
}
/// @nexus.proto "temporal.api.workflowservice.v1.SignalWithStartWorkflowExecutionResponse"
record signal-with-start-workflow-response {
run-id: option<string>,
started: option<bool>,
/// @nexus.omit
signal-link: placeholder,
}
signal-with-start-workflow-execution: func(
request: signal-with-start-workflow-request,
) -> signal-with-start-workflow-response;
}
"#;
let temp_dir = unique_temp_dir("flatten-conflict");
fs::create_dir_all(&temp_dir).unwrap();
let input_path = temp_dir.join("conflict.wit");
fs::write(&input_path, wit).unwrap();
let root = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
let spec = crate::parser::load_api_spec_from_wit_for_language_with_inputs(
Language::Python,
&example_input_paths(&root, input_path.clone()),
)
.unwrap();
let descriptors =
DescriptorIndex::load(&root.join("advanced/samples/descriptors/temporal_api.bin"))
.unwrap();
let error = generate_source(
Language::Python,
spec.clone(),
&descriptors,
&crate::SupportFiles::default(),
)
.unwrap_err();
fs::remove_dir_all(temp_dir).unwrap();
assert!(matches!(
error,
Error::FlattenedApiFieldConflict {
type_name,
field,
conflicting_field,
} if type_name == "SignalWithStartWorkflowRequest"
&& field == "static_summary"
&& conflicting_field.contains("static_summary")
));
}
#[test]
fn allows_missing_service_endpoint() {
let wit = r#"
package temporal:nexus@1.0.0;
world system {
export example-service;
}
interface example-service {
use nexus:temporal-types/model@1.0.0.{retry-policy};
example-operation: func(request: retry-policy) -> retry-policy;
}
"#;
let spec = crate::parser::parse_api_spec_from_wit_for_language_with_inputs(
Language::Python,
wit,
PathBuf::from("inline.wit"),
&[linked_inputs_path(&PathBuf::from(env!(
"CARGO_MANIFEST_DIR"
)))],
)
.unwrap();
let descriptors = DescriptorIndex::load(
&PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("advanced/samples/descriptors/temporal_api.bin"),
)
.unwrap();
let output = generate_source(
Language::Python,
spec.clone(),
&descriptors,
&SupportFiles::default(),
)
.unwrap();
assert!(
output.contains("class ExampleServiceClient:\n def __init__(self, endpoint: str)")
);
assert!(output.contains("async def example_operation(\n self,\n request:"));
assert!(
output.contains(
"self._nexus_client: typing.Any = temporalio.workflow.create_nexus_client("
)
);
assert!(output.contains("endpoint=endpoint,"));
assert!(output.contains("return await self._nexus_client.start_operation("));
}
}