diff --git a/CHANGELOG.md b/CHANGELOG.md index f94345441..9ca6a42dd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,12 @@ release tags add a leading `v` to the package version. ## Unreleased +- Forwardable Fortran optional arguments use a linear number of contained + procedures and converge on one native call site instead of enumerating + presence combinations. Descriptor categories that cannot be forwarded, + including optional assumed-rank arrays, preserve `PRESENT()` through direct + present/absent call leaves. + - Optional Fortran callbacks and optional reference dummies inside callback interfaces preserve `PRESENT()` through source and generated-contract builds. diff --git a/docs/developer/packages/codegen/fortran-bridge.md b/docs/developer/packages/codegen/fortran-bridge.md index 3984433dc..94fb637e2 100644 --- a/docs/developer/packages/codegen/fortran-bridge.md +++ b/docs/developer/packages/codegen/fortran-bridge.md @@ -72,6 +72,24 @@ them. 3. It runs the selected writeback and cleanup finalizers, wrapping derived result or carrier lifecycles when the plan requires them. +Forwardable optional arguments use a linear number of contained procedures and +converge on one native call site instead of enumerating presence combinations. +Each procedure introduces one planned optional actual and forwards the optional +dummies already introduced. Fortran therefore propagates absence through its +own optional-dummy rules. When the binding has already entered Fortran-owned +descriptors through its inverted consumer chain, that outer chain remains +active until this inner native call returns. +Mutable deferred-length character descriptors remain on direct present/absent +call leaves because compiler descriptor updates do not propagate reliably +through another optional dummy. Optional assumed-rank arrays also remain on +direct leaves because Fortran cannot declare the local assumed-rank pointer +that descriptor transport would require. Their completed entrypoint ABI carries +an explicit presence value beside the caller's descriptor or a valid rank-zero +ordinary placeholder. This avoids compiler-dependent rank loss at an optional +assumed-rank `bind(C)` dummy; the placeholder is never passed to the native +procedure. Other optionals in the same procedure still use the forwarding +chain. + The entrypoint record exposes a `bind(C)` name shared with the C binding. For a standalone native procedure, the bridge record explicitly selects its external declaration; for a module procedure, it supplies the native module use. Those diff --git a/docs/user/guide/optional-arguments.md b/docs/user/guide/optional-arguments.md index 665dd61ed..563dd46a7 100644 --- a/docs/user/guide/optional-arguments.md +++ b/docs/user/guide/optional-arguments.md @@ -137,6 +137,8 @@ Result: - Providing a concrete value makes the argument **present**. - Use **keyword arguments** when skipping earlier optional parameters. - Optional arrays and derived types also accept `None` to indicate absence. +- Optional assumed-rank arrays accept ranks 1 through 15 when present; omission + and `None` preserve `present(...) == .false.`. - Optional `intent(out)` / `intent(inout)` arguments remain visible in Python so you can control `present(...)`. - An optional argument without `intent` uses the same conservative diff --git a/prik/codegen/c/binding.py b/prik/codegen/c/binding.py index 11ec4a24b..4ca13a377 100644 --- a/prik/codegen/c/binding.py +++ b/prik/codegen/c/binding.py @@ -56,6 +56,7 @@ NativeDescriptorHandoffABI, EntrypointProjectionAction, EntrypointPassingConvention, + EntrypointOptionalityAction, OptionalMode, OverloadMatchKind, PythonExceptionKind, @@ -7807,6 +7808,8 @@ def _descriptor_array_argument_declarations( CDeclaration(f"{names.value_name}_extents_out", self.ARRAY_EXTENTS_RECORD), *(CDeclaration(name, "int64_t", CodeExpression("0")) for name in names.extent_names), ] + if plan.entrypoint.pass_descriptor_presence: + declarations.append(CDeclaration(names.present_name, "void *", CodeExpression("NULL"))) if plan.transformations: declarations.append( CDeclaration( @@ -10457,7 +10460,12 @@ def _lower_entrypoint_call(self, plan: FunctionPlan, context: _CFunctionContext) CExpressionStatement(CodeExpression(f"prik_native_array_backend_release_call({backend})")) for backend in reversed(leased) ) - return (*acquire_nodes, *body, *release_nodes) + return ( + *self._explicit_descriptor_presence_nodes(plan, context), + *acquire_nodes, + *body, + *release_nodes, + ) def _lower_entrypoint_call_with_live_descriptors( self, @@ -10495,6 +10503,66 @@ def _lower_entrypoint_call_with_live_descriptors( ), ) + def _explicit_descriptor_presence_nodes( + self, + plan: FunctionPlan, + context: _CFunctionContext, + ) -> tuple: + """Supply a valid rank-zero ordinary descriptor beside explicit presence.""" + nodes = [] + for owner_path in context.inverted_descriptors: + argument = self._argument_by_owner(plan, owner_path) + if ( + argument.entrypoint.optionality + is not EntrypointOptionalityAction.EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR + ): + continue + if argument.array is None: + raise ValueError(f"Placeholder descriptor {argument.owner_path!r} has no array handoff") + names = context.arguments[owner_path] + placeholder = f"{names.value_name}_placeholder" + placeholder_type = ( + "char" + if argument.datatype_family is DatatypeFamily.STRING + else PrimitiveScalarTypeRegistry.type_for(argument.semantic_type_name).array_c_spelling + ) + nodes.extend( + ( + CDeclaration(placeholder, placeholder_type, CodeExpression("0")), + CExpressionStatement( + CodeExpression( + f"{names.present_name} = {names.object_name} != Py_None ? " + f"(void *){names.object_name} : NULL" + ) + ), + CIf( + CodeExpression(f"{names.value_name} == NULL"), + body=( + CIf( + CodeExpression( + f"CFI_establish((CFI_cdesc_t *)&{names.value_name}_section, &{placeholder}, " + f"CFI_attribute_other, {self._native_array_cfi_type(argument)}, " + f"{self._native_array_expected_element_size(argument)}, 0, NULL) != CFI_SUCCESS" + ), + body=( + CExpressionStatement( + CodeExpression( + 'PyErr_SetString(PyExc_RuntimeError, "Could not create an absent ' + f'descriptor for argument {argument.binding.python_name}")' + ) + ), + CReturn(CodeExpression("NULL")), + ), + ), + CExpressionStatement( + CodeExpression(f"{names.value_name} = (CFI_cdesc_t *)&{names.value_name}_section") + ), + ), + ), + ) + ) + return tuple(nodes) + @staticmethod def _array_crosses_as_descriptor(argument: ArgumentTransferPlan) -> bool: """Report whether completed policy hands this array over as a descriptor.""" @@ -14366,7 +14434,10 @@ def _ordinary_entrypoint_argument_parameters( if self._array_crosses_as_descriptor(argument): # Extents and strides travel inside the descriptor, so the # address and the fields beside it are not needed. - return (CParameter(name, "CFI_cdesc_t *"),) + parameters = [CParameter(name, "CFI_cdesc_t *")] + if argument.entrypoint.pass_descriptor_presence: + parameters.append(CParameter(f"{name}_present", "void *")) + return tuple(parameters) return self._array_entrypoint_argument_parameters(argument, name) if argument.entrypoint.handoff_mode is ArgumentHandoffMode.NATIVE_DESCRIPTOR: handle = argument.native_array_handle diff --git a/prik/codegen/fortran/bridge.py b/prik/codegen/fortran/bridge.py index e1035b9ba..caafc57c4 100644 --- a/prik/codegen/fortran/bridge.py +++ b/prik/codegen/fortran/bridge.py @@ -56,6 +56,7 @@ NativeEntrypointAction, NativeInvocationKind, EntrypointPassingConvention, + EntrypointOptionalityAction, EntrypointProjectionAction, OptionalMode, ScalarLogicalABI, @@ -4084,13 +4085,31 @@ def _lower_argument_array_descriptor( here: the dummy is the array, with the bounds and directions the caller described, and it is handed to the native procedure as it stands. """ - return ( + parameters = [ self._array_descriptor_parameter( plan, plan.entrypoint.parameter_name, - optional=plan.entrypoint.optional_mode is not OptionalMode.REQUIRED, - ), - ) + optional=( + plan.entrypoint.optional_mode is not OptionalMode.REQUIRED + and not plan.entrypoint.pass_descriptor_presence + ), + target=( + plan.entrypoint.optional_mode is not OptionalMode.REQUIRED + and not plan.entrypoint.pass_descriptor_presence + and plan.array is not None + and plan.array.rank is not None + ), + ) + ] + if plan.entrypoint.pass_descriptor_presence: + parameters.append( + FortranParameter( + f"bound_{plan.entrypoint.parameter_name}_present", + "type(c_ptr)", + ("value",), + ) + ) + return tuple(parameters) def _array_descriptor_parameter( self, @@ -4098,6 +4117,7 @@ def _array_descriptor_parameter( name: str, *, optional: bool, + target: bool = False, ) -> FortranParameter: """Declare one interoperable ordinary-array descriptor dummy.""" array = plan.array @@ -4114,6 +4134,8 @@ def _array_descriptor_parameter( if optional: # C omits it by passing no descriptor, which is what optional means # for an interoperable dummy. + if target: + attributes.append("target") attributes.append("optional") return FortranParameter(name, element_type, tuple(attributes)) @@ -4185,26 +4207,25 @@ def _function_body( tuple[FortranAssignment | FortranCall | FortranIf | FortranSelectCase, ...], tuple[FortranFunction, ...], ]: - """Build one native-call leaf plus linear optional-derived dispatch.""" + """Build one native-call leaf plus linear optional forwarding.""" result_name = self._native_direct_result_name(plan, result_name) - derived_optional = tuple( + optional = tuple( argument for argument in sorted(plan.arguments, key=lambda item: item.projected_call_slot.native_position) - if argument.derived_call is not None - and argument.entrypoint.optional_mode in {OptionalMode.NULLABLE_VALUE, OptionalMode.DESCRIPTOR} + if argument.entrypoint.optional_mode in {OptionalMode.NULLABLE_VALUE, OptionalMode.DESCRIPTOR} + ) + forwarded = tuple(argument for argument in optional if not self._requires_direct_optional_call(argument)) + if not forwarded: + return self._native_dispatch_body(plan, result_name), () + return ( + ( + *self._optional_descriptor_transport_initializers(forwarded), + FortranCall(self._optional_forwarding_step_name(0)), + ), + self._optional_forwarding_procedures(plan, forwarded, result_name), ) - if derived_optional: - forwarded = self._contained_optional_descriptor_arguments(plan) - procedures = self._derived_optional_dispatch_procedures( - plan, - derived_optional, - result_name, - forwarded, - ) - return (self._contained_optional_descriptor_call_tree(forwarded, 0, ()),), procedures - return self._ordinary_function_body(plan, result_name), () - def _ordinary_function_body( + def _native_dispatch_body( self, plan: FunctionPlan, result_name: str | None, @@ -4212,8 +4233,9 @@ def _ordinary_function_body( present: frozenset[str] = frozenset(), replacements: dict[str, str] | None = None, ) -> tuple[FortranAssignment | FortranCall | FortranIf | FortranSelectCase, ...]: - """Build the existing rank and non-derived optional call tree.""" + """Build rank or polymorphic dispatch around one native call site.""" replacements = dict(replacements or {}) + direct_optional = self._direct_call_optional_arguments(plan) polymorphic = self._polymorphic_arguments(plan) if polymorphic: return ( @@ -4224,6 +4246,7 @@ def _ordinary_function_body( present, result_name, replacements, + direct_optional, ), ) assumed_rank = self._assumed_rank_arguments(plan) @@ -4236,12 +4259,10 @@ def _ordinary_function_body( replacements, result_name, present=present, + direct_optional=direct_optional, ), ) - optional = self._non_derived_optional_arguments(plan) - if not optional: - return (self._native_invocation(plan, present, result_name, replacements),) - return (self._optional_call_tree(plan, optional, 0, present, result_name, replacements),) + return (self._direct_optional_call_tree(plan, direct_optional, 0, present, result_name, replacements),) @staticmethod def _polymorphic_arguments(plan: FunctionPlan) -> tuple[ArgumentTransferPlan, ...]: @@ -4263,56 +4284,6 @@ def _assumed_rank_arguments(plan: FunctionPlan) -> tuple[ArgumentTransferPlan, . and argument.array.entrypoint_abi is ArrayEntrypointABI.RAW_ADDRESS ) - @staticmethod - def _non_derived_optional_arguments( - plan: FunctionPlan, - ) -> tuple[ArgumentTransferPlan, ...]: - """Return optional arguments handled by the ordinary presence tree.""" - return tuple( - argument - for argument in sorted(plan.arguments, key=lambda item: item.projected_call_slot.native_position) - if argument.entrypoint.optional_mode in {OptionalMode.NULLABLE_VALUE, OptionalMode.DESCRIPTOR} - and argument.derived_call is None - ) - - def _contained_optional_descriptor_arguments( - self, - plan: FunctionPlan, - ) -> tuple[ArgumentTransferPlan, ...]: - """Return descriptor optionals that must not be host-associated on ifx.""" - return tuple( - argument - for argument in sorted(plan.arguments, key=lambda item: item.projected_call_slot.native_position) - if argument.derived_call is None - and argument.entrypoint.optional_mode in {OptionalMode.NULLABLE_VALUE, OptionalMode.DESCRIPTOR} - and argument.entrypoint.handoff_mode is ArgumentHandoffMode.ARRAY_BUFFER - and self._array_crosses_as_descriptor(argument) - ) - - def _contained_optional_descriptor_call_tree( - self, - arguments: tuple[ArgumentTransferPlan, ...], - index: int, - passed: tuple[CodeExpression, ...], - ) -> FortranCall | FortranIf: - """Enter the contained chain without forwarding an absent descriptor.""" - if index == len(arguments): - return FortranCall(self._derived_optional_step_name(0), passed) - argument = arguments[index] - local_name = self._forwarded_optional_descriptor_parameter_name(argument) - actual_name = argument.entrypoint.parameter_name - return FortranIf( - CodeExpression(f"present({actual_name})"), - body=( - self._contained_optional_descriptor_call_tree( - arguments, - index + 1, - (*passed, CodeExpression(f"{local_name}={actual_name}")), - ), - ), - else_body=(self._contained_optional_descriptor_call_tree(arguments, index + 1, passed),), - ) - def _polymorphic_call_tree( self, plan: FunctionPlan, @@ -4321,10 +4292,18 @@ def _polymorphic_call_tree( present: frozenset[str], result_name: str | None, replacements: dict[str, str], + direct_optional: tuple[ArgumentTransferPlan, ...], ) -> FortranAssignment | FortranCall | FortranIf | FortranSelectCase: """Dispatch N enumerated scalar inputs without speculative native calls.""" if index == len(arguments): - return self._native_invocation(plan, present, result_name, replacements) + return self._direct_optional_call_tree( + plan, + direct_optional, + 0, + present, + result_name, + replacements, + ) argument = arguments[index] cases = [] for variant in argument.polymorphic.variants: @@ -4340,6 +4319,7 @@ def _polymorphic_call_tree( present, result_name, replacements, + direct_optional, ), ), ) @@ -4354,59 +4334,58 @@ def _polymorphic_variant_name(argument: ArgumentTransferPlan, abi_code: int) -> """Name one bridge-local typed pointer from its stable plan code.""" return f"{argument.entrypoint.parameter_name}_polymorphic_{abi_code}" - def _derived_optional_dispatch_procedures( + def _optional_forwarding_procedures( self, plan: FunctionPlan, optional: tuple[ArgumentTransferPlan, ...], result_name: str | None, - forwarded: tuple[ArgumentTransferPlan, ...], ) -> tuple[FortranFunction, ...]: - """Propagate N optional derived dummies with O(N) adapter procedures.""" + """Propagate N optional dummies through O(N) contained procedures.""" procedures = [] - forwarded_parameters = tuple(self._forwarded_optional_descriptor_parameter(item) for item in forwarded) - forwarded_passed = tuple( - CodeExpression(self._forwarded_optional_descriptor_parameter_name(item)) for item in forwarded - ) for index, argument in enumerate(optional): carried = optional[:index] - parameters = (*forwarded_parameters, *(self._derived_optional_parameter(item) for item in carried)) - passed = ( - *forwarded_passed, - *(CodeExpression(self._derived_optional_parameter_name(item)) for item in carried), + parameters = tuple(self._optional_forwarding_parameter(plan, item) for item in carried) + passed = tuple( + CodeExpression( + f"{self._optional_forwarding_parameter_name(item)}={self._optional_forwarding_parameter_name(item)}" + ) + for item in carried ) - expression = CodeExpression(self._native_argument_expression(argument)) + expression = self._optional_forwarding_actual_expression(argument) procedures.append( FortranFunction( - name=self._derived_optional_step_name(index), + name=self._optional_forwarding_step_name(index), parameters=parameters, body=( FortranIf( - CodeExpression(self._presence_condition(argument)), + CodeExpression(self._optional_forwarding_presence_condition(argument)), body=( + *self._present_preparation(argument), FortranCall( - self._derived_optional_step_name(index + 1), - (*passed, expression), + self._optional_forwarding_step_name(index + 1), + ( + *passed, + CodeExpression( + f"{self._optional_forwarding_parameter_name(argument)}={expression.text}" + ), + ), ), ), - else_body=(FortranCall(self._derived_optional_step_name(index + 1), passed),), + else_body=(FortranCall(self._optional_forwarding_step_name(index + 1), passed),), ), ), is_subroutine=True, ) ) replacements = { - **{argument.owner_path: self._derived_optional_parameter_name(argument) for argument in optional}, - **{ - argument.owner_path: self._forwarded_optional_descriptor_parameter_name(argument) - for argument in forwarded - }, + argument.owner_path: self._optional_forwarding_parameter_name(argument) for argument in optional } present = frozenset(argument.owner_path for argument in optional) procedures.append( FortranFunction( - name=self._derived_optional_step_name(len(optional)), - parameters=(*forwarded_parameters, *(self._derived_optional_parameter(item) for item in optional)), - body=self._ordinary_function_body( + name=self._optional_forwarding_step_name(len(optional)), + parameters=tuple(self._optional_forwarding_parameter(plan, item) for item in optional), + body=self._native_dispatch_body( plan, result_name, present=present, @@ -4417,29 +4396,206 @@ def _derived_optional_dispatch_procedures( ) return tuple(procedures) - def _forwarded_optional_descriptor_parameter( + @staticmethod + def _requires_direct_optional_call(argument: ArgumentTransferPlan) -> bool: + """Keep non-forwardable Fortran descriptors on a direct call leaf.""" + if argument.entrypoint.optionality is EntrypointOptionalityAction.EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR: + return True + character = argument.bridge.character_local + return bool( + argument.mutates_native + and ( + (character is not None and character.deferred_length and character.descriptor_kind is not None) + or ( + argument.datatype_family is DatatypeFamily.STRING + and argument.native_array_handle is not None + and argument.array is not None + and argument.array.itemsize is None + ) + ) + ) + + def _direct_call_optional_arguments( + self, + plan: FunctionPlan, + ) -> tuple[ArgumentTransferPlan, ...]: + """Return optionals whose descriptor changes cannot cross a forwarding dummy.""" + return tuple( + argument + for argument in sorted(plan.arguments, key=lambda item: item.projected_call_slot.native_position) + if argument.entrypoint.optional_mode in {OptionalMode.NULLABLE_VALUE, OptionalMode.DESCRIPTOR} + and self._requires_direct_optional_call(argument) + ) + + def _direct_optional_call_tree( + self, + plan: FunctionPlan, + optional: tuple[ArgumentTransferPlan, ...], + index: int, + present: frozenset[str], + result_name: str | None, + replacements: dict[str, str], + ) -> FortranAssignment | FortranCall | FortranIf: + """Keep descriptor-changing character actuals direct for compiler correctness.""" + if index == len(optional): + return self._native_invocation(plan, present, result_name, replacements) + argument = optional[index] + return FortranIf( + CodeExpression(self._presence_condition(argument)), + body=( + *self._present_preparation(argument), + self._direct_optional_call_tree( + plan, + optional, + index + 1, + present | {argument.owner_path}, + result_name, + replacements, + ), + ), + else_body=( + self._direct_optional_call_tree( + plan, + optional, + index + 1, + present, + result_name, + replacements, + ), + ), + ) + + def _optional_descriptor_transport_initializers( self, + optional: tuple[ArgumentTransferPlan, ...], + ) -> tuple[FortranNullify | FortranIf, ...]: + """Capture optional entrypoint descriptors as ordinary local pointers.""" + nodes = [] + for argument in optional: + if not self._array_crosses_as_descriptor(argument): + continue + transport = self._optional_descriptor_transport_name(argument) + nodes.extend( + ( + FortranNullify(transport), + FortranIf( + CodeExpression(f"present({argument.entrypoint.parameter_name})"), + body=( + FortranPointerAssignment( + transport, + CodeExpression(argument.entrypoint.parameter_name), + ), + ), + ), + ) + ) + return tuple(nodes) + + def _optional_forwarding_presence_condition(self, argument: ArgumentTransferPlan) -> str: + """Read presence from the entrypoint transport selected for one optional.""" + if self._array_crosses_as_descriptor(argument): + return f"associated({self._optional_descriptor_transport_name(argument)})" + return self._presence_condition(argument) + + def _optional_forwarding_actual_expression(self, argument: ArgumentTransferPlan) -> CodeExpression: + """Return the prepared actual introduced by one forwarding step.""" + if self._array_crosses_as_descriptor(argument): + return CodeExpression(self._optional_descriptor_transport_name(argument)) + return CodeExpression(self._native_argument_expression(argument)) + + def _optional_forwarding_parameter( + self, + plan: FunctionPlan, argument: ArgumentTransferPlan, ) -> FortranParameter: - """Declare one optional descriptor passed into every contained step.""" - return self._array_descriptor_parameter( - argument, - self._forwarded_optional_descriptor_parameter_name(argument), - optional=True, + """Declare one optional dummy matching the prepared native actual.""" + name = self._optional_forwarding_parameter_name(argument) + if argument.callback is not None: + return FortranParameter( + name, + f"procedure({argument.callback.prototype.interface_symbol})", + ("optional",), + ) + if argument.derived_call is not None: + return self._derived_native_parameter(argument, name, optional=True) + if self._array_crosses_as_descriptor(argument): + return self._array_descriptor_parameter(argument, name, optional=True) + if parameter := self._optional_native_array_parameter(argument, name): + return parameter + if parameter := self._optional_character_descriptor_parameter(argument, name): + return parameter + return self._ordinary_optional_forwarding_parameter(plan, argument, name) + + def _optional_native_array_parameter( + self, + argument: ArgumentTransferPlan, + name: str, + ) -> FortranParameter | None: + """Declare one optional native-array descriptor forwarding dummy.""" + handle = argument.native_array_handle + if handle is None or argument.entrypoint.handoff_mode is not ArgumentHandoffMode.NATIVE_DESCRIPTOR: + return None + attribute = "allocatable" if handle.descriptor_kind is NativeArrayDescriptorKind.ALLOCATABLE else "pointer" + element_type = ( + self._native_array_argument_element_type(argument) + if argument.array is not None and argument.array.itemsize is None + else self._array_element_fortran_type(argument) + ) + return FortranParameter( + name, + element_type, + ("optional", attribute, self._array_dimension_attribute(handle.array.rank)), ) + def _optional_character_descriptor_parameter( + self, + argument: ArgumentTransferPlan, + name: str, + ) -> FortranParameter | None: + """Declare one optional scalar character descriptor forwarding dummy.""" + character_local = argument.bridge.character_local + if argument.object_kind is not ObjectKind.STRING or character_local is None: + return None + descriptor = character_local.descriptor_kind + if descriptor is None: + return None + if not character_local.deferred_length and argument.character_length is None: + raise ValueError(f"Fixed-length character optional {argument.owner_path!r} has no length") + length = ":" if character_local.deferred_length else str(argument.character_length) + attribute = "allocatable" if descriptor is NativeArrayDescriptorKind.ALLOCATABLE else "pointer" + return FortranParameter( + name, + f"character(kind=c_char, len={length})", + (attribute, "optional"), + ) + + def _ordinary_optional_forwarding_parameter( + self, + plan: FunctionPlan, + argument: ArgumentTransferPlan, + name: str, + ) -> FortranParameter: + """Add optionality to an ordinary prepared native actual.""" + parameter = self._external_interface_parameter(plan, argument, name=name) + attributes = list(parameter.attributes) + if "optional" not in attributes: + attributes.append("optional") + if argument.object_kind is ObjectKind.SCALAR and argument.entrypoint.optional_mode is OptionalMode.DESCRIPTOR: + attribute = "pointer" if argument.projected_call_slot.value_kind == "pointer" else "allocatable" + if attribute not in attributes: + attributes.append(attribute) + type_name = argument.scalar_native_type or parameter.type_name + return replace(parameter, type_name=type_name, attributes=tuple(attributes)) + @staticmethod - def _forwarded_optional_descriptor_parameter_name(argument: ArgumentTransferPlan) -> str: - """Name an optional descriptor local to the contained dispatch chain.""" + def _optional_forwarding_parameter_name(argument: ArgumentTransferPlan) -> str: + """Name one optional dummy carried through the contained chain.""" return f"prik_optional_{argument.entrypoint.parameter_name}" - def _derived_optional_parameter(self, argument: ArgumentTransferPlan) -> FortranParameter: - """Mirror the completed native dummy category and add OPTIONAL.""" - return self._derived_native_parameter( - argument, - self._derived_optional_parameter_name(argument), - optional=True, - ) + @staticmethod + def _optional_descriptor_transport_name(argument: ArgumentTransferPlan) -> str: + """Name the local pointer that carries an entrypoint descriptor safely.""" + return f"prik_optional_{argument.entrypoint.parameter_name}_transport" def _derived_native_parameter( self, @@ -4468,41 +4624,9 @@ def _derived_native_parameter( ) @staticmethod - def _derived_optional_parameter_name(argument: ArgumentTransferPlan) -> str: - """Return the local optional-presence parameter name for one derived argument.""" - return f"prik_optional_{argument.entrypoint.parameter_name}" - - @staticmethod - def _derived_optional_step_name(index: int) -> str: - """Return the deterministic nested-procedure name for one optional derived dispatch case.""" - return f"prik_derived_optional_step_{index}" - - def _optional_call_tree( - self, - plan: FunctionPlan, - optional: tuple[ArgumentTransferPlan, ...], - index: int, - present: frozenset[str], - result_name: str | None, - replacements: dict[str, str], - ) -> FortranAssignment | FortranCall | FortranIf: - """Return an exhaustive native-call tree for optional presence states.""" - if index == len(optional): - return self._native_invocation(plan, present, result_name, replacements) - argument = optional[index] - present_roles = present | {argument.owner_path} - return FortranIf( - condition=CodeExpression( - f"present({replacements[argument.owner_path]})" - if self._array_crosses_as_descriptor(argument) and argument.owner_path in replacements - else self._presence_condition(argument) - ), - body=( - *self._present_preparation(argument), - self._optional_call_tree(plan, optional, index + 1, present_roles, result_name, replacements), - ), - else_body=(self._optional_call_tree(plan, optional, index + 1, present, result_name, replacements),), - ) + def _optional_forwarding_step_name(index: int) -> str: + """Return the deterministic name for one linear optional-forwarding step.""" + return f"prik_optional_step_{index}" def _native_invocation( self, @@ -4705,18 +4829,18 @@ def _assumed_rank_call_tree( result_name: str | None, *, present: frozenset[str] = frozenset(), + direct_optional: tuple[ArgumentTransferPlan, ...] = (), ) -> FortranAssignment | FortranCall | FortranIf | FortranSelectCase: """Dispatch each runtime-rank array through explicit one-to-fifteen branches.""" if index == len(arguments): - optional = tuple( - argument - for argument in sorted(plan.arguments, key=lambda item: item.projected_call_slot.native_position) - if argument.entrypoint.optional_mode in {OptionalMode.NULLABLE_VALUE, OptionalMode.DESCRIPTOR} - and argument.derived_call is None - ) - if optional: - return self._optional_call_tree(plan, optional, 0, present, result_name, replacements) - return self._native_invocation(plan, present, result_name, replacements) + return self._direct_optional_call_tree( + plan, + direct_optional, + 0, + present, + result_name, + replacements, + ) argument = arguments[index] name = argument.entrypoint.parameter_name cases = [] @@ -4730,6 +4854,7 @@ def _assumed_rank_call_tree( replacements, result_name, present=present, + direct_optional=direct_optional, ) del replacements[argument.owner_path] cases.append( @@ -4797,6 +4922,8 @@ def _presence_condition(self, plan: ArgumentTransferPlan) -> str: if handle is not None and handle.handoff.abi is NativeDescriptorHandoffABI.FORTRAN_OWNER: return f"c_associated(bound_{name}_present)" if self._array_crosses_as_descriptor(plan): + if plan.entrypoint.pass_descriptor_presence: + return f"c_associated(bound_{name}_present)" # The dummy is the array itself, and C omits it by passing no # descriptor at all, so Fortran's own inquiry is the condition. return f"present({name})" @@ -4931,6 +5058,23 @@ def _optional_argument_declarations( argument: ArgumentTransferPlan, ) -> tuple[FortranDeclaration, ...]: """Return optional helper declarations for one completed handoff.""" + if argument.entrypoint.optional_mode in { + OptionalMode.NULLABLE_VALUE, + OptionalMode.DESCRIPTOR, + } and self._array_crosses_as_descriptor(argument): + array = argument.array + if array is None: + raise ValueError(f"Optional descriptor transport {argument.owner_path!r} has no array plan") + if array.rank is None: + return () + attributes = (self._array_dimension_attribute(array.rank), "pointer") + return ( + FortranDeclaration( + self._optional_descriptor_transport_name(argument), + self._array_element_fortran_type(argument), + attributes, + ), + ) if argument.callback is not None: return () handle = argument.native_array_handle diff --git a/prik/pipeline/wrapper.py b/prik/pipeline/wrapper.py index 90cb50b37..ff56a966c 100644 --- a/prik/pipeline/wrapper.py +++ b/prik/pipeline/wrapper.py @@ -63,6 +63,7 @@ DerivedWriteback, DeclarationCallableAction, DirectResultABI, + EntrypointOptionalityAction, EntrypointPassingConvention, LifecycleOperation, FIXED_STRING_RESULT_COPY_REASON, @@ -4352,10 +4353,32 @@ def _optional_presence_diagnostics( descriptor_mode = mode in {OptionalMode.REQUIRED_DESCRIPTOR, OptionalMode.DESCRIPTOR} if plan.binding.descriptor_boundary != descriptor_mode: diagnostics.append(self._diagnostic(plan.owner_path, "inconsistent-descriptor-boundary", mode.value)) - if mode is OptionalMode.DESCRIPTOR and plan.entrypoint.presence_role is None: + if plan.entrypoint.pass_descriptor_presence and plan.entrypoint.presence_role is None: diagnostics.append(self._diagnostic(plan.owner_path, "missing-descriptor-presence-role", mode.value)) - if mode is not OptionalMode.DESCRIPTOR and plan.entrypoint.presence_role is not None: + if not plan.entrypoint.pass_descriptor_presence and plan.entrypoint.presence_role is not None: diagnostics.append(self._diagnostic(plan.owner_path, "unexpected-descriptor-presence-role", mode.value)) + if plan.entrypoint.pass_descriptor_presence and plan.entrypoint.optionality not in { + EntrypointOptionalityAction.EXPLICIT_NATIVE_PRESENCE, + EntrypointOptionalityAction.EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR, + }: + diagnostics.append( + self._diagnostic( + plan.owner_path, + "inconsistent-explicit-descriptor-presence", + plan.entrypoint.optionality.value, + ) + ) + if plan.entrypoint.optionality is EntrypointOptionalityAction.EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR: + array = plan.array + if ( + not plan.entrypoint.pass_descriptor_presence + or array is None + or array.rank is not None + or array.entrypoint_abi is not ArrayEntrypointABI.C_DESCRIPTOR + ): + diagnostics.append( + self._diagnostic(plan.owner_path, "invalid-placeholder-descriptor-presence", mode.value) + ) return tuple(diagnostics) def _optional_native_diagnostics( diff --git a/prik/planning/planner.py b/prik/planning/planner.py index dfb6bc8df..a450fe125 100644 --- a/prik/planning/planner.py +++ b/prik/planning/planner.py @@ -2165,7 +2165,7 @@ def _argument_presence_role( """Return the explicit optional or descriptor presence handoff role.""" if native_array_handle is not None: return native_array_handle.handoff.presence_role - if policy.optional_mode is OptionalMode.DESCRIPTOR: + if policy.entrypoint_pass_descriptor_presence: return f"{policy.owner_path}:present" return None diff --git a/prik/policy/construction.py b/prik/policy/construction.py index 2ff497a01..f19a30db4 100644 --- a/prik/policy/construction.py +++ b/prik/policy/construction.py @@ -2014,6 +2014,15 @@ def _complete_entrypoint_argument_route( ) -> ArgumentPolicy: """Project a selected route into one argument's completed ABI metadata.""" uses_adapter = action is NativeEntrypointAction.GENERATED_FORTRAN_ADAPTER + placeholder_descriptor_presence = ( + uses_adapter + and argument.entrypoint_optionality is EntrypointOptionalityAction.NULL_C_DESCRIPTOR_POINTER + and argument.array is not None + and argument.array.rank is None + ) + explicit_descriptor_presence = ( + uses_adapter and argument.optional_mode is OptionalMode.DESCRIPTOR + ) or placeholder_descriptor_presence return replace( argument, entrypoint_pass_character_length=( @@ -2037,15 +2046,17 @@ def _complete_entrypoint_argument_route( and argument.array is not None and argument.array.entrypoint_abi is ArrayEntrypointABI.RAW_ADDRESS ), - entrypoint_pass_descriptor_presence=(uses_adapter and argument.optional_mode is OptionalMode.DESCRIPTOR), + entrypoint_pass_descriptor_presence=explicit_descriptor_presence, entrypoint_pass_derived_transaction=(uses_adapter and argument.derived_call is not None), entrypoint_pass_callback_parameter=( argument.callback is not None and (action is NativeEntrypointAction.DIRECT_C_ABI or argument.optional_mode is OptionalMode.NULLABLE_VALUE) ), entrypoint_optionality=( - EntrypointOptionalityAction.EXPLICIT_NATIVE_PRESENCE - if uses_adapter and argument.optional_mode is OptionalMode.DESCRIPTOR + EntrypointOptionalityAction.EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR + if placeholder_descriptor_presence + else EntrypointOptionalityAction.EXPLICIT_NATIVE_PRESENCE + if explicit_descriptor_presence else argument.entrypoint_optionality ), ) @@ -5150,9 +5161,6 @@ def _array_storage_boundary_blockers( blockers.append(f"argument {argument.name!r} ordinary array must be non-descriptor storage") if decision.nullable and not argument.optional: blockers.append(f"argument {argument.name!r} ordinary array is nullable without optional presence") - array_policy = _array_handoff_policy(argument.semantic_type) - if argument.optional and array_policy is not None and array_policy.rank is None: - blockers.append(f"argument {argument.name!r} optional assumed-rank combination is not supported") return tuple(blockers) diff --git a/prik/policy/models.py b/prik/policy/models.py index c9ce72bda..41d681ae2 100644 --- a/prik/policy/models.py +++ b/prik/policy/models.py @@ -107,6 +107,7 @@ class EntrypointOptionalityAction(str, Enum): NULL_POINTER = "null_pointer" NULL_C_DESCRIPTOR_POINTER = "null_c_descriptor_pointer" EXPLICIT_NATIVE_PRESENCE = "explicit_native_presence" + EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR = "explicit_presence_with_placeholder_descriptor" ADAPTER_SIDE_FORTRAN_OMISSION = "adapter_side_fortran_omission" BLOCKED = "blocked" diff --git a/tests/fortran/arrays/codegen/test_specialized_array_roles.py b/tests/fortran/arrays/codegen/test_specialized_array_roles.py index 639227d3f..ab45e7cbe 100644 --- a/tests/fortran/arrays/codegen/test_specialized_array_roles.py +++ b/tests/fortran/arrays/codegen/test_specialized_array_roles.py @@ -7,6 +7,7 @@ from prik.policy.completion import complete_semantic_policies from prik.policy.models import ( ArrayEntrypointABI, + EntrypointOptionalityAction, EntrypointPassingConvention, NativeArraySourceKind, OptionalMode, @@ -22,6 +23,7 @@ def _later_array_plan(): from prik.contracts import Float64, String def optional(values: Float64[:] = ...) -> None: ... +def optional_any_rank(values: Float64[...] = ...) -> None: ... def any_rank(values: Float64[...]) -> Float64: ... def labels(values: String[8][:]) -> None: ... def labels_any_width(values: String[...][:]) -> None: ... @@ -49,6 +51,8 @@ def hidden_labels() -> String[4][2]: ... def test_optional_assumed_rank_and_character_arrays_have_explicit_distinct_roles(): functions = {function.binding.python_name: function for function in _later_array_plan().namespaces[0].functions} optional = functions["optional"].arguments[0] + optional_assumed_argument = functions["optional_any_rank"].arguments[0] + optional_assumed = optional_assumed_argument.array assumed_argument = functions["any_rank"].arguments[0] assumed = assumed_argument.array character_argument = functions["labels"].arguments[0] @@ -64,6 +68,17 @@ def test_optional_assumed_rank_and_character_arrays_have_explicit_distinct_roles assert optional.entrypoint.optional_mode is OptionalMode.NULLABLE_VALUE assert optional.native_array_actual is not None assert optional.native_array_actual.accepted_sources == handle_sources + assert optional_assumed_argument.binding.optional_mode is OptionalMode.NULLABLE_VALUE + assert optional_assumed_argument.entrypoint.optional_mode is OptionalMode.NULLABLE_VALUE + assert ( + optional_assumed_argument.entrypoint.optionality + is EntrypointOptionalityAction.EXPLICIT_PRESENCE_WITH_PLACEHOLDER_DESCRIPTOR + ) + assert optional_assumed_argument.entrypoint.presence_role is not None + assert optional_assumed is not None + assert optional_assumed.rank is None + assert optional_assumed.entrypoint_abi is ArrayEntrypointABI.C_DESCRIPTOR + assert optional_assumed_argument.entrypoint.passing is EntrypointPassingConvention.C_DESCRIPTOR_POINTER assert assumed is not None assert assumed.rank is None assert assumed.contiguous is False @@ -105,12 +120,27 @@ def test_optional_assumed_rank_and_character_lowering_follow_named_plan_fields() ) in c_source assert "NPY_FLOAT64, 1, 15, PRIK_ARRAY_LAYOUT_SIGNED_STRIDED_F" in c_source assert "bound_values_rank = (int64_t)PyArray_NDIM" in c_source + assert "void bind_c_optional_any_rank(CFI_cdesc_t * values, void * values_present);" in c_source + assert "bound_values_present = bound_values_obj != Py_None ? (void *)bound_values_obj : NULL;" in c_source + assert ( + "CFI_establish((CFI_cdesc_t *)&bound_values_section, &bound_values_placeholder, " + "CFI_attribute_other, CFI_type_double, sizeof(double), 0, NULL)" + ) in c_source # Runtime character width is part of the raw bridge ABI. The shared binder # returns it for either a NumPy array or a native handle. assert "bound_values_itemsize" in c_source assert "&bound_values_itemsize, CFI_type_char" in c_source assert "real(c_double), dimension(..) :: values" in bridge_source assert "select case (values_rank)" not in bridge_source + optional_any_rank = bridge_source.split("subroutine bind_c_optional_any_rank", maxsplit=1)[1].split( + "end subroutine bind_c_optional_any_rank", maxsplit=1 + )[0] + assert "real(c_double), dimension(..) :: values" in optional_any_rank + assert "type(c_ptr), value :: bound_values_present" in optional_any_rank + assert "if (c_associated(bound_values_present)) then" in optional_any_rank + assert "call native_optional_any_rank(values=values)" in optional_any_rank + assert "call native_optional_any_rank()" in optional_any_rank + assert "prik_optional_values_transport" not in optional_any_rank assert "character(kind=c_char, len=8), pointer, contiguous, dimension(:) :: values" in bridge_source assert max(map(len, bridge_source.splitlines())) <= 132 diff --git a/tests/fortran/arrays/end_to_end/fixtures/contracts/fassumed_rank_f90/fassumed_rank_f90.pyi b/tests/fortran/arrays/end_to_end/fixtures/contracts/fassumed_rank_f90/fassumed_rank_f90.pyi index 07b1c4045..3ef37ba52 100644 --- a/tests/fortran/arrays/end_to_end/fixtures/contracts/fassumed_rank_f90/fassumed_rank_f90.pyi +++ b/tests/fortran/arrays/end_to_end/fixtures/contracts/fassumed_rank_f90/fassumed_rank_f90.pyi @@ -1,4 +1,12 @@ -from prik.contracts import Float64, Int32 +from prik.contracts import Allocatable, Annotated, Float64, Int32, Pointer, PointerAssociation + +def allocatable_handle() -> Allocatable[Float64[:]]: ... + +def pointer_handle() -> Annotated[Pointer[Float64[:]], PointerAssociation("runtime")]: ... + +def optional_rank( + values: Float64[...] = ... +) -> Int32: ... def rank_weighted_sum( values: Float64[...] @@ -13,4 +21,11 @@ def rank_pair_score( right: Float64[...] ) -> Int32: ... -__all__ = ["rank_weighted_sum", "bump_assumed_rank", "rank_pair_score"] +__all__ = [ + "allocatable_handle", + "pointer_handle", + "optional_rank", + "rank_weighted_sum", + "bump_assumed_rank", + "rank_pair_score", +] diff --git a/tests/fortran/arrays/end_to_end/fixtures/native/fassumed_rank_f90.f90 b/tests/fortran/arrays/end_to_end/fixtures/native/fassumed_rank_f90.f90 index 4732c4ae2..198b25975 100644 --- a/tests/fortran/arrays/end_to_end/fixtures/native/fassumed_rank_f90.f90 +++ b/tests/fortran/arrays/end_to_end/fixtures/native/fassumed_rank_f90.f90 @@ -1,6 +1,31 @@ module fassumed_rank_f90 + real(8), target :: pointer_values(2) = [3.0_8, 4.0_8] + private :: pointer_values contains + function allocatable_handle() result(values) + real(8), allocatable :: values(:) + + allocate(values(2)) + values = [1.0_8, 2.0_8] + end function allocatable_handle + + function pointer_handle() result(values) + real(8), pointer :: values(:) + + values => pointer_values + end function pointer_handle + + integer function optional_rank(values) result(observed) + real(8), intent(in), optional :: values(..) + + if (present(values)) then + observed = rank(values) + else + observed = -1 + end if + end function optional_rank + real(8) function rank_weighted_sum(values) result(total) real(8), intent(in) :: values(..) diff --git a/tests/fortran/arrays/end_to_end/test_assumed_rank_arrays.py b/tests/fortran/arrays/end_to_end/test_assumed_rank_arrays.py index 6360e2366..0e2170b2d 100644 --- a/tests/fortran/arrays/end_to_end/test_assumed_rank_arrays.py +++ b/tests/fortran/arrays/end_to_end/test_assumed_rank_arrays.py @@ -57,6 +57,26 @@ def test_assumed_rank_arguments_dispatch_to_runtime_rank( module.rank_weighted_sum(rank16) +def test_optional_assumed_rank_preserves_fortran_presence(assumed_rank_module): + module = assumed_rank_module + + assert module.optional_rank() == -1 + assert module.optional_rank(None) == -1 + assert module.optional_rank(np.ones(3, dtype=np.float64)) == 1 + assert module.optional_rank(np.ones((2, 3), dtype=np.float64, order="F")) == 2 + + allocatable = module.allocatable_handle() + pointer = module.pointer_handle() + try: + assert allocatable.allocated is True + assert pointer.associated is True + assert module.optional_rank(allocatable) == 1 + assert module.optional_rank(pointer) == 1 + finally: + allocatable.close() + pointer.close() + + def test_assumed_rank_bridge_dispatches_each_runtime_rank_argument( assumed_rank_module, ): diff --git a/tests/fortran/callbacks/codegen/test_callback_planning.py b/tests/fortran/callbacks/codegen/test_callback_planning.py index 58d59255d..866b3003a 100644 --- a/tests/fortran/callbacks/codegen/test_callback_planning.py +++ b/tests/fortran/callbacks/codegen/test_callback_planning.py @@ -278,8 +278,9 @@ def test_optional_callback_uses_the_ordinary_presence_plan(): c_source, bridge = _sources(plan) assert "bound_callback_obj != Py_None ? prik_callback_trampoline_" in c_source assert "if (c_associated(callback)) then" in bridge - assert "native_apply_value_callback(callback=prik_callback_adapter_" in bridge - assert "native_apply_value_callback(value=value)" in bridge + assert "procedure(prik_value_callback_" in bridge + assert "callback=prik_optional_callback" in bridge + assert bridge.count("result = native_apply_value_callback(") == 1 def test_direct_bind_c_callback_generates_no_fortran_callback_adapter(): diff --git a/tests/fortran/infrastructure/semantic_pyi/pipeline/test_calls_and_policy_metadata.py b/tests/fortran/infrastructure/semantic_pyi/pipeline/test_calls_and_policy_metadata.py index 429cb76ae..636569d3f 100644 --- a/tests/fortran/infrastructure/semantic_pyi/pipeline/test_calls_and_policy_metadata.py +++ b/tests/fortran/infrastructure/semantic_pyi/pipeline/test_calls_and_policy_metadata.py @@ -444,10 +444,9 @@ def update(scale: Float64 | None = ..., target: Float64 | None = ...) -> None: . assert "bound_target_present" in bridge_source assert "if (c_associated(bound_scale_present)) then" in bridge_source assert "if (c_associated(bound_target_present)) then" in bridge_source - assert "call native_update()" in bridge_source - assert "call native_update(scale=scale_descriptor)" in bridge_source - assert "call native_update(target=target_descriptor)" in bridge_source - assert "call native_update(scale=scale_descriptor, target=target_descriptor" in bridge_source + assert "scale=prik_optional_scale" in bridge_source + assert "target=prik_optional_target" in bridge_source + assert bridge_source.count("call native_update(") == 1 assert "bound_scale_obj = NULL;" in c_wrapper assert "bound_target_obj = NULL;" in c_wrapper diff --git a/tests/fortran/optional_arguments/codegen/test_optional_lowering.py b/tests/fortran/optional_arguments/codegen/test_optional_lowering.py index 4f294cfab..9eb21411c 100644 --- a/tests/fortran/optional_arguments/codegen/test_optional_lowering.py +++ b/tests/fortran/optional_arguments/codegen/test_optional_lowering.py @@ -55,8 +55,8 @@ def test_optional_scalar_lowering_distinguishes_absent_or_none_from_value(): assert "bound_factor_nullable = &bound_factor;" in c_source assert "bind_c_optional_scale(base, bound_factor)" in fortran_source assert "if (c_associated(bound_factor)) then" in fortran_source - assert "result = optional_scale(base=base, factor=factor)" in fortran_source - assert "result = optional_scale(base=base)" in fortran_source + assert "result = optional_scale(base=base, factor=prik_optional_factor)" in fortran_source + assert fortran_source.count("result = optional_scale(") == 1 def test_optional_descriptor_lowering_records_presence_and_nullable_value_handoffs(): @@ -86,8 +86,8 @@ def alloc_state(value: Annotated[Float64, Immutable] | None = ...) -> Int32: ... assert "bind_c_alloc_state(bound_value_nullable, bound_value_present)" in c_source assert "type(c_ptr), value :: bound_value_present" in fortran_source assert "if (c_associated(bound_value_present)) then" in fortran_source - assert "result = native_alloc_state(value=value_descriptor)" in fortran_source - assert "result = native_alloc_state()" in fortran_source + assert "result = native_alloc_state(value=prik_optional_value)" in fortran_source + assert fortran_source.count("result = native_alloc_state(") == 1 def test_optional_arguments_with_hidden_literals_materialize_the_literal_in_the_binding(): @@ -105,10 +105,10 @@ def optional_literal(value: Annotated[Float64, Immutable] | None = ...) -> Float assert "double bind_c_optional_literal(int32_t literal_0, double * value);" in c_source assert "bind_c_optional_literal(1, bound_value_nullable);" in c_source assert "function bind_c_optional_literal(literal_0, bound_value)" in fortran_source - assert "native_optional_literal(literal_0, value=value)" in fortran_source + assert "native_optional_literal(literal_0, value=prik_optional_value)" in fortran_source -def test_optional_descriptor_is_passed_into_contained_derived_dispatch(): +def test_optional_descriptor_is_forwarded_explicitly_into_the_native_call(): """A contained procedure receives, rather than host-associates, the descriptor.""" module = pyi_file_to_semantic_module(OPTIONAL_MIXED_CONTRACT, module_name="foptional_f90") fortran_source = _source(_artifacts(module), ".f90") @@ -117,11 +117,33 @@ def test_optional_descriptor_is_passed_into_contained_derived_dispatch(): )[0] contained = summarize.split(" contains", maxsplit=1)[1] - assert "if (present(values)) then" in fortran_source - assert "call prik_derived_optional_step_0(prik_optional_values=values)" in fortran_source + assert "call prik_optional_step_0()" in summarize + assert "if (present(values)) then" in summarize.split(" contains", maxsplit=1)[0] + assert "prik_optional_values_transport => values" in summarize assert "real(c_double), dimension(:), optional :: prik_optional_values" in contained - assert "if (present(prik_optional_values)) then" in contained + assert "prik_optional_values=prik_optional_values" in contained assert "present(values)" not in contained + assert contained.count("result = native_summarize(") == 1 + + +def test_many_optional_scalars_generate_one_native_call_site(): + """Forwardable optionals use linear procedures and converge on one call site.""" + argument_count = 24 + arguments = ",\n ".join( + f"value_{index}: Annotated[Int32, Immutable] | None = ..." for index in range(argument_count) + ) + module = parse_pyi_text( + f""" +def many_optional( + {arguments}, +) -> Int32: ... +""", + module_name="many_optional", + ) + fortran_source = _source(_artifacts(module), ".f90") + + assert fortran_source.count("result = native_many_optional(") == 1 + assert "value_23=prik_optional_value_23" in fortran_source def test_required_descriptor_keeps_python_presence_separate_from_native_state_and_copyout(): diff --git a/tests/fortran/optional_arguments/end_to_end/fixtures/contracts/foptional_f90/foptional_f90.pyi b/tests/fortran/optional_arguments/end_to_end/fixtures/contracts/foptional_f90/foptional_f90.pyi index d226df026..1b440da59 100644 --- a/tests/fortran/optional_arguments/end_to_end/fixtures/contracts/foptional_f90/foptional_f90.pyi +++ b/tests/fortran/optional_arguments/end_to_end/fixtures/contracts/foptional_f90/foptional_f90.pyi @@ -36,4 +36,11 @@ def optional_status( status: Int32[()] = ... ) -> tuple[Int32, Returns["status", Int32[()]] | None]: ... -__all__ = ["Sample", "summarize", "mutate_optional", "fill_optional", "optional_status"] +@native_call([Addr(Arg(0)), Addr(Arg(1)), Addr(Arg(2))]) +def three_optional( + first: Int32 = ..., + second: Int32 = ..., + third: Int32 = ... +) -> Int32: ... + +__all__ = ["Sample", "summarize", "mutate_optional", "fill_optional", "optional_status", "three_optional"] diff --git a/tests/fortran/optional_arguments/end_to_end/fixtures/native/foptional_f90.f90 b/tests/fortran/optional_arguments/end_to_end/fixtures/native/foptional_f90.f90 index 50ab1e2c4..c7457657d 100644 --- a/tests/fortran/optional_arguments/end_to_end/fixtures/native/foptional_f90.f90 +++ b/tests/fortran/optional_arguments/end_to_end/fixtures/native/foptional_f90.f90 @@ -53,4 +53,13 @@ integer function optional_status(base, status) optional_status = base if (present(status)) status = base + 50 end function optional_status + + integer function three_optional(first, second, third) + integer, intent(in), optional :: first, second, third + + three_optional = 0 + if (present(first)) three_optional = three_optional + 1 + if (present(second)) three_optional = three_optional + 2 + if (present(third)) three_optional = three_optional + 4 + end function three_optional end module foptional_f90 diff --git a/tests/fortran/optional_arguments/end_to_end/test_optional_runtime.py b/tests/fortran/optional_arguments/end_to_end/test_optional_runtime.py index d1b1b5134..45dfe08c6 100644 --- a/tests/fortran/optional_arguments/end_to_end/test_optional_runtime.py +++ b/tests/fortran/optional_arguments/end_to_end/test_optional_runtime.py @@ -121,6 +121,10 @@ def test_optional_arguments_drive_fortran_present_behavior( assert module.summarize(np.int32(5), item=item, values=values, label="abc") == np.int32(21) assert module.summarize(np.int32(5), None, values=values, item=item) == np.int32(18) + assert module.three_optional() == np.int32(0) + assert module.three_optional(None, np.int32(20)) == np.int32(2) + assert module.three_optional(np.int32(10), np.int32(20), np.int32(30)) == np.int32(7) + mutable = np.array([1.0, 2.0], dtype=np.float64) assert module.mutate_optional() is None assert module.mutate_optional(None, np.float64(100.0)) is None diff --git a/tests/fortran/strings/codegen/test_fixed_string_writeback.py b/tests/fortran/strings/codegen/test_fixed_string_writeback.py index 166e308ba..a9db89381 100644 --- a/tests/fortran/strings/codegen/test_fixed_string_writeback.py +++ b/tests/fortran/strings/codegen/test_fixed_string_writeback.py @@ -211,8 +211,8 @@ def optional_identity(label: String = ...) -> None: ... assert "character(kind=c_char, len=label_length), pointer :: label" in bridge_source assert "if (c_associated(bound_label)) then" in bridge_source assert "call c_f_pointer(bound_label, label)" in bridge_source - assert "call native_optional(label=label)" in bridge_source - assert "call native_optional()" in bridge_source + assert "call native_optional(label=prik_optional_label)" in bridge_source + assert bridge_source.count("call native_optional(") == 1 # A mutating callee wrote the binding's bytes, so nothing is copied back. assert "label_bytes" not in bridge_source assert "transfer(" not in bridge_source diff --git a/tests/fortran/subroutines/codegen/test_hidden_scalar_outputs.py b/tests/fortran/subroutines/codegen/test_hidden_scalar_outputs.py index 4a8e441ee..16477e2b1 100644 --- a/tests/fortran/subroutines/codegen/test_hidden_scalar_outputs.py +++ b/tests/fortran/subroutines/codegen/test_hidden_scalar_outputs.py @@ -63,5 +63,5 @@ def scale( assert "real(c_double) :: x" in native_interface assert "real(c_double) :: result" in native_interface assert "integer(c_int32_t), optional :: mode" in native_interface - assert "call SCALE_OUT(x=x, result=result, mode=mode)" in fortran_source - assert "call SCALE_OUT(x=x, result=result)" in fortran_source + assert "call SCALE_OUT(x=x, result=result, mode=prik_optional_mode)" in fortran_source + assert fortran_source.count("call SCALE_OUT(") == 1