Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
98 changes: 69 additions & 29 deletions extensions/variants.go
Original file line number Diff line number Diff line change
Expand Up @@ -159,6 +159,12 @@ func EvaluateTypeExpression(urn string, nullHandling NullabilityHandling, return
return outType.WithNullability(types.NullabilityNullable), nil
}

if nullHandling == DeclaredOutputNullability {
if anyReturn, ok := returnTypeExpr.(*types.AnyType); ok && anyReturn.Nullability == types.NullabilityRequired {
return outType.WithNullability(types.NullabilityRequired), nil
}
}

return outType, nil
}

Expand Down Expand Up @@ -190,7 +196,7 @@ func matchArguments(nullability NullabilityHandling, paramTypeList FuncParameter
return types.AreSyncTypeParametersMatching(funcDefArgList, actualTypes), nil
}

if err := ValidateConstrainedAnyTypeConsistency(funcDefArgList, actualTypes, variadicBehavior); err != nil {
if err := validateConstrainedAnyTypeConsistency(nullability, funcDefArgList, actualTypes, variadicBehavior); err != nil {
return false, err
}
return true, nil
Expand Down Expand Up @@ -675,30 +681,49 @@ func getLeafParameterizedParams(abstractTypes interface{}) []string {
panic("invalid non-leaf, non-parameterized type param")
}

// validateAnyTypeBinding validates and records the binding of an AnyType parameter to a concrete type.
// It recursively handles nested types (lists, maps, structs).
func validateAnyTypeBinding(paramType types.FuncDefArgType, argType types.Type, bindings map[string]types.Type) error {
// bindAnyType binds a constrained AnyType name to a concrete argument type.
// When ignoreNullability is true, the binding uses the required version of the argument type.
func bindAnyType(anyType *types.AnyType, argType types.Type, bindings map[string]types.Type, ignoreNullability bool) error {
// Plain "any" is unconstrained - each occurrence can be a different type.
if anyType.Name == "any" {
return nil
}

boundType := argType
if boundType.GetNullability() == types.NullabilityUnspecified {
boundType = boundType.WithNullability(types.NullabilityRequired)
}
if ignoreNullability {
boundType = boundType.WithNullability(types.NullabilityRequired)
}

if existingType, exists := bindings[anyType.Name]; exists {
if !existingType.Equals(boundType) {
return fmt.Errorf("%w: type parameter %s cannot be both %s and %s",
substraitgo.ErrInvalidType, anyType.Name,
existingType.ShortString(), boundType.ShortString())
}
} else {
bindings[anyType.Name] = boundType
}
return nil
}

// validateTopLevelAnyTypeBinding handles top-level argument bindings, where MIRROR and DECLARED_OUTPUT ignore outer argument nullability.
func validateTopLevelAnyTypeBinding(nullHandling NullabilityHandling, paramType types.FuncDefArgType, argType types.Type, bindings map[string]types.Type) error {
if anyType, ok := paramType.(*types.AnyType); ok {
return bindAnyType(anyType, argType, bindings,
nullHandling != DiscreteNullability || anyType.Nullability != types.NullabilityRequired)
}
return validateAnyTypeBinding(nullHandling, paramType, argType, bindings)
}

// validateAnyTypeBinding validates nested AnyType bindings, preserving nullability unless the nested AnyType is declared nullable.
// It recursively handles nested types (lists, maps, structs, and functions).
func validateAnyTypeBinding(nullHandling NullabilityHandling, paramType types.FuncDefArgType, argType types.Type, bindings map[string]types.Type) error {
switch p := paramType.(type) {
case *types.AnyType:
// Plain "any" is unconstrained - each occurrence can be a different type
if p.Name == "any" {
return nil
}
if existingType, exists := bindings[p.Name]; exists {
// Compare base types ignoring nullability. Nullability is enforced separately
// at the individual argument level by MatchWithNullability/MatchWithoutNullability
// (see matchArgumentAtCommon in variants.go). Here we only validate that all uses
// of the same type parameter (e.g., any1) resolve to the same base type.
existingBase := existingType.WithNullability(types.NullabilityRequired)
argBase := argType.WithNullability(types.NullabilityRequired)
if !existingBase.Equals(argBase) {
return fmt.Errorf("%w: type parameter %s cannot be both %s and %s",
substraitgo.ErrInvalidType, p.Name,
existingType.ShortString(), argType.ShortString())
}
} else {
bindings[p.Name] = argType
}
return bindAnyType(p, argType, bindings, p.Nullability != types.NullabilityRequired)
case *types.ParameterizedListType:
if listType, ok := argType.(*types.ListType); ok {
params := listType.GetParameters()
Expand All @@ -707,7 +732,7 @@ func validateAnyTypeBinding(paramType types.FuncDefArgType, argType types.Type,
if !ok {
return fmt.Errorf("%w: invalid list element type", substraitgo.ErrInvalidType)
}
if err := validateAnyTypeBinding(p.Type, elementType, bindings); err != nil {
if err := validateAnyTypeBinding(nullHandling, p.Type, elementType, bindings); err != nil {
return err
}
}
Expand All @@ -724,10 +749,10 @@ func validateAnyTypeBinding(paramType types.FuncDefArgType, argType types.Type,
if !ok {
return fmt.Errorf("%w: invalid map value type", substraitgo.ErrInvalidType)
}
if err := validateAnyTypeBinding(p.Key, keyType, bindings); err != nil {
if err := validateAnyTypeBinding(nullHandling, p.Key, keyType, bindings); err != nil {
return err
}
if err := validateAnyTypeBinding(p.Value, valueType, bindings); err != nil {
if err := validateAnyTypeBinding(nullHandling, p.Value, valueType, bindings); err != nil {
return err
}
}
Expand All @@ -742,20 +767,35 @@ func validateAnyTypeBinding(paramType types.FuncDefArgType, argType types.Type,
if !ok {
return fmt.Errorf("%w: invalid struct field type at position %d", substraitgo.ErrInvalidType, i)
}
if err := validateAnyTypeBinding(fieldType, elementType, bindings); err != nil {
if err := validateAnyTypeBinding(nullHandling, fieldType, elementType, bindings); err != nil {
return err
}
}
}
}
}
case *types.ParameterizedFuncType:
if funcType, ok := argType.(*types.FuncType); ok && len(funcType.ParameterTypes) == len(p.Parameters) {
for i, parameterType := range p.Parameters {
if err := validateAnyTypeBinding(nullHandling, parameterType, funcType.ParameterTypes[i], bindings); err != nil {
return err
}
}
if err := validateAnyTypeBinding(nullHandling, p.Return, funcType.ReturnType, bindings); err != nil {
return err
}
}
}
return nil
}

// ValidateConstrainedAnyTypeConsistency validates that all uses of the same AnyN parameter
// (e.g., any1, any2, etc.) resolve to the same concrete type across all arguments
// (e.g., any1, any2, etc.) resolve to the same concrete type across all arguments.
func ValidateConstrainedAnyTypeConsistency(funcParameters []types.FuncDefArgType, argumentTypes []types.Type, variadicBehavior *VariadicBehavior) error {
return validateConstrainedAnyTypeConsistency(MirrorNullability, funcParameters, argumentTypes, variadicBehavior)
}

func validateConstrainedAnyTypeConsistency(nullHandling NullabilityHandling, funcParameters []types.FuncDefArgType, argumentTypes []types.Type, variadicBehavior *VariadicBehavior) error {
// For variadic functions, expand the parameter list to match actual arguments
expandedFuncParameters := funcParameters
if variadicBehavior != nil && len(argumentTypes) > len(funcParameters) {
Expand All @@ -774,7 +814,7 @@ func ValidateConstrainedAnyTypeConsistency(funcParameters []types.FuncDefArgType

// Validate each argument against its parameter type
for i := 0; i < len(expandedFuncParameters) && i < len(argumentTypes); i++ {
if err := validateAnyTypeBinding(expandedFuncParameters[i], argumentTypes[i], bindings); err != nil {
if err := validateTopLevelAnyTypeBinding(nullHandling, expandedFuncParameters[i], argumentTypes[i], bindings); err != nil {
return err
}
}
Expand Down
125 changes: 125 additions & 0 deletions extensions/variants_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -442,6 +442,92 @@ func valArg(typ types.FuncDefArgType) extensions.ValueArg {
return extensions.ValueArg{Value: &parser.TypeExpression{ValueType: typ}}
}

func mustParseFuncDefArg(t *testing.T, typeExpression string) types.FuncDefArgType {
t.Helper()

typ, err := parser.ParseType(typeExpression)
require.NoError(t, err)
return typ
}

func mustParseConcreteType(t *testing.T, typeExpression string) types.Type {
t.Helper()

typ, err := mustParseFuncDefArg(t, typeExpression).ReturnType(nil, nil)
require.NoError(t, err)
return typ
}

func TestNullabilityAndAnyTypeBindingDocExamples(t *testing.T) {
tests := []struct {
definition string
nullabilityHandling extensions.NullabilityHandling
argTypes []string
returnType string
matches bool
expectedReturnType string
}{
{"f(any1, any1) -> any1", extensions.MirrorNullability, []string{"i32", "i32"}, "any1", true, "i32"},
{"f(any1, any1) -> any1", extensions.MirrorNullability, []string{"i32?", "i32"}, "any1", true, "i32?"},
{"f(any1, any1) -> any1", extensions.MirrorNullability, []string{"i32", "i32?"}, "any1", true, "i32?"},
{"f(any1, any1) -> any1", extensions.MirrorNullability, []string{"i32?", "i32?"}, "any1", true, "i32?"},
{"h(list<any1>, list<any1>) -> list<any1>", extensions.MirrorNullability, []string{"list<i32>", "list<i32>"}, "list<any1>", true, "list<i32>"},
{"h(list<any1>, list<any1>) -> list<any1>", extensions.MirrorNullability, []string{"list?<i32>", "list<i32>"}, "list<any1>", true, "list?<i32>"},
{"h(list<any1>, list<any1>) -> list<any1>", extensions.MirrorNullability, []string{"list<i32?>", "list<i32?>"}, "list<any1>", true, "list<i32?>"},
{"h(list<any1>, list<any1>) -> list<any1>", extensions.MirrorNullability, []string{"list<i32>", "list<i32?>"}, "list<any1>", false, ""},
{"k(func<any1 -> any1>) -> any1", extensions.MirrorNullability, []string{"func<i32 -> i32>"}, "any1", true, "i32"},
{"k(func<any1 -> any1>) -> any1", extensions.MirrorNullability, []string{"func<i32 -> fp64>"}, "any1", false, ""},
{"j(any1, list<any1?>) -> any1", extensions.MirrorNullability, []string{"i32", "list<i32?>"}, "any1", true, "i32"},
{"j(any1, list<any1?>) -> any1", extensions.MirrorNullability, []string{"i32", "list<i32>"}, "any1", false, ""},
{"j(any1, list<any1?>) -> any1", extensions.MirrorNullability, []string{"i32", "list<fp64?>"}, "any1", false, ""},
{"d(any1, any1) -> any1?", extensions.DeclaredOutputNullability, []string{"i32", "i32"}, "any1?", true, "i32?"},
{"d(any1, any1) -> any1?", extensions.DeclaredOutputNullability, []string{"i32?", "i32"}, "any1?", true, "i32?"},
{"d(any1, any1) -> any1?", extensions.DeclaredOutputNullability, []string{"i32?", "i32?"}, "any1?", true, "i32?"},
{"d(any1, any1) -> any1", extensions.DeclaredOutputNullability, []string{"i32", "i32"}, "any1", true, "i32"},
{"d(any1, any1) -> any1", extensions.DeclaredOutputNullability, []string{"i32?", "i32"}, "any1", true, "i32"},
{"d(any1, any1) -> any1", extensions.DeclaredOutputNullability, []string{"i32?", "i32?"}, "any1", true, "i32"},
{"g(any1, any1) -> any1", extensions.DiscreteNullability, []string{"i32", "i32"}, "any1", true, "i32"},
{"g(any1, any1) -> any1", extensions.DiscreteNullability, []string{"i32?", "i32?"}, "any1", false, ""},
{"g(any1, any1) -> any1", extensions.DiscreteNullability, []string{"i32", "i32?"}, "any1", false, ""},
{"g2(any1, any1?) -> any1?", extensions.DiscreteNullability, []string{"i32", "i32?"}, "any1?", true, "i32?"},
{"g2(any1, any1?) -> any1?", extensions.DiscreteNullability, []string{"i32", "i32"}, "any1?", false, ""},
{"g2(any1, any1?) -> any1?", extensions.DiscreteNullability, []string{"i32?", "i32?"}, "any1?", false, ""},
}

for _, tt := range tests {
t.Run(tt.definition+" with "+strings.Join(tt.argTypes, ", "), func(t *testing.T) {
funcParams := strings.TrimSuffix(strings.TrimPrefix(tt.definition[strings.Index(tt.definition, "("):strings.Index(tt.definition, ")")], "("), ")")
paramExpressions := strings.Split(funcParams, ", ")
parameters := make(extensions.FuncParameterList, len(paramExpressions))
for i, paramExpression := range paramExpressions {
parameters[i] = valArg(mustParseFuncDefArg(t, paramExpression))
}

arguments := make([]types.Type, len(tt.argTypes))
for i, argType := range tt.argTypes {
arguments[i] = mustParseConcreteType(t, argType)
}

actual, err := extensions.EvaluateTypeExpression(
"extension:test:def",
tt.nullabilityHandling,
mustParseFuncDefArg(t, tt.returnType),
parameters,
nil,
arguments,
extensions.NewSet())

if !tt.matches {
require.Error(t, err)
return
}

require.NoError(t, err)
require.True(t, mustParseConcreteType(t, tt.expectedReturnType).Equals(actual))
})
}
}

func TestResolveType(t *testing.T) {
// Test TypeReference setting logic for user-defined types
tests := []struct {
Expand Down Expand Up @@ -827,6 +913,45 @@ func TestValidateConstrainedAnyTypeConsistency(t *testing.T) {
require.Contains(t, err.Error(), "type parameter any1 cannot be both i32 and i64")
})

t.Run("nested function types with matching parameter and return types", func(t *testing.T) {
// func(func<any1 -> any1>) with (func<i32 -> i32>)
params := []types.FuncDefArgType{
&types.ParameterizedFuncType{
Parameters: []types.FuncDefArgType{&types.AnyType{Name: "any1"}},
Return: &types.AnyType{Name: "any1"},
},
}
args := []types.Type{
&types.FuncType{
ParameterTypes: []types.Type{&types.Int32Type{}},
ReturnType: &types.Int32Type{},
},
}

err := extensions.ValidateConstrainedAnyTypeConsistency(params, args, nil)
require.NoError(t, err)
})

t.Run("nested function types with mismatched return type fail", func(t *testing.T) {
// func(func<any1 -> any1>) with (func<i32 -> fp64>)
params := []types.FuncDefArgType{
&types.ParameterizedFuncType{
Parameters: []types.FuncDefArgType{&types.AnyType{Name: "any1"}},
Return: &types.AnyType{Name: "any1"},
},
}
args := []types.Type{
&types.FuncType{
ParameterTypes: []types.Type{&types.Int32Type{}},
ReturnType: &types.Float64Type{},
},
}

err := extensions.ValidateConstrainedAnyTypeConsistency(params, args, nil)
require.Error(t, err)
require.Contains(t, err.Error(), "type parameter any1 cannot be both i32 and fp64")
})

t.Run("single any1 parameter is always valid", func(t *testing.T) {
// func(any1) with (i32) - no constraint checking needed
params := []types.FuncDefArgType{
Expand Down
Loading
Loading