From 5dfb86cf629560f1b1cc7b0a2309b1389f6c0088 Mon Sep 17 00:00:00 2001 From: quobix Date: Mon, 21 Sep 2026 11:23:00 -0400 Subject: [PATCH] Tighten the generated SDK runtime and finish the gosdk rename Follow-up to #626, which merged before these review residuals landed. Generated runtime: - Read response bodies through readBody. bytes.Buffer.ReadFrom reserves bytes.MinRead before every read, so Grow(Content-Length) alone still reallocated and a 12-byte body cost 1.6KB. Reserve the headroom when the length is declared and fall back to io.ReadAll when it is not. - Apply credentials in place. authorize cloned the whole request for every complete alternative; it now clones only when a later alternative could still run, which is the one case where a failed Apply can leak partial writes. One Widgets.List call against a stub doer goes from 2067ns, 5560B and 45 allocs to 1696ns, 3656B and 38 allocs. Emitter: - Rename generator/sdk/golang to generator/sdk/gosdk so the directory matches the package, and carry the name through the doc comment, the README, the generated file headers and the error prefix. No tag contains the merge yet, so the import path is still free to move. - Share one shape walk (forEachShapeChild) between reachability and collection so the two cannot disagree about what is generated. - Give generator/golang a single scalar mapping. ScalarType now takes the JSON type and format, and goType uses the same builtinScalarType, so SDK parameters and model fields cannot drift. Tests: the generated -race suite now proves a failed alternative does not leak into its fallback (it fails with the clone removed) and that Content-Length is a preallocation hint, never a bound. Co-Authored-By: Claude Fable 5.1 --- generator/golang/coverage_test.go | 27 ++- generator/golang/generator.go | 21 ++- generator/golang/to_go.go | 45 +++-- generator/sdk/README.md | 4 +- .../sdk/{golang => gosdk}/coverage_test.go | 0 generator/sdk/{golang => gosdk}/doc.go | 2 +- generator/sdk/{golang => gosdk}/generator.go | 156 ++++++++---------- .../sdk/{golang => gosdk}/generator_test.go | 78 ++++++++- generator/sdk/{golang => gosdk}/options.go | 0 .../{golang => gosdk}/templates/client.tmpl | 83 +++++++--- .../templates/resources.tmpl | 2 +- .../templates/workflows.tmpl | 2 +- generator/sdk/{golang => gosdk}/view.go | 0 13 files changed, 268 insertions(+), 152 deletions(-) rename generator/sdk/{golang => gosdk}/coverage_test.go (100%) rename generator/sdk/{golang => gosdk}/doc.go (63%) rename generator/sdk/{golang => gosdk}/generator.go (80%) rename generator/sdk/{golang => gosdk}/generator_test.go (88%) rename generator/sdk/{golang => gosdk}/options.go (100%) rename generator/sdk/{golang => gosdk}/templates/client.tmpl (87%) rename generator/sdk/{golang => gosdk}/templates/resources.tmpl (98%) rename generator/sdk/{golang => gosdk}/templates/workflows.tmpl (94%) rename generator/sdk/{golang => gosdk}/view.go (100%) diff --git a/generator/golang/coverage_test.go b/generator/golang/coverage_test.go index 86b586f4..1c49b3b9 100644 --- a/generator/golang/coverage_test.go +++ b/generator/golang/coverage_test.go @@ -366,11 +366,28 @@ func TestSharedNameAndReferenceHelpers(t *testing.T) { if got := RefName("#/components/schemas/Tenant~1Record~0V2"); got != "Tenant/Record~V2" { t.Fatalf("decoded reference name = %q", got) } - if got := NewGenerator(WithFormatMapping("date-time", "time.Time", "time")).ScalarType("date-time", "string"); got != "time.Time" { - t.Fatalf("mapped scalar type = %q", got) - } - if got := NewGenerator().ScalarType("unknown", "string"); got != "string" { - t.Fatalf("fallback scalar type = %q", got) + mapped := NewGenerator(WithFormatMapping("date-time", "time.Time", "time")) + for _, test := range []struct { + jsonType string + format string + want string + scalar bool + }{ + {jsonType: "string", format: "date-time", want: "time.Time", scalar: true}, + {jsonType: "string", format: "uuid", want: "string", scalar: true}, + {jsonType: "integer", format: "int32", want: "int32", scalar: true}, + {jsonType: "integer", format: "int64", want: "int64", scalar: true}, + {jsonType: "integer", want: "int", scalar: true}, + {jsonType: "number", format: "float", want: "float32", scalar: true}, + {jsonType: "number", format: "double", want: "float64", scalar: true}, + {jsonType: "boolean", want: "bool", scalar: true}, + {jsonType: "array"}, + {jsonType: "object"}, + } { + got, scalar := mapped.ScalarType(test.jsonType, test.format) + if got != test.want || scalar != test.scalar { + t.Fatalf("ScalarType(%q, %q) = %q, %t; want %q, %t", test.jsonType, test.format, got, scalar, test.want, test.scalar) + } } } diff --git a/generator/golang/generator.go b/generator/golang/generator.go index 7b94ed2e..f1c10cfe 100644 --- a/generator/golang/generator.go +++ b/generator/golang/generator.go @@ -113,13 +113,22 @@ type GeneratedField struct { Type string } -// ScalarType returns the configured Go type for an OpenAPI string format, or -// fallback when the generator has no mapping for format. -func (g *Generator) ScalarType(format, fallback string) string { - if mapping, ok := g.formatMappings[format]; ok { - return mapping.goType +// ScalarType returns the Go type this generator renders for a scalar JSON +// Schema type and format, including configured string format mappings. The +// second result is false when jsonType is not a scalar. SDK emitters call it so +// parameter types and model field types come from one mapping. +func (g *Generator) ScalarType(jsonType, format string) (string, bool) { + kind := kindForJSONType(jsonType) + switch kind { + case KindString: + if mapping, ok := g.formatMappings[format]; ok { + return mapping.goType, true + } + return builtinScalarType(kind, format), true + case KindInteger, KindNumber, KindBoolean: + return builtinScalarType(kind, format), true } - return fallback + return "", false } // NewGenerator creates a Go model generator. diff --git a/generator/golang/to_go.go b/generator/golang/to_go.go index 8e707d65..514a4525 100644 --- a/generator/golang/to_go.go +++ b/generator/golang/to_go.go @@ -418,24 +418,9 @@ func (g *Generator) goType(ir *SchemaIR, required bool, field bool) string { case KindMap: typ = "map[string]" + g.goType(ir.AdditionalProperties, true, false) case KindString: - typ = g.formatType(ir.Format, "string") - case KindInteger: - switch ir.Format { - case "int32": - typ = "int32" - case "int64": - typ = "int64" - default: - typ = "int" - } - case KindNumber: - if ir.Format == "float" { - typ = "float32" - } else { - typ = "float64" - } - case KindBoolean: - typ = "bool" + typ = g.formatType(ir.Format, builtinScalarType(ir.Kind, ir.Format)) + case KindInteger, KindNumber, KindBoolean: + typ = builtinScalarType(ir.Kind, ir.Format) case KindEnum: if ir.Name != "" { typ = ir.Name @@ -462,6 +447,30 @@ func (g *Generator) formatType(format, fallback string) string { return fallback } +// builtinScalarType is the one mapping from a scalar kind and format to a Go +// builtin. goType and ScalarType both use it, so generated models and SDK +// parameters cannot drift apart. +func builtinScalarType(kind Kind, format string) string { + switch kind { + case KindInteger: + switch format { + case "int32": + return "int32" + case "int64": + return "int64" + } + return "int" + case KindNumber: + if format == "float" { + return "float32" + } + return "float64" + case KindBoolean: + return "bool" + } + return "string" +} + func pointerDepth(typ string, ir *SchemaIR, required, optionalPointers, nullablePointer, optionalNullableDoublePointer bool) int { compound := typ == "any" || strings.HasPrefix(typ, "[]") || strings.HasPrefix(typ, "map[") nullable := ir != nil && ir.Nullable && nullablePointer diff --git a/generator/sdk/README.md b/generator/sdk/README.md index f0094936..5110aa65 100644 --- a/generator/sdk/README.md +++ b/generator/sdk/README.md @@ -4,7 +4,7 @@ Arazzo. Language emitters consume that prepared contract. Generated clients do not import libopenapi or parse specifications at runtime. -The first emitter is `generator/sdk/golang`. It generates: +The first emitter is `generator/sdk/gosdk`. It generates: - reachable models through `generator/golang`; - one shared HTTP client and security runtime; @@ -27,7 +27,7 @@ if err != nil { return err } -result, err := sdkgolang.GenerateContract(contract, sdkgolang.Options{ +result, err := gosdk.GenerateContract(contract, gosdk.Options{ PackageName: "exampleapi", }) ``` diff --git a/generator/sdk/golang/coverage_test.go b/generator/sdk/gosdk/coverage_test.go similarity index 100% rename from generator/sdk/golang/coverage_test.go rename to generator/sdk/gosdk/coverage_test.go diff --git a/generator/sdk/golang/doc.go b/generator/sdk/gosdk/doc.go similarity index 63% rename from generator/sdk/golang/doc.go rename to generator/sdk/gosdk/doc.go index 16b95d28..722f616b 100644 --- a/generator/sdk/golang/doc.go +++ b/generator/sdk/gosdk/doc.go @@ -1,6 +1,6 @@ // Copyright 2026 Princess B33f Heavy Industries / Dave Shanley // SPDX-License-Identifier: MIT -// Package golang emits an idiomatic, resource-oriented Go SDK from a prepared +// Package gosdk emits an idiomatic, resource-oriented Go SDK from a prepared // OpenAPI client contract. package gosdk diff --git a/generator/sdk/golang/generator.go b/generator/sdk/gosdk/generator.go similarity index 80% rename from generator/sdk/golang/generator.go rename to generator/sdk/gosdk/generator.go index 278f2d08..70aff48d 100644 --- a/generator/sdk/golang/generator.go +++ b/generator/sdk/gosdk/generator.go @@ -49,7 +49,7 @@ var workflowsTemplate = template.Must(template.New("workflows").Funcs(templateFu // finite Arazzo workflow helpers. func Generate(document *highv3.Document, options Options) (*sdk.Result, error) { if len(options.Workflows) > 0 { - return nil, errors.New("sdk/golang: Generate does not accept workflows; prepare the contract and call GenerateContract") + return nil, errors.New("gosdk: Generate does not accept workflows; prepare the contract and call GenerateContract") } contract, err := sdk.Prepare(document, options.Prepare) if err != nil { @@ -61,14 +61,14 @@ func Generate(document *highv3.Document, options Options) (*sdk.Result, error) { // GenerateContract emits a Go SDK from a prepared contract. func GenerateContract(contract *sdk.Contract, options Options) (*sdk.Result, error) { if contract == nil { - return nil, errors.New("sdk/golang: contract is required") + return nil, errors.New("gosdk: contract is required") } if len(contract.Operations) == 0 { - return nil, errors.New("sdk/golang: contract contains no selected operations") + return nil, errors.New("gosdk: contract contains no selected operations") } packageName := options.packageName() if !token.IsIdentifier(packageName) || token.Lookup(packageName).IsKeyword() { - return nil, fmt.Errorf("sdk/golang: invalid package name %q", packageName) + return nil, fmt.Errorf("gosdk: invalid package name %q", packageName) } var sdkEmitter *emitter modelOptions := append([]modelgen.Option(nil), options.Models...) @@ -86,7 +86,7 @@ func GenerateContract(contract *sdk.Contract, options Options) (*sdk.Result, err } models, err := modelsGenerator.RenderSchemas(sdkEmitter.schemas) if err != nil { - return nil, fmt.Errorf("sdk/golang: generate models: %w", err) + return nil, fmt.Errorf("gosdk: generate models: %w", err) } if err := resolveWorkflowFieldTypes(models.Types, view.Workflows); err != nil { return nil, err @@ -122,11 +122,11 @@ func GenerateContract(contract *sdk.Contract, options Options) (*sdk.Result, err func render(tmpl *template.Template, value any) ([]byte, error) { var output bytes.Buffer if err := tmpl.Execute(&output, value); err != nil { - return nil, fmt.Errorf("sdk/golang: execute %s template: %w", tmpl.Name(), err) + return nil, fmt.Errorf("gosdk: execute %s template: %w", tmpl.Name(), err) } formatted, err := format.Source(output.Bytes()) if err != nil { - return nil, fmt.Errorf("sdk/golang: format %s: %w\n%s", tmpl.Name(), err, output.Bytes()) + return nil, fmt.Errorf("gosdk: format %s: %w\n%s", tmpl.Name(), err, output.Bytes()) } return formatted, nil } @@ -200,28 +200,7 @@ func (e *emitter) findReachableComponents() { if schema == nil { return } - for _, child := range schema.AllOf { - walk(child) - } - for _, child := range schema.OneOf { - walk(child) - } - for _, child := range schema.AnyOf { - walk(child) - } - if schema.Items != nil && schema.Items.IsA() { - walk(schema.Items.A) - } - if schema.AdditionalProperties != nil && schema.AdditionalProperties.IsA() { - walk(schema.AdditionalProperties.A) - } - for _, children := range []*orderedmap.Map[string, *highbase.SchemaProxy]{schema.Properties, schema.PatternProperties} { - if children != nil { - for _, child := range children.FromOldest() { - walk(child) - } - } - } + forEachShapeChild(schema, walk) } for _, operation := range e.contract.Operations { if operation == nil { @@ -285,7 +264,7 @@ func (e *emitter) prepareView() (*clientView, error) { view.DefaultServer = operation.Servers[0] } if len(operation.Servers) > 0 && view.DefaultServer != "" && operation.Servers[0] != view.DefaultServer { - return nil, fmt.Errorf("sdk/golang: operation %q uses unsupported operation-specific server %q", operation.ID, operation.Servers[0]) + return nil, fmt.Errorf("gosdk: operation %q uses unsupported operation-specific server %q", operation.ID, operation.Servers[0]) } for _, alternative := range operation.Security { if alternative == nil { @@ -296,7 +275,7 @@ func (e *emitter) prepareView() (*clientView, error) { continue } if _, ok := knownSchemes[scheme.Name]; !ok { - return nil, fmt.Errorf("sdk/golang: operation %q references undefined security scheme %q", operation.ID, scheme.Name) + return nil, fmt.Errorf("gosdk: operation %q references undefined security scheme %q", operation.ID, scheme.Name) } } } @@ -310,7 +289,7 @@ func (e *emitter) prepareView() (*clientView, error) { } methodName := modelgen.PublicName(operation.Name) if prior, exists := resourceMethods[resourceName][methodName]; exists { - return nil, fmt.Errorf("sdk/golang: operations %q and %q both map to %s.%s", prior, operation.ID, resourceName, methodName) + return nil, fmt.Errorf("gosdk: operations %q and %q both map to %s.%s", prior, operation.ID, resourceName, methodName) } resourceMethods[resourceName][methodName] = operation.ID operationView, err := e.prepareOperation(operation, methodName) @@ -356,38 +335,38 @@ func (e *emitter) prepareOperation(operation *sdk.Operation, methodName string) continue } if parameter.Schema == nil { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q has no schema", operation.ID, parameter.Name) + return view, fmt.Errorf("gosdk: operation %q parameter %q has no schema", operation.ID, parameter.Name) } if !parameterSchemaSupported(parameter.Schema) { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q uses unsupported object or tuple serialization", operation.ID, parameter.Name) + return view, fmt.Errorf("gosdk: operation %q parameter %q uses unsupported object or tuple serialization", operation.ID, parameter.Name) } if parameter.In == "cookie" { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q uses unsupported cookie serialization", operation.ID, parameter.Name) + return view, fmt.Errorf("gosdk: operation %q parameter %q uses unsupported cookie serialization", operation.ID, parameter.Name) } if parameter.In != "path" && parameter.In != "query" && parameter.In != "header" { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q has unsupported location %q", operation.ID, parameter.Name, parameter.In) + return view, fmt.Errorf("gosdk: operation %q parameter %q has unsupported location %q", operation.ID, parameter.Name, parameter.In) } if parameter.AllowReserved { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q uses unsupported allowReserved serialization", operation.ID, parameter.Name) + return view, fmt.Errorf("gosdk: operation %q parameter %q uses unsupported allowReserved serialization", operation.ID, parameter.Name) } if !supportedStyle(parameter.In, parameter.Style) { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q uses unsupported %s style %q", operation.ID, parameter.Name, parameter.In, parameter.Style) + return view, fmt.Errorf("gosdk: operation %q parameter %q uses unsupported %s style %q", operation.ID, parameter.Name, parameter.In, parameter.Style) } typeName, err := e.schemaType(parameter.Schema, operation.ID+modelgen.PublicName(parameter.Name)+"Parameter") if err != nil { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q: %w", operation.ID, parameter.Name, err) + return view, fmt.Errorf("gosdk: operation %q parameter %q: %w", operation.ID, parameter.Name, err) } if !parameter.Required { typeName = "*" + typeName } fieldName := modelgen.PublicName(parameter.Name) if prior, ok := fieldNames[fieldName]; ok { - return view, fmt.Errorf("sdk/golang: operation %q parameters %q and %q have the same Go field name %q", operation.ID, prior, parameter.Name, fieldName) + return view, fmt.Errorf("gosdk: operation %q parameters %q and %q have the same Go field name %q", operation.ID, prior, parameter.Name, fieldName) } fieldNames[fieldName] = parameter.Name encoder, array, err := e.parameterEncoder(parameter.Schema) if err != nil { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q: %w", operation.ID, parameter.Name, err) + return view, fmt.Errorf("gosdk: operation %q parameter %q: %w", operation.ID, parameter.Name, err) } view.Parameters = append(view.Parameters, parameterView{ Name: parameter.Name, FieldName: fieldName, Type: typeName, In: parameter.In, @@ -400,18 +379,18 @@ func (e *emitter) prepareOperation(operation *sdk.Operation, methodName string) } if operation.RequestBody != nil { if prior, exists := fieldNames["Body"]; exists { - return view, fmt.Errorf("sdk/golang: operation %q parameter %q collides with request body field Body", operation.ID, prior) + return view, fmt.Errorf("gosdk: operation %q parameter %q collides with request body field Body", operation.ID, prior) } mediaType, schema, err := jsonMedia(operation.RequestBody.Content) if err != nil { - return view, fmt.Errorf("sdk/golang: operation %q request body: %w", operation.ID, err) + return view, fmt.Errorf("gosdk: operation %q request body: %w", operation.ID, err) } if schema == nil { - return view, fmt.Errorf("sdk/golang: operation %q request body has no schema", operation.ID) + return view, fmt.Errorf("gosdk: operation %q request body has no schema", operation.ID) } typeName, err := e.schemaType(schema, operation.ID+"Request") if err != nil { - return view, fmt.Errorf("sdk/golang: operation %q request body: %w", operation.ID, err) + return view, fmt.Errorf("gosdk: operation %q request body: %w", operation.ID, err) } fieldType := typeName if !operation.RequestBody.Required { @@ -428,12 +407,12 @@ func (e *emitter) prepareOperation(operation *sdk.Operation, methodName string) view.SuccessStatuses = successes view.SuccessCondition, err = statusCondition(successes) if err != nil { - return view, fmt.Errorf("sdk/golang: operation %q success responses: %w", operation.ID, err) + return view, fmt.Errorf("gosdk: operation %q success responses: %w", operation.ID, err) } for index := range errorResponses { errorResponses[index].Condition, err = statusCondition([]string{errorResponses[index].Status}) if err != nil { - return view, fmt.Errorf("sdk/golang: operation %q error response %s: %w", operation.ID, errorResponses[index].Status, err) + return view, fmt.Errorf("gosdk: operation %q error response %s: %w", operation.ID, errorResponses[index].Status, err) } } view.ErrorResponses = errorResponses @@ -455,11 +434,11 @@ func (e *emitter) prepareOperation(operation *sdk.Operation, methodName string) func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string]*operationView) (workflowView, error) { if workflow == nil || workflow.Operation == nil { - return workflowView{}, errors.New("sdk/golang: workflow and operation are required") + return workflowView{}, errors.New("gosdk: workflow and operation are required") } operation, ok := operations[workflow.Operation.ID] if !ok { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q operation %q is not generated", workflow.ID, workflow.Operation.ID) + return workflowView{}, fmt.Errorf("gosdk: workflow %q operation %q is not generated", workflow.ID, workflow.Operation.ID) } view := workflowView{ ID: workflow.ID, MethodName: e.names.Claim(modelgen.PublicName(workflow.ID), "Workflow"), @@ -471,7 +450,7 @@ func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string] } inputSchema, inputShape, err := workflowInputSchema(workflow.Inputs) if err != nil { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q inputs: %w", workflow.ID, err) + return workflowView{}, fmt.Errorf("gosdk: workflow %q inputs: %w", workflow.ID, err) } inputKey := "__workflow_" + view.InputType e.schemas.Set(inputKey, inputSchema) @@ -479,7 +458,7 @@ func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string] e.collectSchema(inputSchema) view.SuccessCondition, err = statusCondition(workflow.SuccessStatuses) if err != nil { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q success responses: %w", workflow.ID, err) + return workflowView{}, fmt.Errorf("gosdk: workflow %q success responses: %w", workflow.ID, err) } operationSuccess := make(map[string]struct{}, len(operation.SuccessStatuses)) for _, status := range operation.SuccessStatuses { @@ -487,7 +466,7 @@ func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string] } for _, status := range workflow.SuccessStatuses { if _, ok := operationSuccess[status]; !ok { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q accepts status %s but operation %q does not", workflow.ID, status, workflow.Operation.ID) + return workflowView{}, fmt.Errorf("gosdk: workflow %q accepts status %s but operation %q does not", workflow.ID, status, workflow.Operation.ID) } } declaredInputs := make(map[string]struct{}, inputShape.Properties.Len()) @@ -504,10 +483,10 @@ func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string] } parameter, exists := operationParameters[binding.In+"\x00"+binding.Name] if !exists { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q parameter %s %q is not generated", workflow.ID, binding.In, binding.Name) + return workflowView{}, fmt.Errorf("gosdk: workflow %q parameter %s %q is not generated", workflow.ID, binding.In, binding.Name) } if _, exists := declaredInputs[binding.Input]; !exists { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q parameter input %q is not declared by workflow inputs", workflow.ID, binding.Input) + return workflowView{}, fmt.Errorf("gosdk: workflow %q parameter input %q is not declared by workflow inputs", workflow.ID, binding.Input) } view.ParameterFields = append(view.ParameterFields, workflowAssignmentView{ Target: parameter.FieldName, Input: binding.Input, ExpectedType: parameter.Type, @@ -515,7 +494,7 @@ func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string] } if len(workflow.PayloadBindings) > 0 { if operation.Body == nil || !operation.Body.Required { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q requires a declared required JSON request body", workflow.ID) + return workflowView{}, fmt.Errorf("gosdk: workflow %q requires a declared required JSON request body", workflow.ID) } view.BodyType = operation.Body.Type for _, binding := range workflow.PayloadBindings { @@ -523,10 +502,10 @@ func (e *emitter) prepareWorkflow(workflow *sdk.Workflow, operations map[string] continue } if strings.Contains(binding.Input, ".") { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q payload input %q is nested; only direct workflow inputs are supported", workflow.ID, binding.Input) + return workflowView{}, fmt.Errorf("gosdk: workflow %q payload input %q is nested; only direct workflow inputs are supported", workflow.ID, binding.Input) } if _, exists := declaredInputs[binding.Input]; !exists { - return workflowView{}, fmt.Errorf("sdk/golang: workflow %q payload input %q is not declared by workflow inputs", workflow.ID, binding.Input) + return workflowView{}, fmt.Errorf("gosdk: workflow %q payload input %q is not declared by workflow inputs", workflow.ID, binding.Input) } view.BodyFields = append(view.BodyFields, workflowAssignmentView{ Input: binding.Input, ExpectedModel: operation.Body.Type, ExpectedField: binding.Property, @@ -573,16 +552,16 @@ func resolveWorkflowFieldTypes(types []*modelgen.GeneratedType, workflows []work workflow := &workflows[workflowIndex] inputType := typesByName[workflow.InputType] if inputType == nil || inputType.Kind != modelgen.KindObject { - return fmt.Errorf("sdk/golang: workflow %q inputs %q are not a struct model", workflow.ID, workflow.InputType) + return fmt.Errorf("gosdk: workflow %q inputs %q are not a struct model", workflow.ID, workflow.InputType) } for assignmentIndex := range workflow.ParameterFields { assignment := &workflow.ParameterFields[assignmentIndex] inputField, ok := generatedField(typesByName, inputType, assignment.Input, make(map[string]struct{})) if !ok { - return fmt.Errorf("sdk/golang: workflow %q input %q is not present on %s", workflow.ID, assignment.Input, workflow.InputType) + return fmt.Errorf("gosdk: workflow %q input %q is not present on %s", workflow.ID, assignment.Input, workflow.InputType) } if inputField.Type != assignment.ExpectedType { - return fmt.Errorf("sdk/golang: workflow %q input %q has type %s, operation parameter requires %s", workflow.ID, assignment.Input, inputField.Type, assignment.ExpectedType) + return fmt.Errorf("gosdk: workflow %q input %q has type %s, operation parameter requires %s", workflow.ID, assignment.Input, inputField.Type, assignment.ExpectedType) } assignment.Source = inputField.Name } @@ -590,15 +569,15 @@ func resolveWorkflowFieldTypes(types []*modelgen.GeneratedType, workflows []work assignment := &workflow.BodyFields[assignmentIndex] inputField, ok := generatedField(typesByName, inputType, assignment.Input, make(map[string]struct{})) if !ok { - return fmt.Errorf("sdk/golang: workflow %q input %q is not present on %s", workflow.ID, assignment.Input, workflow.InputType) + return fmt.Errorf("gosdk: workflow %q input %q is not present on %s", workflow.ID, assignment.Input, workflow.InputType) } bodyType := typesByName[assignment.ExpectedModel] bodyField, ok := generatedField(typesByName, bodyType, assignment.ExpectedField, make(map[string]struct{})) if !ok { - return fmt.Errorf("sdk/golang: workflow %q payload property %q is not present on %s", workflow.ID, assignment.ExpectedField, assignment.ExpectedModel) + return fmt.Errorf("gosdk: workflow %q payload property %q is not present on %s", workflow.ID, assignment.ExpectedField, assignment.ExpectedModel) } if inputField.Type != bodyField.Type { - return fmt.Errorf("sdk/golang: workflow %q input %q has type %s, payload property %q requires %s", workflow.ID, assignment.Input, inputField.Type, assignment.ExpectedField, bodyField.Type) + return fmt.Errorf("gosdk: workflow %q input %q has type %s, payload property %q requires %s", workflow.ID, assignment.Input, inputField.Type, assignment.ExpectedField, bodyField.Type) } assignment.Source = inputField.Name assignment.Target = bodyField.Name @@ -641,20 +620,20 @@ func (e *emitter) prepareResponses(operation *sdk.Operation) (string, []string, isSuccess := statusIsSuccess(status) _, schema, err := jsonMedia(response.Content) if err != nil { - return "", nil, nil, fmt.Errorf("sdk/golang: operation %q response %s: %w", operation.ID, status, err) + return "", nil, nil, fmt.Errorf("gosdk: operation %q response %s: %w", operation.ID, status, err) } typeName := "" if schema != nil { typeName, err = e.schemaType(schema, operation.ID+statusName(status)+"Response") if err != nil { - return "", nil, nil, fmt.Errorf("sdk/golang: operation %q response %s: %w", operation.ID, status, err) + return "", nil, nil, fmt.Errorf("gosdk: operation %q response %s: %w", operation.ID, status, err) } } if isSuccess { successes = append(successes, status) if typeName != "" { if selectedType != "" && selectedType != typeName { - return "", nil, nil, fmt.Errorf("sdk/golang: operation %q has incompatible success response types %s and %s", operation.ID, selectedType, typeName) + return "", nil, nil, fmt.Errorf("gosdk: operation %q has incompatible success response types %s and %s", operation.ID, selectedType, typeName) } selectedType = typeName } @@ -663,7 +642,7 @@ func (e *emitter) prepareResponses(operation *sdk.Operation) (string, []string, errorResponses = append(errorResponses, errorResponseView{Status: status, Type: typeName}) } if len(successes) == 0 { - return "", nil, nil, fmt.Errorf("sdk/golang: operation %q has no declared 2xx response", operation.ID) + return "", nil, nil, fmt.Errorf("gosdk: operation %q has no declared 2xx response", operation.ID) } if selectedType != "" { responseType = selectedType @@ -699,27 +678,15 @@ func (e *emitter) schemaType(proxy *highbase.SchemaProxy, hint string) (string, } } if len(types) == 1 { - switch types[0] { - case "string": - if mapped := e.models.ScalarType(schema.Format, "string"); mapped != "string" { + if goType, scalar := e.models.ScalarType(types[0], schema.Format); scalar { + // A string format mapped to a named Go type is emitted as a model so + // models.gen.go owns the import. + if types[0] == "string" && goType != "string" { return e.addInlineSchema(hint, proxy), nil } - return "string", nil - case "integer": - if schema.Format == "int32" { - return "int32", nil - } - if schema.Format == "int64" { - return "int64", nil - } - return "int", nil - case "number": - if schema.Format == "float" { - return "float32", nil - } - return "float64", nil - case "boolean": - return "bool", nil + return goType, nil + } + switch types[0] { case "array": if schema.Items == nil || !schema.Items.IsA() { return "[]any", nil @@ -789,13 +756,13 @@ func (e *emitter) collectSchema(proxy *highbase.SchemaProxy) { ref := proxy.GetReference() if name, component := componentSchemaRefName(ref); component { if err := e.ensureComponent(name); err != nil && e.collectErr == nil { - e.collectErr = fmt.Errorf("sdk/golang: %w", err) + e.collectErr = fmt.Errorf("gosdk: %w", err) } } else if e.collectErr == nil { if strings.HasPrefix(ref, "#/") { - e.collectErr = fmt.Errorf("sdk/golang: schema reference %q is not a components/schemas reference", ref) + e.collectErr = fmt.Errorf("gosdk: schema reference %q is not a components/schemas reference", ref) } else { - e.collectErr = fmt.Errorf("sdk/golang: external schema reference %q is not supported by SDK generation", ref) + e.collectErr = fmt.Errorf("gosdk: external schema reference %q is not supported by SDK generation", ref) } } return @@ -804,7 +771,14 @@ func (e *emitter) collectSchema(proxy *highbase.SchemaProxy) { if schema == nil { return } - visit := e.collectSchema + forEachShapeChild(schema, e.collectSchema) +} + +// forEachShapeChild visits the child schemas that generator/golang renders into +// Go shape. Validation-only keywords are deliberately absent: a $ref beneath +// one of them must not make a component reachable. Reachability and collection +// share this walk so the two can never disagree about what is generated. +func forEachShapeChild(schema *highbase.Schema, visit func(*highbase.SchemaProxy)) { for _, child := range schema.AllOf { visit(child) } @@ -906,7 +880,7 @@ func (e *emitter) parameterEncoder(proxy *highbase.SchemaProxy) (string, bool, e } switch nonNullType { case "string": - if mapped := e.models.ScalarType(schema.Format, "string"); mapped != "string" { + if mapped, _ := e.models.ScalarType("string", schema.Format); mapped != "string" { return "", false, fmt.Errorf("custom scalar format %q maps to %s and is not supported for parameter serialization", schema.Format, mapped) } return "encodeString", false, nil diff --git a/generator/sdk/golang/generator_test.go b/generator/sdk/gosdk/generator_test.go similarity index 88% rename from generator/sdk/golang/generator_test.go rename to generator/sdk/gosdk/generator_test.go index cb5b286a..0f140de4 100644 --- a/generator/sdk/golang/generator_test.go +++ b/generator/sdk/gosdk/generator_test.go @@ -62,7 +62,8 @@ func TestGeneratedSDKExecutesResourceOperations(t *testing.T) { } } for _, expected := range []string{ - "buffer.Grow(capacity)", + "buffer.Grow(length + bytes.MinRead)", + "target := request", "var securityWidgetsList = [][]securityRequirement", "response.StatusCode == 201", `input.Body, "application/json"`, @@ -844,4 +845,79 @@ func TestGeneratedClient(t *testing.T) { close(errorsSeen) for err := range errorsSeen { t.Error(err) } } + +type doerFunc func(*http.Request) (*http.Response, error) + +func (do doerFunc) Do(request *http.Request) (*http.Response, error) { return do(request) } + +func unavailableCredential() Credential { + return CredentialFunc(func(context.Context, SecurityScheme, []string, *http.Request) error { + return errors.New("tenant key unavailable") + }) +} + +// A credential that fails part-way through one alternative must not leave its +// partial writes on the request the fallback alternative sends. +func TestFailedAlternativeDoesNotLeakIntoFallback(t *testing.T) { + var authorization, query string + recorder := doerFunc(func(request *http.Request) (*http.Response, error) { + authorization, query = request.Header.Get("Authorization"), request.URL.RawQuery + return &http.Response{StatusCode: http.StatusOK, Header: make(http.Header), Body: io.NopCloser(strings.NewReader("{}"))}, nil + }) + client, err := NewClient("https://api.example.test/v1", WithHTTPClient(recorder), + WithCredential("bearerAuth", BearerToken("token")), + WithCredential("tenantKey", unavailableCredential()), + WithCredential("serviceKey", APIKey("service")), + ) + if err != nil { t.Fatal(err) } + if _, err := client.Widgets.List(context.Background(), &ListWidgetsParams{TenantID: "tenant-1"}); err != nil { + t.Fatalf("fallback credential alternative failed: %v", err) + } + if authorization != "" || query != "api_key=service" { + t.Fatalf("failed alternative leaked into the fallback request: Authorization=%q query=%q", authorization, query) + } + + withoutFallback, err := NewClient("https://api.example.test/v1", WithHTTPClient(recorder), + WithCredential("bearerAuth", BearerToken("token")), + WithCredential("tenantKey", unavailableCredential()), + ) + if err != nil { t.Fatal(err) } + if _, err := withoutFallback.Widgets.List(context.Background(), &ListWidgetsParams{TenantID: "tenant-1"}); err == nil || !strings.Contains(err.Error(), "tenant key unavailable") { + t.Fatalf("expected the credential failure to surface, got %v", err) + } +} + +// Content-Length is a preallocation hint and never a bound: the limit holds and +// the body arrives whole whether the length is declared, absent, or wrong. +func TestResponseBodyReadPaths(t *testing.T) { + payload := "{\"items\":[]}" + for _, test := range []struct { + name string + contentLength int64 + body string + wantErr string + }{ + {name: "declared length", contentLength: int64(len(payload)), body: payload}, + {name: "unknown length", contentLength: -1, body: payload}, + {name: "understated length", contentLength: 2, body: payload}, + {name: "declared length over the limit", contentLength: 65, body: strings.Repeat(" ", 65), wantErr: "exceeds 64 bytes"}, + {name: "unknown length over the limit", contentLength: -1, body: strings.Repeat(" ", 65), wantErr: "exceeds 64 bytes"}, + } { + t.Run(test.name, func(t *testing.T) { + client, err := NewClient("https://api.example.test/v1", WithMaxResponseBody(64), + WithCredential("serviceKey", APIKey("service")), + WithHTTPClient(doerFunc(func(*http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, ContentLength: test.contentLength, Header: make(http.Header), Body: io.NopCloser(strings.NewReader(test.body))}, nil + })), + ) + if err != nil { t.Fatal(err) } + page, err := client.Widgets.List(context.Background(), &ListWidgetsParams{TenantID: "tenant-1"}) + if test.wantErr != "" { + if err == nil || !strings.Contains(err.Error(), test.wantErr) { t.Fatalf("expected %q, got %v", test.wantErr, err) } + return + } + if err != nil || page.Value.Items == nil { t.Fatalf("body did not arrive whole: %#v %v", page, err) } + }) + } +} ` diff --git a/generator/sdk/golang/options.go b/generator/sdk/gosdk/options.go similarity index 100% rename from generator/sdk/golang/options.go rename to generator/sdk/gosdk/options.go diff --git a/generator/sdk/golang/templates/client.tmpl b/generator/sdk/gosdk/templates/client.tmpl similarity index 87% rename from generator/sdk/golang/templates/client.tmpl rename to generator/sdk/gosdk/templates/client.tmpl index abf481ec..ff7db90b 100644 --- a/generator/sdk/golang/templates/client.tmpl +++ b/generator/sdk/gosdk/templates/client.tmpl @@ -1,4 +1,4 @@ -// Code generated by libopenapi generator/sdk/golang. DO NOT EDIT. +// Code generated by libopenapi generator/sdk/gosdk. DO NOT EDIT. package {{.PackageName}} @@ -320,37 +320,39 @@ func (client *Client) authorize(ctx context.Context, alternatives [][]securityRe return nil } var credentialError error - for _, alternative := range alternatives { + for index, alternative := range alternatives { if len(alternative) == 0 { return nil } - complete := true - for _, requirement := range alternative { - if client.credentials[requirement.Name] == nil { - complete = false - break - } - } - if !complete { + if !client.hasCredentials(alternative) { continue } - candidate := request.Clone(ctx) + // A failed Apply can leave partial writes on the request. That only + // matters when a later alternative could still run, so the request is + // cloned for that case alone and credentials are otherwise applied in place. + target := request + if client.hasFallback(alternatives[index+1:]) { + target = request.Clone(ctx) + } + applied := true for _, requirement := range alternative { scheme, ok := securitySchemes[requirement.Name] if !ok { return fmt.Errorf("security scheme %q is not defined", requirement.Name) } - if err := client.credentials[requirement.Name].Apply(ctx, scheme, requirement.Scopes, candidate); err != nil { + if err := client.credentials[requirement.Name].Apply(ctx, scheme, requirement.Scopes, target); err != nil { credentialError = fmt.Errorf("apply security scheme %q: %w", requirement.Name, err) - complete = false + applied = false break } } - if !complete { + if !applied { continue } - request.Header = candidate.Header - request.URL = candidate.URL + if target != request { + request.Header = target.Header + request.URL = target.URL + } return nil } if credentialError != nil { @@ -359,6 +361,26 @@ func (client *Client) authorize(ctx context.Context, alternatives [][]securityRe return errors.New("no complete credential alternative is configured") } +func (client *Client) hasCredentials(alternative []securityRequirement) bool { + for _, requirement := range alternative { + if client.credentials[requirement.Name] == nil { + return false + } + } + return true +} + +// hasFallback reports whether any remaining alternative could still authorize +// the request: an empty requirement, or one whose credentials are all configured. +func (client *Client) hasFallback(alternatives [][]securityRequirement) bool { + for _, alternative := range alternatives { + if client.hasCredentials(alternative) { + return true + } + } + return false +} + func (client *Client) do(request *http.Request, alternatives [][]securityRequirement) (*http.Response, []byte, error) { if err := client.authorize(request.Context(), alternatives, request); err != nil { return nil, nil, err @@ -372,25 +394,34 @@ func (client *Client) do(request *http.Request, alternatives [][]securityRequire if readLimit < int64(^uint64(0)>>1) { readLimit++ } - limited := io.LimitReader(response.Body, readLimit) - var buffer bytes.Buffer - if response.ContentLength > 0 && response.ContentLength <= client.maxResponseBody { - capacity := int(response.ContentLength) - if int64(capacity) == response.ContentLength { - buffer.Grow(capacity) - } - } - _, err = buffer.ReadFrom(limited) + body, err := readBody(io.LimitReader(response.Body, readLimit), response.ContentLength, client.maxResponseBody) if err != nil { return nil, nil, fmt.Errorf("read response body: %w", err) } - body := buffer.Bytes() if int64(len(body)) > client.maxResponseBody { return nil, nil, fmt.Errorf("response body exceeds %d bytes", client.maxResponseBody) } return response, body, nil } +// readBody buffers a response body. bytes.Buffer.ReadFrom reserves +// bytes.MinRead before every read, including the final one that reports EOF, so +// a declared length is preallocated with that headroom. Without it the buffer +// reallocates anyway and a small body costs more than io.ReadAll, which is the +// path taken when the length is unknown, over the limit, or not addressable. +func readBody(reader io.Reader, contentLength, limit int64) ([]byte, error) { + length := int(contentLength) + if contentLength <= 0 || contentLength > limit || int64(length) != contentLength || length+bytes.MinRead < length { + return io.ReadAll(reader) + } + var buffer bytes.Buffer + buffer.Grow(length + bytes.MinRead) + if _, err := buffer.ReadFrom(reader); err != nil { + return nil, err + } + return buffer.Bytes(), nil +} + // Response contains a decoded value and its HTTP metadata. type Response[T any] struct { StatusCode int diff --git a/generator/sdk/golang/templates/resources.tmpl b/generator/sdk/gosdk/templates/resources.tmpl similarity index 98% rename from generator/sdk/golang/templates/resources.tmpl rename to generator/sdk/gosdk/templates/resources.tmpl index f80987b4..7d5ce68a 100644 --- a/generator/sdk/golang/templates/resources.tmpl +++ b/generator/sdk/gosdk/templates/resources.tmpl @@ -1,4 +1,4 @@ -// Code generated by libopenapi generator/sdk/golang. DO NOT EDIT. +// Code generated by libopenapi generator/sdk/gosdk. DO NOT EDIT. package {{.PackageName}} diff --git a/generator/sdk/golang/templates/workflows.tmpl b/generator/sdk/gosdk/templates/workflows.tmpl similarity index 94% rename from generator/sdk/golang/templates/workflows.tmpl rename to generator/sdk/gosdk/templates/workflows.tmpl index a8433c5f..6c4bfc01 100644 --- a/generator/sdk/golang/templates/workflows.tmpl +++ b/generator/sdk/gosdk/templates/workflows.tmpl @@ -1,4 +1,4 @@ -// Code generated by libopenapi generator/sdk/golang. DO NOT EDIT. +// Code generated by libopenapi generator/sdk/gosdk. DO NOT EDIT. package {{.PackageName}} diff --git a/generator/sdk/golang/view.go b/generator/sdk/gosdk/view.go similarity index 100% rename from generator/sdk/golang/view.go rename to generator/sdk/gosdk/view.go