diff --git a/dispatch.go b/dispatch.go index 887bffb..643e13d 100644 --- a/dispatch.go +++ b/dispatch.go @@ -43,11 +43,17 @@ func newDispatcher( } // dispatch looks up an enabled ID, re-binds, authorizes, invokes, and converts typed output. +// +// remainingIntermediateBytes is the unused native-result value-body budget. +// ConvertOutput uses min(maxValueBytes, remainingIntermediateBytes) as its +// node and materialization limit. ErrValueLimit maps to resource failure; +// other conversion failures map to capability failure. func (dispatch *dispatcher) dispatch( ctx context.Context, subject authz.Subject, id string, args map[string]any, + remainingIntermediateBytes int, ) (any, error) { if dispatch == nil || dispatch.catalog == nil || dispatch.authorizer == nil { return nil, execution.ErrInternal @@ -98,19 +104,16 @@ func (dispatch *dispatcher) dispatch( if outcome.err != nil { return nil, outcome.err } - converted, conversionErr := entry.Plan.ConvertOutput(outcome.output) - if conversionErr != nil { - return nil, fmt.Errorf("%w: %w", execution.ErrCapabilityFailure, conversionErr) - } - if validationErr := binding.ValidateValue( - converted, + converted, conversionErr := entry.Plan.ConvertOutput( + outcome.output, dispatch.maxValueDepth, - dispatch.maxValueBytes, - ); validationErr != nil { - if errors.Is(validationErr, binding.ErrValueLimit) { - return nil, fmt.Errorf("%w: %w", execution.ErrResourceLimit, validationErr) + min(dispatch.maxValueBytes, remainingIntermediateBytes), + ) + if conversionErr != nil { + if errors.Is(conversionErr, binding.ErrValueLimit) { + return nil, fmt.Errorf("%w: %w", execution.ErrResourceLimit, conversionErr) } - return nil, fmt.Errorf("%w: %w", execution.ErrCapabilityFailure, validationErr) + return nil, fmt.Errorf("%w: %w", execution.ErrCapabilityFailure, conversionErr) } return converted, nil } diff --git a/dispatch_test.go b/dispatch_test.go index 173d08f..29a9834 100644 --- a/dispatch_test.go +++ b/dispatch_test.go @@ -130,8 +130,7 @@ func TestDispatchBindsAuthorizesThenInvokes(t *testing.T) { "limit": int64(25), "enabled": true, "weight": 2.5, - }, - ) + }, 64*1024) require.NoError(t, err) assert.Equal(t, []string{"authorize", "handler"}, events) @@ -197,8 +196,7 @@ func TestDispatchTranslatesEveryBindValueFailureInternally(t *testing.T) { t.Context(), authz.Subject{ID: "subject-1"}, "cap.lookup", - tt.arguments, - ) + tt.arguments, 64*1024) require.ErrorIs(t, err, worker.ErrProtocol) require.NotErrorIs(t, err, execution.ErrInvalidArguments) @@ -220,8 +218,7 @@ func TestDispatchRejectsUnknownIDsBeforeAuthorization(t *testing.T) { t.Context(), authz.Subject{ID: "subject-1"}, "cap.missing", - map[string]any{"value": "alpha"}, - ) + map[string]any{"value": "alpha"}, 64*1024) require.ErrorIs(t, err, worker.ErrProtocol) assert.Zero(t, handlerCalls.Load()) @@ -336,8 +333,7 @@ func TestDispatchClassifiesPolicyAndHandlerFailures(t *testing.T) { t.Context(), authz.Subject{ID: "subject-1"}, "cap.lookup", - map[string]any{"value": "alpha"}, - ) + map[string]any{"value": "alpha"}, 64*1024) require.ErrorIs(t, err, tt.target) }) @@ -380,8 +376,7 @@ func TestDispatchReturnsFreshCanonicalMaps(t *testing.T) { t.Context(), authz.Subject{ID: "subject-1"}, "cap.lookup", - decoded, - ) + decoded, 64*1024) require.NoError(t, err) require.NotNil(t, authorized) @@ -451,8 +446,7 @@ func TestDispatchCancellationAfterAllowPreventsInvoke(t *testing.T) { ctx, authz.Subject{ID: "subject-1"}, "cap.lookup", - map[string]any{"value": "alpha"}, - ) + map[string]any{"value": "alpha"}, 64*1024) result <- dispatchOutcome{value: value, err: err} }() @@ -510,8 +504,7 @@ func TestDispatchClassifiesParentOutputLimits(t *testing.T) { t.Context(), authz.Subject{ID: "subject-1"}, "cap.lookup", - map[string]any{"value": "alpha"}, - ) + map[string]any{"value": "alpha"}, 64*1024) require.ErrorIs(t, err, execution.ErrResourceLimit) }) @@ -540,8 +533,7 @@ func TestDispatchCancellationDuringHandlerReturnsPromptly(t *testing.T) { ctx, authz.Subject{ID: "subject-1"}, "cap.lookup", - map[string]any{"value": "alpha"}, - ) + map[string]any{"value": "alpha"}, 64*1024) result <- dispatchOutcome{value: value, err: err} }() @@ -606,3 +598,120 @@ func newWidenedDispatchSubject(t *testing.T, authorizer authz.Authorizer, invoke dispatch: newDispatcher(capabilityCatalog, authorizer, 16, 64*1024), } } + +// overflowDispatchOutput covers unsigned values above MaxInt64. +type overflowDispatchOutput struct { + // Count is an unsigned 64-bit integer. + Count uint64 `json:"count"` +} + +// nanDispatchOutput covers non-finite floating-point results. +type nanDispatchOutput struct { + // Score is a finite floating-point field. + Score float64 `json:"score"` +} + +// TestDispatchClassifiesInvalidCompositeOutputs proves invalid runtime values stay capability failures. +func TestDispatchClassifiesInvalidCompositeOutputs(t *testing.T) { + tests := []struct { + // name identifies the invalid handler output. + name string + + // plan compiles the handler contract. + plan func(*testing.T) *binding.Plan + + // invoke returns an invalid registered output. + invoke catalog.Invoker + }{ + { + name: "unsigned overflow", + plan: func(t *testing.T) *binding.Plan { + t.Helper() + plan, err := binding.CompileFor[dispatchInput, overflowDispatchOutput]() + require.NoError(t, err) + return plan + }, + invoke: func(context.Context, authz.Subject, any) (any, error) { + return overflowDispatchOutput{Count: uint64(1) << 63}, nil + }, + }, + { + name: "NaN", + plan: func(t *testing.T) *binding.Plan { + t.Helper() + plan, err := binding.CompileFor[dispatchInput, nanDispatchOutput]() + require.NoError(t, err) + return plan + }, + invoke: func(context.Context, authz.Subject, any) (any, error) { + return nanDispatchOutput{Score: math.NaN()}, nil + }, + }, + { + name: "infinity", + plan: func(t *testing.T) *binding.Plan { + t.Helper() + plan, err := binding.CompileFor[dispatchInput, nanDispatchOutput]() + require.NoError(t, err) + return plan + }, + invoke: func(context.Context, authz.Subject, any) (any, error) { + return nanDispatchOutput{Score: math.Inf(1)}, nil + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + authorizer := authzmocks.NewMockAuthorizer(t) + authorizer.EXPECT().Authorize(mock.Anything, mock.Anything).Return(nil).Once() + capabilityCatalog, err := catalog.Build([]catalog.Registration{{ + ID: "cap.lookup", + Name: "records.lookup", + Summary: "Return one record.", + Description: "Returns the supplied record value.", + Plan: tt.plan(t), + Invoke: tt.invoke, + }}, catalog.Options{ + MaxSearchQueryBytes: 256, + MaxSearchResults: 20, + }) + require.NoError(t, err) + subject := &dispatchSubject{dispatch: newDispatcher(capabilityCatalog, authorizer, 16, 64*1024)} + + _, err = subject.dispatch.dispatch( + t.Context(), + authz.Subject{ID: "subject-1"}, + "cap.lookup", + map[string]any{"value": "alpha"}, + 64*1024, + ) + + require.ErrorIs(t, err, execution.ErrCapabilityFailure) + require.NotErrorIs(t, err, execution.ErrResourceLimit) + }) + } +} + +// TestDispatchMapsValueLimitToResourceFailure proves remaining-byte conversion limits are resource failures. +func TestDispatchMapsValueLimitToResourceFailure(t *testing.T) { + authorizer := authzmocks.NewMockAuthorizer(t) + authorizer.EXPECT().Authorize(mock.Anything, mock.Anything).Return(nil).Once() + subject := newDispatchSubject( + t, + authorizer, + func(context.Context, authz.Subject, any) (any, error) { + return dispatchOutput{Value: "alpha"}, nil + }, + ) + + _, err := subject.dispatch.dispatch( + t.Context(), + authz.Subject{ID: "subject-1"}, + "cap.lookup", + map[string]any{"value": "alpha"}, + 1, + ) + + require.ErrorIs(t, err, execution.ErrResourceLimit) +} diff --git a/internal/binding/doc.go b/internal/binding/doc.go index 67e2be4..c2ac84c 100644 --- a/internal/binding/doc.go +++ b/internal/binding/doc.go @@ -6,5 +6,6 @@ // ValidateValue, FromStarlark, and ToStarlark enforce that matrix plus positive // depth and materialization limits. [json.Number] and other numeric types are // rejected. Plan.InputShape remains the only descriptor source for compiled -// input fields. +// input fields. Plan.OutputShape remains a flat FieldShape slice whose Type +// strings carry nested list, dict, struct, and nullable notation. package binding diff --git a/internal/binding/output.go b/internal/binding/output.go index b76da91..f285c9f 100644 --- a/internal/binding/output.go +++ b/internal/binding/output.go @@ -2,47 +2,298 @@ package binding import ( "fmt" + "math" "reflect" + "slices" + "strconv" + "strings" ) // ConvertOutput converts the plan's exact handler output to a process-neutral object. -func (plan *Plan) ConvertOutput(output any) (map[string]any, error) { +// +// maxDepth and maxNodes must be positive. The root struct is depth 1. Pointers +// add no depth. Nested structs, lists, and maps add one depth through their +// value node. Destination maps, lists, and sorted map-key slices are allocated +// only after their child counts fit the remaining materialization budget. +func (plan *Plan) ConvertOutput(output any, maxDepth int, maxNodes int) (map[string]any, error) { if plan == nil { return nil, fmt.Errorf("%w: nil plan", ErrInvalidPlan) } + converter, err := newValueConverter(maxDepth, maxNodes) + if err != nil { + return nil, err + } value := reflect.ValueOf(output) if !value.IsValid() || value.Type() != plan.outputType { return nil, fmt.Errorf("%w: handler output type does not match the compiled plan", ErrUnsupportedValue) } + converted, err := plan.convertNode(plan.outputRoot, value, "output", 1, &converter) + if err != nil { + return nil, err + } + object, ok := converted.(map[string]any) + if !ok { + return nil, fmt.Errorf("%w: handler output type does not match the compiled plan", ErrUnsupportedValue) + } + return object, nil +} + +// convertNode converts one compiled node to a process-neutral value. +func (plan *Plan) convertNode( + index int, + value reflect.Value, + path string, + depth int, + converter *valueConverter, +) (any, error) { + node := plan.outputNodes[index] + switch node.kind { + case outputNodePointer: + return plan.convertPointer(node, value, path, depth, converter) + case outputNodeString: + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return value.String(), nil + case outputNodeInt: + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return value.Int(), nil + case outputNodeUint: + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return convertOutputUint(value, path) + case outputNodeBool: + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return value.Bool(), nil + case outputNodeFloat: + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return convertOutputFloat(value, path) + case outputNodeBytes: + return convertOutputBytes(value, path, depth, converter) + case outputNodeList: + return plan.convertList(node, value, path, depth, converter) + case outputNodeMap: + return plan.convertMap(node, value, path, depth, converter) + case outputNodeStruct: + return plan.convertStruct(node, value, path, depth, converter) + default: + return nil, fmt.Errorf("%w: %s has an unknown compiled kind", ErrInvalidPlan, path) + } +} + +// convertPointer follows a pointer without adding depth and maps nil to None. +func (plan *Plan) convertPointer( + node outputNode, + value reflect.Value, + path string, + depth int, + converter *valueConverter, +) (any, error) { + if value.IsNil() { + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return nil, nil //nolint:nilnil // A nil pointer is process-neutral None. + } + return plan.convertNode(node.elem, value.Elem(), path, depth, converter) +} - converted := make(map[string]any, len(plan.outputFields)) - for _, field := range plan.outputFields { - item, err := convertOutputField(field, value.Field(field.index)) +// convertStruct materializes included fields after preflighting their count. +func (plan *Plan) convertStruct( + node outputNode, + value reflect.Value, + path string, + depth int, + converter *valueConverter, +) (any, error) { + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + included := 0 + for _, field := range node.fields { + if omitOutputField(field, value.Field(field.index)) { + continue + } + included++ + } + if _, err := converter.containerLength(included); err != nil { + return nil, err + } + object := make(map[string]any, included) + for _, field := range node.fields { + fieldValue := value.Field(field.index) + if omitOutputField(field, fieldValue) { + continue + } + converted, err := plan.convertNode( + field.node, + fieldValue, + outputRuntimeField(path, field.name), + depth+1, + converter, + ) if err != nil { return nil, err } - converted[field.name] = item + object[field.name] = converted } - return converted, nil + return object, nil } -// convertOutputField converts one field according to its compiled output kind. -func convertOutputField(field outputField, value reflect.Value) (any, error) { - switch field.kind { - case fieldString: - return value.String(), nil - case fieldInt64: - return value.Int(), nil - case fieldBool: - return value.Bool(), nil - case fieldFloat64: - float := value.Float() - if !isFiniteFloat(float) { - return nil, fmt.Errorf("%w: output field %q is not finite", ErrUnsupportedValue, field.name) +// convertList materializes array and non-nil slice elements after preflight. +func (plan *Plan) convertList( + node outputNode, + value reflect.Value, + path string, + depth int, + converter *valueConverter, +) (any, error) { + if value.Kind() == reflect.Slice && value.IsNil() { + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return nil, nil //nolint:nilnil // A nil slice is process-neutral None. + } + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + length, err := converter.containerLength(value.Len()) + if err != nil { + return nil, err + } + items := make([]any, length) + for index := range length { + converted, err := plan.convertNode( + node.elem, + value.Index(index), + outputRuntimeIndex(path, index), + depth+1, + converter, + ) + if err != nil { + return nil, err + } + items[index] = converted + } + return items, nil +} + +// convertMap materializes a string-keyed map in sorted key order after preflight. +func (plan *Plan) convertMap( + node outputNode, + value reflect.Value, + path string, + depth int, + converter *valueConverter, +) (any, error) { + if value.IsNil() { + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return nil, nil //nolint:nilnil // A nil map is process-neutral None. + } + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + length, err := converter.containerLength(value.Len()) + if err != nil { + return nil, err + } + keys := value.MapKeys() + slices.SortFunc(keys, compareMapKeys) + object := make(map[string]any, length) + for _, key := range keys { + name := key.String() + converted, err := plan.convertNode( + node.elem, + value.MapIndex(key), + outputRuntimeKey(path, name), + depth+1, + converter, + ) + if err != nil { + return nil, err + } + object[name] = converted + } + return object, nil +} + +// convertOutputBytes materializes a byte slice or array as a list of integers. +func convertOutputBytes(value reflect.Value, path string, depth int, converter *valueConverter) (any, error) { + if value.Kind() == reflect.Slice && value.IsNil() { + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + return nil, nil //nolint:nilnil // A nil byte slice is process-neutral None. + } + if err := converter.consumeNode(depth); err != nil { + return nil, err + } + length, err := converter.containerLength(value.Len()) + if err != nil { + return nil, err + } + items := make([]any, length) + for index := range length { + if err := converter.consumeNode(depth + 1); err != nil { + return nil, fmt.Errorf("%w: %s[%d]", err, path, index) } - return float, nil - case fieldOptionalString, fieldOptionalInt64, fieldOptionalBool, fieldOptionalFloat64: - return nil, fmt.Errorf("%w: output field %q has an invalid compiled kind", ErrInvalidPlan, field.name) + item, err := convertOutputUint(value.Index(index), path) + if err != nil { + return nil, err + } + items[index] = item + } + return items, nil +} + +// convertOutputUint normalizes an unsigned integer when it fits int64. +func convertOutputUint(value reflect.Value, path string) (int64, error) { + integer := value.Uint() + if integer > uint64(math.MaxInt64) { + return 0, fmt.Errorf("%w: %s overflows int64", ErrUnsupportedValue, path) + } + return int64(integer), nil +} + +// convertOutputFloat normalizes a floating-point value when it is finite. +func convertOutputFloat(value reflect.Value, path string) (float64, error) { + float := value.Float() + if !isFiniteFloat(float) { + return 0, fmt.Errorf("%w: %s is not finite", ErrUnsupportedValue, path) } - return nil, fmt.Errorf("%w: output field %q has an unknown compiled kind", ErrInvalidPlan, field.name) + return float, nil +} + +// omitOutputField reports whether a nil pointer+omitempty field is excluded. +func omitOutputField(field outputStructField, value reflect.Value) bool { + return field.omitempty && value.Kind() == reflect.Pointer && value.IsNil() +} + +// compareMapKeys orders reflect map keys by their string form. +func compareMapKeys(left reflect.Value, right reflect.Value) int { + return strings.Compare(left.String(), right.String()) +} + +// outputRuntimeField joins a runtime path with a struct field name. +func outputRuntimeField(path string, name string) string { + return path + "." + name +} + +// outputRuntimeIndex joins a runtime path with a list index. +func outputRuntimeIndex(path string, index int) string { + return path + "[" + strconv.Itoa(index) + "]" +} + +// outputRuntimeKey joins a runtime path with a quoted map key. +func outputRuntimeKey(path string, key string) string { + return path + "[" + strconv.Quote(key) + "]" } diff --git a/internal/binding/output_compile_test.go b/internal/binding/output_compile_test.go new file mode 100644 index 0000000..cff2dce --- /dev/null +++ b/internal/binding/output_compile_test.go @@ -0,0 +1,186 @@ +package binding + +import ( + "encoding/json" + "reflect" + "testing" + "unsafe" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// TestCompileAcceptsCompositeOutputs proves the reflected output universe registers. +func TestCompileAcceptsCompositeOutputs(t *testing.T) { + tests := []struct { + // name identifies the accepted output type. + name string + + // outputType is compiled as a capability output. + outputType reflect.Type + }{ + {name: "nested structs slices maps and pointers", outputType: reflect.TypeFor[compositeOutput]()}, + {name: "every integer and float kind", outputType: reflect.TypeFor[numericOutput]()}, + {name: "named scalar aliases", outputType: reflect.TypeFor[aliasedOutput]()}, + {name: "byte slice and array", outputType: reflect.TypeFor[bytesOutput]()}, + {name: "string-alias map keys", outputType: reflect.TypeFor[aliasKeyOutput]()}, + {name: "pointer without omitempty", outputType: reflect.TypeOf(struct { + // Value is a required nullable integer. + Value *int64 `json:"value"` + }{})}, + {name: "pointer with omitempty", outputType: reflect.TypeOf(struct { + // Value is an optional integer. + Value *int64 `json:"value,omitempty"` + }{})}, + {name: "optional nested struct", outputType: reflect.TypeOf(struct { + // Item is an optional nested object. + Item *nestedItem `json:"item,omitempty"` + }{})}, + {name: "fixed string array", outputType: reflect.TypeOf(struct { + // Tags is a fixed-length string array. + Tags [2]string `json:"tags"` + }{})}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + plan, err := Compile(reflect.TypeFor[representativeInput](), tt.outputType) + + require.NoError(t, err) + assert.Equal(t, tt.outputType, plan.OutputType()) + }) + } +} + +// TestCompileRejectsUnsupportableOutputGraphs proves unsupportable output graphs fail at registration. +func TestCompileRejectsUnsupportableOutputGraphs(t *testing.T) { + tests := []struct { + // name identifies the invalid output shape. + name string + + // outputType is compiled as a capability output. + outputType reflect.Type + + // contains is the expected diagnostic fragment. + contains string + }{ + {name: "pointer output", outputType: reflect.TypeFor[*representativeOutput](), contains: "non-pointer struct"}, + {name: "non-pointer omitempty", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value string `json:"value,omitempty"` + }{}), contains: "cannot use omitempty"}, + {name: "duplicate JSON name", outputType: duplicateJSONNameType(), contains: "duplicate output name"}, + {name: "interface field", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value any `json:"value"` + }{}), contains: "value"}, + {name: "nested interface", outputType: reflect.TypeFor[nestedInterfaceOutput](), contains: "items/value"}, + {name: "raw message", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value json.RawMessage `json:"value"` + }{}), contains: "value"}, + {name: "value json marshaler", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value valueJSONMarshaler `json:"value"` + }{}), contains: "value"}, + {name: "pointer json marshaler", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value pointerJSONMarshaler `json:"value"` + }{}), contains: "value"}, + {name: "value text marshaler", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value valueTextMarshaler `json:"value"` + }{}), contains: "value"}, + {name: "pointer text marshaler", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value pointerTextMarshaler `json:"value"` + }{}), contains: "value"}, + {name: "named uint8 marshaler slice", outputType: reflect.TypeOf(struct { + // Value is a slice of named uint8 marshalers. + Value []marshalerByte `json:"value"` + }{}), contains: "value"}, + {name: "named uint8 marshaler array", outputType: reflect.TypeOf(struct { + // Value is an array of named uint8 marshalers. + Value [2]marshalerByte `json:"value"` + }{}), contains: "value"}, + {name: "func field", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value func() `json:"value"` + }{}), contains: "value"}, + {name: "chan field", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value chan string `json:"value"` + }{}), contains: "value"}, + {name: "complex field", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value complex128 `json:"value"` + }{}), contains: "value"}, + {name: "unsafe pointer", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value unsafe.Pointer `json:"value"` + }{}), contains: "value"}, + {name: "uintptr field", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value uintptr `json:"value"` + }{}), contains: "value"}, + {name: "non-string map key", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value map[int]string `json:"value"` + }{}), contains: "unsupported map key type"}, + {name: "embedded field", outputType: reflect.TypeOf(struct { + // nestedItem is intentionally anonymous to exercise rejection. + nestedItem + }{}), contains: "embedded fields"}, + {name: "unexported field", outputType: reflect.TypeOf(struct { + // value is intentionally unexported. + value string + }{}), contains: "must be exported"}, + {name: "ignored field", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value string `json:"-"` + }{}), contains: "ignored JSON fields"}, + {name: "unknown tag option", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value string `json:"value,string"` + }{}), contains: "unsupported JSON tag option"}, + {name: "unknown struct tag", outputType: reflect.TypeOf(struct { + // Value is the field under validation. + Value string `json:"value" xml:"value"` + }{}), contains: "only one json struct tag"}, + {name: "direct cycle", outputType: reflect.TypeFor[directCycleOutput](), contains: "next"}, + {name: "indirect cycle", outputType: reflect.TypeFor[indirectCycleBranch](), contains: "leaf"}, + {name: "cycle through slice", outputType: reflect.TypeFor[sliceCycleOutput](), contains: "items"}, + {name: "cycle through map", outputType: reflect.TypeFor[mapCycleOutput](), contains: "children"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := Compile(reflect.TypeFor[representativeInput](), tt.outputType) + + require.Error(t, err) + require.ErrorIs(t, err, ErrInvalidPlan) + assert.Contains(t, err.Error(), tt.contains) + switch tt.name { + case "interface field", "nested interface", "raw message", + "value json marshaler", "pointer json marshaler", + "value text marshaler", "pointer text marshaler", + "named uint8 marshaler slice", "named uint8 marshaler array", + "func field", "chan field", "complex field", "unsafe pointer", + "uintptr field": + assert.Contains(t, err.Error(), "unsupported type") + case "direct cycle": + assert.Contains(t, err.Error(), "cycl") + assert.Contains(t, err.Error(), "next") + case "indirect cycle": + assert.Contains(t, err.Error(), "cycl") + assert.Contains(t, err.Error(), "leaf") + case "cycle through slice": + assert.Contains(t, err.Error(), "cycl") + assert.Contains(t, err.Error(), "items") + case "cycle through map": + assert.Contains(t, err.Error(), "cycl") + assert.Contains(t, err.Error(), "children") + } + }) + } +} diff --git a/internal/binding/output_shape_test.go b/internal/binding/output_shape_test.go new file mode 100644 index 0000000..c6db225 --- /dev/null +++ b/internal/binding/output_shape_test.go @@ -0,0 +1,47 @@ +package binding + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +// TestOutputShapeUsesDeterministicNestedNotation proves discovery Type strings stay flat and exact. +func TestOutputShapeUsesDeterministicNestedNotation(t *testing.T) { + plan, err := CompileFor[representativeInput, notationOutput]() + require.NoError(t, err) + + assert.Equal(t, []FieldShape{ + {Name: "items", Type: "list[{title: str, score: float}]", Required: true}, + {Name: "by_id", Type: "dict[str, {title: str, score: float}]", Required: true}, + {Name: "tags", Type: "list[str]", Required: true}, + {Name: "note", Type: "str | None", Required: true}, + {Name: "extra", Type: "str", Required: false}, + {Name: "nested", Type: "{value: str, detail?: str}", Required: true}, + {Name: "values", Type: "list[str | None] | None", Required: true}, + {Name: "alias", Type: "dict[str, bool]", Required: true}, + {Name: "payload", Type: "list[int]", Required: true}, + }, plan.OutputShape()) +} + +// TestOutputShapeDeduplicatesNoneAndPreservesDeclarationOrder proves stacked pointers append None once. +func TestOutputShapeDeduplicatesNoneAndPreservesDeclarationOrder(t *testing.T) { + plan, err := CompileFor[representativeInput, noneDedupOutput]() + require.NoError(t, err) + + assert.Equal(t, []FieldShape{ + {Name: "value", Type: "str | None", Required: true}, + }, plan.OutputShape()) + + composite, err := CompileFor[representativeInput, compositeOutput]() + require.NoError(t, err) + assert.Equal(t, []FieldShape{ + {Name: "items", Type: "list[{title: str, score: float}]", Required: true}, + {Name: "tags", Type: "list[str]", Required: true}, + {Name: "by_id", Type: "dict[str, {title: str, score: float}]", Required: true}, + {Name: "note", Type: "str | None", Required: true}, + {Name: "extra", Type: "str", Required: false}, + {Name: "payload", Type: "list[int]", Required: true}, + }, composite.OutputShape()) +} diff --git a/internal/binding/output_test.go b/internal/binding/output_test.go index f0bb406..40a50a7 100644 --- a/internal/binding/output_test.go +++ b/internal/binding/output_test.go @@ -2,19 +2,30 @@ package binding import ( "math" + "runtime" + "runtime/debug" + "strconv" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) +const ( + // generousOutputDepth is a depth budget that never binds representative values. + generousOutputDepth = 16 + + // generousOutputNodes is a node budget that never binds representative values. + generousOutputNodes = 1024 +) + // TestConvertOutputUsesCompiledFields proves handler output conversion follows the immutable output plan. func TestConvertOutputUsesCompiledFields(t *testing.T) { plan, err := CompileFor[representativeInput, representativeOutput]() require.NoError(t, err) output := representativeOutput{Name: "alpha", Count: 3, Active: true, Score: 1.5} - converted, err := plan.ConvertOutput(output) + converted, err := plan.ConvertOutput(output, generousOutputDepth, generousOutputNodes) require.NoError(t, err) assert.Equal(t, map[string]any{ @@ -31,7 +42,7 @@ func TestConvertOutputPreservesExactKinds(t *testing.T) { require.NoError(t, err) output := representativeOutput{Name: "alpha", Count: 3, Active: true, Score: 1.0} - converted, err := plan.ConvertOutput(output) + converted, err := plan.ConvertOutput(output, generousOutputDepth, generousOutputNodes) require.NoError(t, err) assert.Equal(t, map[string]any{ @@ -67,10 +78,422 @@ func TestConvertOutputRejectsTypeDriftAndNonFiniteFloats(t *testing.T) { for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - _, err := plan.ConvertOutput(tt.output) + _, err := plan.ConvertOutput(tt.output, generousOutputDepth, generousOutputNodes) require.Error(t, err) require.ErrorIs(t, err, ErrUnsupportedValue) }) } } + +// TestConvertOutputAcceptsTheCompositeUniverse proves nested values become process-neutral data. +func TestConvertOutputAcceptsTheCompositeUniverse(t *testing.T) { + note := "keep" + extra := "more" + plan, err := CompileFor[representativeInput, compositeOutput]() + require.NoError(t, err) + output := compositeOutput{ + Items: []nestedItem{{Title: "alpha", Score: 1.5}}, + Tags: [2]string{"a", "b"}, + ByID: map[string]nestedItem{"z": {Title: "zeta", Score: 2.5}}, + Note: ¬e, + Extra: &extra, + Payload: []byte{1, 2}, + } + + converted, err := plan.ConvertOutput(output, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + + assert.Equal(t, map[string]any{ + "items": []any{map[string]any{"title": "alpha", "score": 1.5}}, + "tags": []any{"a", "b"}, + "by_id": map[string]any{"z": map[string]any{"title": "zeta", "score": 2.5}}, + "note": "keep", + "extra": "more", + "payload": []any{int64(1), int64(2)}, + }, converted) +} + +// TestConvertOutputNormalizesEveryNumericKind proves integers become int64 and floats become float64. +func TestConvertOutputNormalizesEveryNumericKind(t *testing.T) { + plan, err := CompileFor[representativeInput, numericOutput]() + require.NoError(t, err) + output := numericOutput{ + I: 1, I8: 2, I16: 3, I32: 4, I64: 5, + U: 6, U8: 7, U16: 8, U32: 9, U64: 10, + F32: 1.5, F64: 2.5, + } + + converted, err := plan.ConvertOutput(output, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + + assert.Equal(t, map[string]any{ + "i": int64(1), + "i8": int64(2), + "i16": int64(3), + "i32": int64(4), + "i64": int64(5), + "u": int64(6), + "u8": int64(7), + "u16": int64(8), + "u32": int64(9), + "u64": int64(10), + "f32": float64(float32(1.5)), + "f64": 2.5, + }, converted) + assert.IsType(t, int64(0), converted["u64"]) + assert.IsType(t, float64(0), converted["f32"]) +} + +// TestConvertOutputAcceptsNamedAliasesAndByteSequences proves aliases and bytes stay in the integer list surface. +func TestConvertOutputAcceptsNamedAliasesAndByteSequences(t *testing.T) { + aliasPlan, err := CompileFor[representativeInput, aliasedOutput]() + require.NoError(t, err) + converted, err := aliasPlan.ConvertOutput(aliasedOutput{ + Name: "alpha", + Count: 3, + Active: true, + Score: 1.5, + }, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Equal(t, map[string]any{ + "name": "alpha", + "count": int64(3), + "active": true, + "score": 1.5, + }, converted) + + bytesPlan, err := CompileFor[representativeInput, bytesOutput]() + require.NoError(t, err) + converted, err = bytesPlan.ConvertOutput(bytesOutput{ + Payload: []byte{255}, + Fixed: [2]byte{1, 2}, + }, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Equal(t, map[string]any{ + "payload": []any{int64(255)}, + "fixed": []any{int64(1), int64(2)}, + }, converted) + + keyPlan, err := CompileFor[representativeInput, aliasKeyOutput]() + require.NoError(t, err) + converted, err = keyPlan.ConvertOutput(aliasKeyOutput{ + ByName: map[namedString]int64{"beta": 9}, + }, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Equal(t, map[string]any{"by_name": map[string]any{"beta": int64(9)}}, converted) +} + +// TestConvertOutputDistinguishesNilAndEmptyContainers proves nil becomes None and empty stays empty. +func TestConvertOutputDistinguishesNilAndEmptyContainers(t *testing.T) { + plan, err := CompileFor[representativeInput, compositeOutput]() + require.NoError(t, err) + + converted, err := plan.ConvertOutput(compositeOutput{ + Items: nil, + ByID: nil, + Payload: nil, + }, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Nil(t, converted["items"]) + assert.Equal(t, []any{"", ""}, converted["tags"]) + assert.Nil(t, converted["by_id"]) + assert.Nil(t, converted["note"]) + _, hasExtra := converted["extra"] + assert.False(t, hasExtra) + assert.Nil(t, converted["payload"]) + + converted, err = plan.ConvertOutput(compositeOutput{ + Items: []nestedItem{}, + ByID: map[string]nestedItem{}, + Payload: []byte{}, + }, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Equal(t, []any{}, converted["items"]) + assert.Equal(t, map[string]any{}, converted["by_id"]) + assert.Equal(t, []any{}, converted["payload"]) +} + +// TestConvertOutputOmitsOptionalNilPointers proves omitempty is distinct from required nullability. +func TestConvertOutputOmitsOptionalNilPointers(t *testing.T) { + plan, err := CompileFor[representativeInput, compositeOutput]() + require.NoError(t, err) + note := "keep" + + converted, err := plan.ConvertOutput(compositeOutput{Note: ¬e}, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Equal(t, "keep", converted["note"]) + _, hasExtra := converted["extra"] + assert.False(t, hasExtra) + + converted, err = plan.ConvertOutput(compositeOutput{}, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + assert.Nil(t, converted["note"]) + _, hasNote := converted["note"] + assert.True(t, hasNote) + _, hasExtra = converted["extra"] + assert.False(t, hasExtra) +} + +// TestConvertOutputPreservesNilPointerElements proves list and map omission never applies to elements. +func TestConvertOutputPreservesNilPointerElements(t *testing.T) { + plan, err := CompileFor[representativeInput, pointerElementOutput]() + require.NoError(t, err) + count := int64(7) + + converted, err := plan.ConvertOutput(pointerElementOutput{ + Values: []*int64{&count, nil}, + ByID: map[string]*nestedItem{"a": {Title: "alpha", Score: 1}, "b": nil}, + }, generousOutputDepth, generousOutputNodes) + require.NoError(t, err) + + assert.Equal(t, []any{int64(7), nil}, converted["values"]) + byID, ok := converted["by_id"].(map[string]any) + require.True(t, ok) + assert.Equal(t, map[string]any{"title": "alpha", "score": float64(1)}, byID["a"]) + assert.Nil(t, byID["b"]) +} + +// TestConvertOutputRejectsUnsignedOverflowAndNonFiniteFloats proves invalid numerics stay capability failures. +func TestConvertOutputRejectsUnsignedOverflowAndNonFiniteFloats(t *testing.T) { + overflowPlan, err := CompileFor[representativeInput, overflowOutput]() + require.NoError(t, err) + _, err = overflowPlan.ConvertOutput( + overflowOutput{Count: uint64(math.MaxInt64) + 1}, + generousOutputDepth, + generousOutputNodes, + ) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) + assert.Contains(t, err.Error(), "output.count") + + floatPlan, err := CompileFor[representativeInput, float32Output]() + require.NoError(t, err) + _, err = floatPlan.ConvertOutput( + float32Output{Score: float32(math.NaN())}, + generousOutputDepth, + generousOutputNodes, + ) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) + assert.Contains(t, err.Error(), "output.score") + + _, err = floatPlan.ConvertOutput( + float32Output{Score: float32(math.Inf(1))}, + generousOutputDepth, + generousOutputNodes, + ) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) + + scalarPlan, err := CompileFor[representativeInput, representativeOutput]() + require.NoError(t, err) + _, err = scalarPlan.ConvertOutput( + representativeOutput{Score: math.Inf(-1)}, + generousOutputDepth, + generousOutputNodes, + ) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) +} + +// TestConvertOutputRejectsWrongRootType proves only the compiled output type converts. +func TestConvertOutputRejectsWrongRootType(t *testing.T) { + plan, err := CompileFor[representativeInput, representativeOutput]() + require.NoError(t, err) + + _, err = plan.ConvertOutput(compositeOutput{}, generousOutputDepth, generousOutputNodes) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) + assert.Contains(t, err.Error(), "handler output type") +} + +// TestConvertOutputSortsMapsAndUsesStableErrorPaths proves map keys sort before traversal. +func TestConvertOutputSortsMapsAndUsesStableErrorPaths(t *testing.T) { + mapPlan, err := CompileFor[representativeInput, mapScoreOutput]() + require.NoError(t, err) + _, err = mapPlan.ConvertOutput(mapScoreOutput{ + ByID: map[string]float64{"z": 1, "a": math.NaN(), "m": 2}, + }, generousOutputDepth, generousOutputNodes) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) + assert.Contains(t, err.Error(), `output.by_id["a"]`) + + listPlan, err := CompileFor[representativeInput, listScoreOutput]() + require.NoError(t, err) + _, err = listPlan.ConvertOutput(listScoreOutput{ + Items: []nestedItem{ + {Title: "zero", Score: 1}, + {Title: "one", Score: 2}, + {Title: "two", Score: 3}, + {Title: "three", Score: math.Inf(1)}, + }, + }, generousOutputDepth, generousOutputNodes) + require.Error(t, err) + require.ErrorIs(t, err, ErrUnsupportedValue) + assert.Contains(t, err.Error(), "output.items[3].score") +} + +// TestConvertOutputEnforcesDepthAndNodeLimits proves root depth is 1 and scalar fields sit at depth 2. +func TestConvertOutputEnforcesDepthAndNodeLimits(t *testing.T) { + scalarPlan, err := CompileFor[representativeInput, representativeOutput]() + require.NoError(t, err) + _, err = scalarPlan.ConvertOutput(representativeOutput{Name: "alpha"}, 1, generousOutputNodes) + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) + assert.Contains(t, err.Error(), "depth") + + _, err = scalarPlan.ConvertOutput(representativeOutput{Name: "alpha"}, 2, generousOutputNodes) + require.NoError(t, err) + + nestedPlan, err := CompileFor[representativeInput, nestedDepthOutput]() + require.NoError(t, err) + nested := nestedDepthOutput{Item: nestedItem{Title: "alpha"}} + _, err = nestedPlan.ConvertOutput(nested, 2, generousOutputNodes) + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) + assert.Contains(t, err.Error(), "depth") + + _, err = nestedPlan.ConvertOutput(nested, 3, generousOutputNodes) + require.NoError(t, err) + + listPlan, err := CompileFor[representativeInput, listScoreOutput]() + require.NoError(t, err) + _, err = listPlan.ConvertOutput(listScoreOutput{ + Items: []nestedItem{{Title: "a"}, {Title: "b"}}, + }, generousOutputDepth, 2) + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) + assert.Contains(t, err.Error(), "node budget") + + _, err = scalarPlan.ConvertOutput(representativeOutput{}, 0, generousOutputNodes) + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) + + _, err = scalarPlan.ConvertOutput(representativeOutput{}, generousOutputDepth, 0) + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) +} + +// TestConvertOutputCountsIncludedOptionalFieldsBeforeAllocation proves omitempty inclusion is preflighted. +func TestConvertOutputCountsIncludedOptionalFieldsBeforeAllocation(t *testing.T) { + plan, err := CompileFor[representativeInput, optionalHeavyOutput]() + require.NoError(t, err) + a, b, c := "a", "b", "c" + + _, err = plan.ConvertOutput(optionalHeavyOutput{}, 1, 1) + require.NoError(t, err) + + _, err = plan.ConvertOutput(optionalHeavyOutput{A: &a, B: &b, C: &c}, 2, 3) + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) + assert.Contains(t, err.Error(), "node budget") +} + +// TestConvertOutputRejectsOversizedContainersBeforeAllocation proves reflected containers +// are rejected before proportional destination materialization. +func TestConvertOutputRejectsOversizedContainersBeforeAllocation(t *testing.T) { + const ( + sourceLen = 128_000 + maxDepth = 4 + maxNodes = 1024 + maxAllocBytes = 1 << 20 + ) + + items := make([]string, sourceLen) + mapped := make(map[string]string, sourceLen) + structs := make([]nestedItem, sourceLen) + for index := range sourceLen { + items[index] = "x" + mapped[strconv.Itoa(index)] = "x" + structs[index] = nestedItem{Title: "x"} + } + + tests := []struct { + // name identifies the oversized reflected container. + name string + + // compile builds the plan under test. + compile func(*testing.T) *Plan + + // output is the oversized handler value. + output any + }{ + { + name: "slice", + compile: func(t *testing.T) *Plan { + t.Helper() + plan, err := CompileFor[representativeInput, hugeSliceOutput]() + require.NoError(t, err) + return plan + }, + output: hugeSliceOutput{Items: items}, + }, + { + name: "array", + compile: func(t *testing.T) *Plan { + t.Helper() + plan, err := CompileFor[representativeInput, hugeArrayOutput]() + require.NoError(t, err) + return plan + }, + output: hugeArrayOutput{}, + }, + { + name: "map", + compile: func(t *testing.T) *Plan { + t.Helper() + plan, err := CompileFor[representativeInput, hugeMapOutput]() + require.NoError(t, err) + return plan + }, + output: hugeMapOutput{Items: mapped}, + }, + { + name: "struct list", + compile: func(t *testing.T) *Plan { + t.Helper() + plan, err := CompileFor[representativeInput, hugeStructOutput]() + require.NoError(t, err) + return plan + }, + output: hugeStructOutput{Items: structs}, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + plan := tt.compile(t) + allocated, err := measureConvertOutputGrowth(t, plan, tt.output, maxDepth, maxNodes) + + require.Error(t, err) + require.ErrorIs(t, err, ErrValueLimit) + assert.Contains(t, err.Error(), "value exceeds byte-derived node budget") + assert.LessOrEqual( + t, + allocated, + uint64(maxAllocBytes), + "ConvertOutput allocated %d bytes converting an oversized %s; destination materialization must not scale with source length %d", + allocated, + tt.name, + sourceLen, + ) + runtime.KeepAlive(tt.output) + }) + } +} + +// measureConvertOutputGrowth reports heap bytes allocated by ConvertOutput. +func measureConvertOutputGrowth(t *testing.T, plan *Plan, output any, maxDepth int, maxNodes int) (uint64, error) { + t.Helper() + + previousGCPercent := debug.SetGCPercent(-1) + defer debug.SetGCPercent(previousGCPercent) + + var before, after runtime.MemStats + runtime.ReadMemStats(&before) + _, err := plan.ConvertOutput(output, maxDepth, maxNodes) + runtime.ReadMemStats(&after) + runtime.KeepAlive(output) + return after.TotalAlloc - before.TotalAlloc, err +} diff --git a/internal/binding/output_types_test.go b/internal/binding/output_types_test.go new file mode 100644 index 0000000..725fa74 --- /dev/null +++ b/internal/binding/output_types_test.go @@ -0,0 +1,323 @@ +package binding + +// namedString is a named string alias accepted as an output scalar or map key. +type namedString string + +// namedInt is a named signed integer alias accepted as an output scalar. +type namedInt int64 + +// namedBool is a named Boolean alias accepted as an output scalar. +type namedBool bool + +// namedFloat is a named floating-point alias accepted as an output scalar. +type namedFloat float64 + +// nestedItem is a nested exported result used by composite output tests. +type nestedItem struct { + // Title is a nested string field. + Title string `json:"title"` + + // Score is a nested finite floating-point field. + Score float64 `json:"score"` +} + +// compositeOutput exercises nested structs, lists, maps, pointers, and bytes. +type compositeOutput struct { + // Items is a list of nested objects. + Items []nestedItem `json:"items"` + + // Tags is a fixed-length string array. + Tags [2]string `json:"tags"` + + // ByID is a string-keyed object map. + ByID map[string]nestedItem `json:"by_id"` + + // Note is a required nullable string. + Note *string `json:"note"` + + // Extra is an optional string omitted when nil. + Extra *string `json:"extra,omitempty"` + + // Payload is a byte slice converted as an integer list. + Payload []byte `json:"payload"` +} + +// numericOutput covers every accepted integer and floating-point kind. +type numericOutput struct { + // I is a platform signed integer. + I int `json:"i"` + + // I8 is an 8-bit signed integer. + I8 int8 `json:"i8"` + + // I16 is a 16-bit signed integer. + I16 int16 `json:"i16"` + + // I32 is a 32-bit signed integer. + I32 int32 `json:"i32"` + + // I64 is a 64-bit signed integer. + I64 int64 `json:"i64"` + + // U is a platform unsigned integer. + U uint `json:"u"` + + // U8 is an 8-bit unsigned integer. + U8 uint8 `json:"u8"` + + // U16 is a 16-bit unsigned integer. + U16 uint16 `json:"u16"` + + // U32 is a 32-bit unsigned integer. + U32 uint32 `json:"u32"` + + // U64 is a 64-bit unsigned integer. + U64 uint64 `json:"u64"` + + // F32 is a 32-bit floating-point value. + F32 float32 `json:"f32"` + + // F64 is a 64-bit floating-point value. + F64 float64 `json:"f64"` +} + +// aliasedOutput covers named aliases of accepted scalar kinds. +type aliasedOutput struct { + // Name is a named string. + Name namedString `json:"name"` + + // Count is a named integer. + Count namedInt `json:"count"` + + // Active is a named Boolean. + Active namedBool `json:"active"` + + // Score is a named float. + Score namedFloat `json:"score"` +} + +// bytesOutput covers byte slices and arrays as integer lists. +type bytesOutput struct { + // Payload is a byte slice. + Payload []byte `json:"payload"` + + // Fixed is a byte array. + Fixed [2]byte `json:"fixed"` +} + +// aliasKeyOutput covers maps whose keys are named string aliases. +type aliasKeyOutput struct { + // ByName maps named string keys onto integers. + ByName map[namedString]int64 `json:"by_name"` +} + +// notationOutput locks exact discovery Type grammar. +type notationOutput struct { + // Items is a list of nested objects. + Items []nestedItem `json:"items"` + + // ByID is a string-keyed object map. + ByID map[string]nestedItem `json:"by_id"` + + // Tags is a fixed-length string array. + Tags [2]string `json:"tags"` + + // Note is a required nullable string. + Note *string `json:"note"` + + // Extra is an optional string. + Extra *string `json:"extra,omitempty"` + + // Nested carries a nested optional field. + Nested nestedOptional `json:"nested"` + + // Values is a nullable list of nullable strings. + Values *[]*string `json:"values"` + + // Alias is a named-string-key Boolean map. + Alias map[namedString]bool `json:"alias"` + + // Payload is a byte slice. + Payload []byte `json:"payload"` +} + +// nestedOptional is a nested object with one optional field. +type nestedOptional struct { + // Value is a required nested string. + Value string `json:"value"` + + // Detail is an optional nested string. + Detail *string `json:"detail,omitempty"` +} + +// noneDedupOutput proves stacked pointers append None once. +type noneDedupOutput struct { + // Value is a double pointer to a string. + Value **string `json:"value"` +} + +// pointerElementOutput covers pointer elements inside lists and maps. +type pointerElementOutput struct { + // Values is a list of nullable integers. + Values []*int64 `json:"values"` + + // ByID is a map of nullable nested objects. + ByID map[string]*nestedItem `json:"by_id"` +} + +// optionalHeavyOutput covers omitempty inclusion accounting. +type optionalHeavyOutput struct { + // A is an optional string. + A *string `json:"a,omitempty"` + + // B is an optional string. + B *string `json:"b,omitempty"` + + // C is an optional string. + C *string `json:"c,omitempty"` +} + +// mapScoreOutput covers deterministic map traversal order. +type mapScoreOutput struct { + // ByID maps identifiers onto scores. + ByID map[string]float64 `json:"by_id"` +} + +// listScoreOutput covers deterministic list index paths. +type listScoreOutput struct { + // Items is a list of nested scored objects. + Items []nestedItem `json:"items"` +} + +// overflowOutput covers unsigned values above MaxInt64. +type overflowOutput struct { + // Count is an unsigned 64-bit integer. + Count uint64 `json:"count"` +} + +// float32Output covers non-finite float32 values. +type float32Output struct { + // Score is a 32-bit floating-point field. + Score float32 `json:"score"` +} + +// hugeSliceOutput is used by the allocation preflight regression. +type hugeSliceOutput struct { + // Items is an oversized reflected slice. + Items []string `json:"items"` +} + +// hugeArrayOutput is used by the allocation preflight regression. +type hugeArrayOutput struct { + // Items is an oversized reflected array. + Items [128000]byte `json:"items"` +} + +// hugeMapOutput is used by the allocation preflight regression. +type hugeMapOutput struct { + // Items is an oversized reflected map. + Items map[string]string `json:"items"` +} + +// hugeStructOutput is used by the allocation preflight regression. +type hugeStructOutput struct { + // Items is an oversized list of nested structs. + Items []nestedItem `json:"items"` +} + +// nestedDepthOutput is a one-level nested object used by depth-limit tests. +type nestedDepthOutput struct { + // Item is a nested object that requires depth 2. + Item nestedItem `json:"item"` +} + +// nestedInterfaceOutput embeds an interface behind a list path. +type nestedInterfaceOutput struct { + // Items is a list of objects with an interface field. + Items []struct { + // Value is the field under validation. + Value any `json:"value"` + } `json:"items"` +} + +// directCycleOutput is a self-referential pointer graph. +type directCycleOutput struct { + // Next continues the cycle. + Next *directCycleOutput `json:"next"` +} + +// indirectCycleBranch is the first type in an indirect cycle. +type indirectCycleBranch struct { + // Leaf continues the cycle. + Leaf *indirectCycleLeaf `json:"leaf"` +} + +// indirectCycleLeaf is the second type in an indirect cycle. +type indirectCycleLeaf struct { + // Branch continues the cycle. + Branch *indirectCycleBranch `json:"branch"` +} + +// sliceCycleOutput is a cycle through a slice element type. +type sliceCycleOutput struct { + // Items continues the cycle. + Items []sliceCycleOutput `json:"items"` +} + +// mapCycleOutput is a cycle through a map value type. +type mapCycleOutput struct { + // Children continues the cycle. + Children map[string]*mapCycleOutput `json:"children"` +} + +// valueJSONMarshaler implements [json.Marshaler] on the value method set. +type valueJSONMarshaler struct { + // Value is an unused host field. + Value string `json:"value"` +} + +// MarshalJSON implements [json.Marshaler]. +func (valueJSONMarshaler) MarshalJSON() ([]byte, error) { + return []byte(`{}`), nil +} + +// pointerJSONMarshaler implements [json.Marshaler] on the pointer method set. +type pointerJSONMarshaler struct { + // Value is an unused host field. + Value string `json:"value"` +} + +// MarshalJSON implements [json.Marshaler]. +func (*pointerJSONMarshaler) MarshalJSON() ([]byte, error) { + return []byte(`{}`), nil +} + +// valueTextMarshaler implements [encoding.TextMarshaler] on the value method set. +type valueTextMarshaler struct { + // Value is an unused host field. + Value string `json:"value"` +} + +// MarshalText implements [encoding.TextMarshaler]. +func (valueTextMarshaler) MarshalText() ([]byte, error) { + return []byte("x"), nil +} + +// pointerTextMarshaler implements [encoding.TextMarshaler] on the pointer method set. +type pointerTextMarshaler struct { + // Value is an unused host field. + Value string `json:"value"` +} + +// MarshalText implements [encoding.TextMarshaler]. +func (*pointerTextMarshaler) MarshalText() ([]byte, error) { + return []byte("x"), nil +} + +// marshalerByte is a named uint8 that implements [encoding.TextMarshaler]. +type marshalerByte uint8 + +// MarshalText implements [encoding.TextMarshaler]. +func (marshalerByte) MarshalText() ([]byte, error) { + return []byte("x"), nil +} diff --git a/internal/binding/plan.go b/internal/binding/plan.go index a5c3411..53b8588 100644 --- a/internal/binding/plan.go +++ b/internal/binding/plan.go @@ -1,6 +1,8 @@ package binding import ( + "encoding" + "encoding/json" "errors" "fmt" "reflect" @@ -22,7 +24,10 @@ var ( ErrValueLimit = errors.New("converted value limit exceeded") ) -// fieldKind identifies one supported direct conversion. +// outputArenaHint is the initial compiled-output arena capacity for typical capability graphs. +const outputArenaHint = 8 + +// fieldKind identifies one supported input conversion. type fieldKind uint8 const ( @@ -36,6 +41,22 @@ const ( fieldOptionalFloat64 ) +// outputNodeKind identifies one compiled output conversion node. +type outputNodeKind uint8 + +const ( + outputNodeString outputNodeKind = iota + 1 + outputNodeInt + outputNodeUint + outputNodeBool + outputNodeFloat + outputNodeBytes + outputNodeList + outputNodeMap + outputNodeStruct + outputNodePointer +) + // inputField is one immutable input-field conversion step. type inputField struct { // name is the Starlark keyword and canonical authorization key. @@ -51,16 +72,52 @@ type inputField struct { required bool } -// outputField is one immutable output-field conversion step. -type outputField struct { - // name is the string key returned to Starlark. +// outputNode is one immutable compiled output conversion node. +type outputNode struct { + // kind selects the conversion and notation strategy. + kind outputNodeKind + + // elem is the child node index for pointers, lists, and maps. + elem int + + // fields are declaration-ordered struct members. + fields []outputStructField + + // notation is the exact model-facing type string for this node. + notation string +} + +// outputStructField is one compiled exported struct member. +type outputStructField struct { + // name is the JSON and Starlark field name. name string - // index is the direct field index in the non-embedded output struct. + // index is the direct field index on the struct type. index int - // kind determines direct conversion behavior. - kind fieldKind + // node is the compiled field type. + node int + + // omitempty omits a nil pointer from the converted object. + omitempty bool +} + +// outputCompiler builds a flat immutable node arena with cycle detection. +type outputCompiler struct { + // nodes is the arena under construction. + nodes []outputNode + + // done maps a completed type onto its arena index. + done map[reflect.Type]int + + // active is the stack of types currently being compiled. + active map[reflect.Type]struct{} + + // jsonMarshaler is the json.Marshaler interface type used for method-set checks. + jsonMarshaler reflect.Type + + // textMarshaler is the encoding.TextMarshaler interface type used for method-set checks. + textMarshaler reflect.Type } // Plan is an immutable compiled input, output, signature, and canonical-argument plan. @@ -77,8 +134,11 @@ type Plan struct { // inputByName maps each accepted keyword to its field-plan index. inputByName map[string]int - // outputFields preserves declaration order for deterministic conversion. - outputFields []outputField + // outputRoot is the arena index of the compiled root struct. + outputRoot int + + // outputNodes is the immutable compiled output type arena. + outputNodes []outputNode } // CompileFor compiles the exact generic input and output types once. @@ -92,16 +152,17 @@ func Compile(inputType reflect.Type, outputType reflect.Type) (*Plan, error) { if err != nil { return nil, err } - outputFields, err := compileOutput(outputType) + root, nodes, err := compileOutput(outputType) if err != nil { return nil, err } return &Plan{ - inputType: inputType, - outputType: outputType, - inputFields: inputFields, - inputByName: inputByName, - outputFields: outputFields, + inputType: inputType, + outputType: outputType, + inputFields: inputFields, + inputByName: inputByName, + outputRoot: root, + outputNodes: nodes, }, nil } @@ -182,43 +243,254 @@ func compileInputKind(fieldType reflect.Type) (fieldKind, bool, bool) { } } -// compileOutput validates and compiles the restricted output struct. -func compileOutput(outputType reflect.Type) ([]outputField, error) { +// compileOutput validates and compiles the restricted output struct into an arena. +func compileOutput(outputType reflect.Type) (int, []outputNode, error) { if outputType == nil || outputType.Kind() != reflect.Struct { - return nil, fmt.Errorf("%w: output must be a non-pointer struct", ErrInvalidPlan) + return 0, nil, fmt.Errorf("%w: output must be a non-pointer struct", ErrInvalidPlan) + } + compiler := outputCompiler{ + nodes: make([]outputNode, 0, outputArenaHint), + done: make(map[reflect.Type]int), + active: make(map[reflect.Type]struct{}), + jsonMarshaler: reflect.TypeFor[json.Marshaler](), + textMarshaler: reflect.TypeFor[encoding.TextMarshaler](), + } + root, err := compiler.compile(outputType, "") + if err != nil { + return 0, nil, err } - fields := make([]outputField, 0, outputType.NumField()) - seen := make(map[string]struct{}, outputType.NumField()) - for index := range outputType.NumField() { - field := outputType.Field(index) + return root, compiler.nodes, nil +} + +// compile returns the arena index for typ, reusing completed nodes and rejecting cycles. +func (compiler *outputCompiler) compile(typ reflect.Type, path string) (int, error) { + if typ == nil { + return 0, fmt.Errorf("%w: output field %q has unsupported type ", ErrInvalidPlan, path) + } + if index, ok := compiler.done[typ]; ok { + return index, nil + } + if _, exists := compiler.active[typ]; exists { + return 0, fmt.Errorf("%w: cyclic type at %q", ErrInvalidPlan, path) + } + compiler.active[typ] = struct{}{} + defer delete(compiler.active, typ) + + index, err := compiler.compileNew(typ, path) + if err != nil { + return 0, err + } + compiler.done[typ] = index + return index, nil +} + +// compileNew appends one newly compiled node for typ. +func (compiler *outputCompiler) compileNew(typ reflect.Type, path string) (int, error) { + if compiler.implementsMarshaler(typ) { + return 0, unsupportedOutputType(path, typ) + } + switch typ.Kind() { //nolint:exhaustive // Unsupported reflect kinds share registration-time rejection. + case reflect.String: + return compiler.append(outputNode{kind: outputNodeString, notation: stringType}), nil + case reflect.Bool: + return compiler.append(outputNode{kind: outputNodeBool, notation: boolType}), nil + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + return compiler.append(outputNode{kind: outputNodeInt, notation: integerType}), nil + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: + return compiler.append(outputNode{kind: outputNodeUint, notation: integerType}), nil + case reflect.Float32, reflect.Float64: + return compiler.append(outputNode{kind: outputNodeFloat, notation: floatType}), nil + case reflect.Pointer: + return compiler.compilePointer(typ, path) + case reflect.Slice: + return compiler.compileSlice(typ, path) + case reflect.Array: + return compiler.compileArray(typ, path) + case reflect.Map: + return compiler.compileMap(typ, path) + case reflect.Struct: + return compiler.compileStruct(typ, path) + default: + return 0, unsupportedOutputType(path, typ) + } +} + +// compilePointer compiles a pointer to a supported node. +func (compiler *outputCompiler) compilePointer(typ reflect.Type, path string) (int, error) { + elem, err := compiler.compile(typ.Elem(), path) + if err != nil { + return 0, err + } + return compiler.append(outputNode{ + kind: outputNodePointer, + elem: elem, + notation: pointerNotation(compiler.nodes[elem].notation), + }), nil +} + +// compileSlice compiles a slice, treating uint8 elements as integer byte lists. +func (compiler *outputCompiler) compileSlice(typ reflect.Type, path string) (int, error) { + if typ.Elem().Kind() == reflect.Uint8 && !compiler.implementsMarshaler(typ.Elem()) { + return compiler.append(outputNode{kind: outputNodeBytes, notation: listNotation(integerType)}), nil + } + elem, err := compiler.compile(typ.Elem(), path) + if err != nil { + return 0, err + } + return compiler.append(outputNode{ + kind: outputNodeList, + elem: elem, + notation: listNotation(compiler.nodes[elem].notation), + }), nil +} + +// compileArray compiles a fixed array, treating uint8 elements as integer byte lists. +func (compiler *outputCompiler) compileArray(typ reflect.Type, path string) (int, error) { + if typ.Elem().Kind() == reflect.Uint8 && !compiler.implementsMarshaler(typ.Elem()) { + return compiler.append(outputNode{kind: outputNodeBytes, notation: listNotation(integerType)}), nil + } + elem, err := compiler.compile(typ.Elem(), path) + if err != nil { + return 0, err + } + return compiler.append(outputNode{ + kind: outputNodeList, + elem: elem, + notation: listNotation(compiler.nodes[elem].notation), + }), nil +} + +// compileMap compiles a string-keyed map. +func (compiler *outputCompiler) compileMap(typ reflect.Type, path string) (int, error) { + keyType := typ.Key() + if compiler.implementsMarshaler(keyType) || keyType.Kind() != reflect.String { + if path == "" { + return 0, fmt.Errorf("%w: output type %s has unsupported map key type %s", ErrInvalidPlan, typ, keyType) + } + return 0, fmt.Errorf("%w: output field %q has unsupported map key type %s", ErrInvalidPlan, path, keyType) + } + elem, err := compiler.compile(typ.Elem(), path) + if err != nil { + return 0, err + } + return compiler.append(outputNode{ + kind: outputNodeMap, + elem: elem, + notation: mapNotation(compiler.nodes[elem].notation), + }), nil +} + +// compileStruct compiles exported non-embedded fields in declaration order. +func (compiler *outputCompiler) compileStruct(typ reflect.Type, path string) (int, error) { + fields := make([]outputStructField, 0, typ.NumField()) + seen := make(map[string]struct{}, typ.NumField()) + for index := range typ.NumField() { + field := typ.Field(index) name, options, err := compileFieldName(field) if err != nil { - return nil, fmt.Errorf("%w: output field %s: %w", ErrInvalidPlan, field.Name, err) + return 0, fmt.Errorf("%w: output field %s: %w", ErrInvalidPlan, outputFieldPath(path, field.Name), err) } - if options.omitempty { - return nil, fmt.Errorf("%w: output field %q cannot use omitempty", ErrInvalidPlan, name) + fieldPath := outputFieldPath(path, name) + if options.omitempty && field.Type.Kind() != reflect.Pointer { + return 0, fmt.Errorf("%w: output field %q cannot use omitempty", ErrInvalidPlan, fieldPath) } if _, exists := seen[name]; exists { - return nil, fmt.Errorf("%w: duplicate output name %q", ErrInvalidPlan, name) + return 0, fmt.Errorf("%w: duplicate output name %q", ErrInvalidPlan, fieldPath) } - - compiled := outputField{name: name, index: index} - switch field.Type.Kind() { //nolint:exhaustive // Unsupported reflect kinds share registration-time rejection. - case reflect.String: - compiled.kind = fieldString - case reflect.Int64: - compiled.kind = fieldInt64 - case reflect.Bool: - compiled.kind = fieldBool - case reflect.Float64: - compiled.kind = fieldFloat64 - default: - return nil, fmt.Errorf("%w: output field %q has unsupported type %s", ErrInvalidPlan, name, field.Type) + node, err := compiler.compile(field.Type, fieldPath) + if err != nil { + return 0, err } seen[name] = struct{}{} - fields = append(fields, compiled) + fields = append(fields, outputStructField{ + name: name, + index: index, + node: node, + omitempty: options.omitempty, + }) + } + return compiler.append(outputNode{ + kind: outputNodeStruct, + fields: fields, + notation: compiler.structNotation(fields), + }), nil +} + +// append stores node and returns its arena index. +func (compiler *outputCompiler) append(node outputNode) int { + compiler.nodes = append(compiler.nodes, node) + return len(compiler.nodes) - 1 +} + +// structNotation renders a declaration-ordered struct literal. +func (compiler *outputCompiler) structNotation(fields []outputStructField) string { + var notation strings.Builder + notation.WriteByte('{') + for index, field := range fields { + if index > 0 { + notation.WriteString(", ") + } + notation.WriteString(field.name) + if field.omitempty { + notation.WriteByte('?') + notation.WriteString(": ") + notation.WriteString(compiler.nodes[compiler.nodes[field.node].elem].notation) + continue + } + notation.WriteString(": ") + notation.WriteString(compiler.nodes[field.node].notation) + } + notation.WriteByte('}') + return notation.String() +} + +// implementsMarshaler reports whether typ or *typ implements a forbidden marshaler. +func (compiler *outputCompiler) implementsMarshaler(typ reflect.Type) bool { + if typ == nil { + return false + } + if typ.Implements(compiler.jsonMarshaler) || typ.Implements(compiler.textMarshaler) { + return true + } + if typ.Kind() == reflect.Pointer { + return false + } + pointer := reflect.PointerTo(typ) + return pointer.Implements(compiler.jsonMarshaler) || pointer.Implements(compiler.textMarshaler) +} + +// unsupportedOutputType classifies a rejected output type at path. +func unsupportedOutputType(path string, typ reflect.Type) error { + if path == "" { + return fmt.Errorf("%w: output type %s is unsupported", ErrInvalidPlan, typ) + } + return fmt.Errorf("%w: output field %q has unsupported type %s", ErrInvalidPlan, path, typ) +} + +// outputFieldPath joins a parent compile path with a JSON field name. +func outputFieldPath(parent string, name string) string { + if parent == "" { + return name + } + return parent + "/" + name +} + +// pointerNotation appends one nullable suffix. +func pointerNotation(elem string) string { + if strings.HasSuffix(elem, noneSuffix) { + return elem } - return fields, nil + return elem + noneSuffix +} + +// listNotation renders list[T]. +func listNotation(elem string) string { + return "list[" + elem + "]" +} + +// mapNotation renders dict[str, T]. +func mapNotation(elem string) string { + return "dict[str, " + elem + "]" } // tagOptions contains the supported JSON tag option state. diff --git a/internal/binding/plan_test.go b/internal/binding/plan_test.go index 61300ef..5f2d060 100644 --- a/internal/binding/plan_test.go +++ b/internal/binding/plan_test.go @@ -254,7 +254,8 @@ func TestCompileRejectsUnsupportedInputShapes(t *testing.T) { } } -// TestCompileRejectsUnsupportedOutputShapes proves outputs cannot smuggle pointers or unsupported kinds. +// TestCompileRejectsUnsupportedOutputShapes proves root outputs stay non-pointer structs +// and non-pointer omitempty remains rejected. func TestCompileRejectsUnsupportedOutputShapes(t *testing.T) { tests := []struct { // name identifies the invalid output shape. @@ -267,11 +268,7 @@ func TestCompileRejectsUnsupportedOutputShapes(t *testing.T) { contains string }{ {name: "pointer output", outputType: reflect.TypeFor[*representativeOutput](), contains: "non-pointer struct"}, - {name: "pointer field", outputType: reflect.TypeOf(struct { - // Value is the field under validation. - Value *int64 `json:"value"` - }{}), contains: "unsupported type"}, - {name: "output omitempty", outputType: reflect.TypeOf(struct { + {name: "non-pointer omitempty", outputType: reflect.TypeOf(struct { // Value is the field under validation. Value string `json:"value,omitempty"` }{}), contains: "cannot use omitempty"}, diff --git a/internal/binding/signature.go b/internal/binding/signature.go index cb05503..3c06238 100644 --- a/internal/binding/signature.go +++ b/internal/binding/signature.go @@ -83,15 +83,28 @@ func (plan *Plan) InputShape() []FieldShape { // OutputShape returns a fresh model-facing description of the compiled output fields. func (plan *Plan) OutputShape() []FieldShape { - shape := make([]FieldShape, len(plan.outputFields)) - for index, field := range plan.outputFields { - shape[index] = FieldShape{ + root := plan.outputNodes[plan.outputRoot] + shape := make([]FieldShape, len(root.fields)) + for index, field := range root.fields { + shape[index] = outputFieldShape(plan, field) + } + return shape +} + +// outputFieldShape renders one root field's flat discovery descriptor. +func outputFieldShape(plan *Plan, field outputStructField) FieldShape { + if field.omitempty { + return FieldShape{ Name: field.name, - Type: outputKindSignature(field.kind), - Required: true, + Type: plan.outputNodes[plan.outputNodes[field.node].elem].notation, + Required: false, } } - return shape + return FieldShape{ + Name: field.name, + Type: plan.outputNodes[field.node].notation, + Required: true, + } } // ValidateInputShape reports whether fields is a combination Plan.InputShape can produce. @@ -192,20 +205,3 @@ func inputKindSignature(kind fieldKind) string { } return unsupportedTypeSignature } - -// outputKindSignature returns the model-facing notation for one supported output conversion. -func outputKindSignature(kind fieldKind) string { - switch kind { - case fieldString: - return stringType - case fieldInt64: - return integerType - case fieldBool: - return boolType - case fieldFloat64: - return floatType - case fieldOptionalString, fieldOptionalInt64, fieldOptionalBool, fieldOptionalFloat64: - return unsupportedTypeSignature - } - return unsupportedTypeSignature -} diff --git a/internal/binding/value.go b/internal/binding/value.go index ae022f5..1cc1aaf 100644 --- a/internal/binding/value.go +++ b/internal/binding/value.go @@ -82,7 +82,6 @@ func newValueConverter(maxDepth int, maxNodes int) (valueConverter, error) { return valueConverter{ maxDepth: maxDepth, remainingNodes: maxNodes, - active: make(map[visitKey]struct{}), }, nil } @@ -416,6 +415,9 @@ func (converter *valueConverter) leaveGoContainer(kind byte, value any) { // enterKey adds a prepared container identity to the active recursion path. func (converter *valueConverter) enterKey(key visitKey) (visitKey, error) { + if converter.active == nil { + converter.active = make(map[visitKey]struct{}) + } if _, exists := converter.active[key]; exists { return visitKey{}, fmt.Errorf("%w: cyclic value", ErrUnsupportedValue) } diff --git a/internal/worker/parent.go b/internal/worker/parent.go index 38346c4..2732fa8 100644 --- a/internal/worker/parent.go +++ b/internal/worker/parent.go @@ -113,7 +113,12 @@ type Limits struct { } // Dispatch invokes one authoritative parent capability operation. -type Dispatch func(context.Context, authz.Subject, string, map[string]any) (any, error) +// +// remainingIntermediateBytes is the unused native-result value-body budget for +// this execution. handleNative passes min(MaxValueBytes, remainingIntermediate); +// the dispatcher uses that as the ConvertOutput node and materialization limit. +// Exact encoded-byte debit stays in writeNativeResult. +type Dispatch func(context.Context, authz.Subject, string, map[string]any, int) (any, error) // ErrProtocol lets the authoritative root dispatcher report an impossible child ID or re-bind mismatch. var ErrProtocol = errors.New("worker protocol violation") @@ -721,7 +726,13 @@ func (r *Runner) handleNative( ) execOutcome { results := make(chan callbackResult, 1) go func() { - value, err := r.dispatch(runCtx, subject, frame.CapabilityID, frame.Arguments) + value, err := r.dispatch( + runCtx, + subject, + frame.CapabilityID, + frame.Arguments, + min(r.limits.MaxValueBytes, conn.remainingIntermediate), + ) results <- callbackResult{value: value, err: err} }() select { diff --git a/internal/worker/parent_test.go b/internal/worker/parent_test.go index 844ded6..f67389c 100644 --- a/internal/worker/parent_test.go +++ b/internal/worker/parent_test.go @@ -58,7 +58,7 @@ func lookupBinding() execution.CapabilityBinding { // nopDispatch rejects every native call as an ordinary capability failure. func nopDispatch() Dispatch { - return func(context.Context, authz.Subject, string, map[string]any) (any, error) { + return func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { return nil, execution.ErrCapabilityFailure } } @@ -275,10 +275,14 @@ func TestRunnerIntermediateBudget(t *testing.T) { limits := testLimits() limits.MaxIntermediateValueBytes = len(body)*2 + 1 var calls int - runner := newTestRunner(t, limits, func(context.Context, authz.Subject, string, map[string]any) (any, error) { - calls++ - return result, nil - }) + runner := newTestRunner( + t, + limits, + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + calls++ + return result, nil + }, + ) got, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, twoCallSource) @@ -290,9 +294,13 @@ func TestRunnerIntermediateBudget(t *testing.T) { t.Run("inclusive edge", func(t *testing.T) { limits := testLimits() limits.MaxIntermediateValueBytes = len(body) * 2 - runner := newTestRunner(t, limits, func(context.Context, authz.Subject, string, map[string]any) (any, error) { - return result, nil - }) + runner := newTestRunner( + t, + limits, + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + return result, nil + }, + ) got, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, twoCallSource) @@ -304,10 +312,14 @@ func TestRunnerIntermediateBudget(t *testing.T) { limits := testLimits() limits.MaxIntermediateValueBytes = len(body)*2 - 1 var calls int - runner := newTestRunner(t, limits, func(context.Context, authz.Subject, string, map[string]any) (any, error) { - calls++ - return result, nil - }) + runner := newTestRunner( + t, + limits, + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + calls++ + return result, nil + }, + ) _, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, twoCallSource) @@ -318,9 +330,13 @@ func TestRunnerIntermediateBudget(t *testing.T) { t.Run("fresh next execution", func(t *testing.T) { limits := testLimits() limits.MaxIntermediateValueBytes = len(body) - runner := newTestRunner(t, limits, func(context.Context, authz.Subject, string, map[string]any) (any, error) { - return result, nil - }) + runner := newTestRunner( + t, + limits, + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + return result, nil + }, + ) _, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, twoCallSource) require.ErrorIs(t, err, execution.ErrResourceLimit) @@ -334,9 +350,13 @@ func TestRunnerIntermediateBudget(t *testing.T) { limits := testLimits() limits.MaxIntermediateValueBytes = len(body) limits.MaxConcurrentExecutions = 2 - runner := newTestRunner(t, limits, func(context.Context, authz.Subject, string, map[string]any) (any, error) { - return result, nil - }) + runner := newTestRunner( + t, + limits, + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + return result, nil + }, + ) var wg sync.WaitGroup errCh := make(chan error, 2) @@ -428,7 +448,7 @@ func TestRunnerNativeForwarding(t *testing.T) { gotArgs map[string]any parent = os.Getpid() ) - dispatch := func(_ context.Context, _ authz.Subject, id string, arguments map[string]any) (any, error) { + dispatch := func(_ context.Context, _ authz.Subject, id string, arguments map[string]any, _ int) (any, error) { gotID = id gotArgs = arguments return map[string]any{"pid": int64(os.Getpid())}, nil @@ -446,9 +466,13 @@ func TestRunnerNativeForwarding(t *testing.T) { // TestRunnerRetainedCallbackError proves ordinary dispatch errors abort and win cleanup. func TestRunnerRetainedCallbackError(t *testing.T) { retained := execution.ErrPermissionDenied - runner := newTestRunner(t, testLimits(), func(context.Context, authz.Subject, string, map[string]any) (any, error) { - return nil, retained - }) + runner := newTestRunner( + t, + testLimits(), + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + return nil, retained + }, + ) _, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, nativeSource) @@ -460,9 +484,13 @@ func TestRunnerRetainedCallbackError(t *testing.T) { // TestRunnerProtocolViolationKills proves ErrProtocol skips native_abort and returns ErrInternal. func TestRunnerProtocolViolationKills(t *testing.T) { - runner := newTestRunner(t, testLimits(), func(context.Context, authz.Subject, string, map[string]any) (any, error) { - return nil, fmt.Errorf("unknown id: %w", ErrProtocol) - }) + runner := newTestRunner( + t, + testLimits(), + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + return nil, fmt.Errorf("unknown id: %w", ErrProtocol) + }, + ) _, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, nativeSource) @@ -498,9 +526,13 @@ func TestRunnerFinalErrorTrailingByteKills(t *testing.T) { // TestRunnerProtocolViolationWritesNoAbort proves ErrProtocol kills without writing native_abort. func TestRunnerProtocolViolationWritesNoAbort(t *testing.T) { - runner := newTestRunner(t, testLimits(), func(context.Context, authz.Subject, string, map[string]any) (any, error) { - return nil, fmt.Errorf("unknown id: %w", ErrProtocol) - }) + runner := newTestRunner( + t, + testLimits(), + func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { + return nil, fmt.Errorf("unknown id: %w", ErrProtocol) + }, + ) var parentWrites bytes.Buffer var childOut bytes.Buffer payload, err := encodeNativeCall("cap.lookup", map[string]any{"value": "alpha"}) @@ -549,7 +581,7 @@ func TestRunnerQueueCancellationSpawnsNoChild(t *testing.T) { started := make(chan struct{}) release := make(chan struct{}) var live atomic.Int32 - dispatch := func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any) (any, error) { + dispatch := func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any, _ int) (any, error) { live.Add(1) close(started) select { @@ -590,7 +622,7 @@ func TestRunnerQueueCancellationSpawnsNoChild(t *testing.T) { // TestRunnerPermitReuseAfterKill proves a later execution works after kill/reap. func TestRunnerPermitReuseAfterKill(t *testing.T) { started := make(chan struct{}) - dispatch := func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any) (any, error) { + dispatch := func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any, _ int) (any, error) { close(started) <-ctx.Done() return nil, ctx.Err() @@ -652,7 +684,7 @@ func TestRunnerParallelCap(t *testing.T) { limits := testLimits() limits.MaxConcurrentExecutions = 2 var current, peak atomic.Int32 - dispatch := func(context.Context, authz.Subject, string, map[string]any) (any, error) { + dispatch := func(context.Context, authz.Subject, string, map[string]any, int) (any, error) { n := current.Add(1) for { prev := peak.Load() @@ -688,7 +720,7 @@ func TestRunnerParallelCap(t *testing.T) { // TestRunnerCancellationDuringDispatch proves cancellation is not joined to the callback. func TestRunnerCancellationDuringDispatch(t *testing.T) { started := make(chan struct{}) - dispatch := func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any) (any, error) { + dispatch := func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any, _ int) (any, error) { close(started) <-ctx.Done() time.Sleep(20 * time.Millisecond) @@ -723,7 +755,7 @@ func TestRunnerRepeatedCleanup(t *testing.T) { runner := newTestRunner( t, testLimits(), - func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any) (any, error) { + func(ctx context.Context, _ authz.Subject, _ string, _ map[string]any, _ int) (any, error) { <-ctx.Done() return nil, ctx.Err() }, @@ -808,3 +840,54 @@ func TestProbeHelperExit(_ *testing.T) { } os.Exit(code) } + +// TestRunnerForwardsMinRemainingIntermediateBytes proves handleNative forwards +// min(MaxValueBytes, remaining intermediate) into Dispatch. +func TestRunnerForwardsMinRemainingIntermediateBytes(t *testing.T) { + const first = "xx" + body, err := encodeNormalizedValue(first) + require.NoError(t, err) + const twoCallSource = "def main():\n records.lookup(value=\"a\")\n return records.lookup(value=\"b\")\n" + + t.Run("remaining is smaller than MaxValueBytes", func(t *testing.T) { + limits := testLimits() + limits.MaxValueBytes = 1024 + limits.MaxIntermediateValueBytes = len(body) * 2 + got := make([]int, 0, 2) + runner := newTestRunner( + t, + limits, + func(_ context.Context, _ authz.Subject, _ string, _ map[string]any, remaining int) (any, error) { + got = append(got, remaining) + return first, nil + }, + ) + + result, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, twoCallSource) + require.NoError(t, err) + assert.Equal(t, first, result) + require.Len(t, got, 2) + assert.Equal(t, []int{len(body) * 2, len(body)}, got) + }) + + t.Run("MaxValueBytes is smaller than remaining", func(t *testing.T) { + limits := testLimits() + limits.MaxValueBytes = 64 + limits.MaxIntermediateValueBytes = 1024 + got := make([]int, 0, 1) + runner := newTestRunner( + t, + limits, + func(_ context.Context, _ authz.Subject, _ string, _ map[string]any, remaining int) (any, error) { + got = append(got, remaining) + return first, nil + }, + ) + + result, err := runner.Execute(context.Background(), authz.Subject{ID: "s"}, nativeSource) + require.NoError(t, err) + assert.Equal(t, first, result) + require.Len(t, got, 1) + assert.Equal(t, 64, got[0]) + }) +}