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