From 888c4d9f15a8f7d932ed7f4ebbe99d3ee2907bf1 Mon Sep 17 00:00:00 2001 From: quobix Date: Mon, 21 Sep 2026 22:56:34 -0400 Subject: [PATCH] fix(gosdk): reject empty JSON responses --- generator/sdk/gosdk/coverage_test.go | 42 ++++++++++++++-- generator/sdk/gosdk/generator.go | 53 +++++++++++++++++--- generator/sdk/gosdk/generator_test.go | 21 +++++++- generator/sdk/gosdk/templates/client.tmpl | 2 +- generator/sdk/gosdk/templates/resources.tmpl | 2 + generator/sdk/gosdk/view.go | 1 + 6 files changed, 107 insertions(+), 14 deletions(-) diff --git a/generator/sdk/gosdk/coverage_test.go b/generator/sdk/gosdk/coverage_test.go index d729f737..1a199e7b 100644 --- a/generator/sdk/gosdk/coverage_test.go +++ b/generator/sdk/gosdk/coverage_test.go @@ -307,27 +307,58 @@ func TestPrepareResponsesContracts(t *testing.T) { for _, test := range tests { t.Run(test.name, func(t *testing.T) { operation := &sdk.Operation{ID: "read", Responses: test.responses} - _, _, _, err := emitter.prepareResponses(operation) + _, _, _, _, err := emitter.prepareResponses(operation) if err == nil || !strings.Contains(err.Error(), test.want) { t.Fatalf("expected error containing %q, got %v", test.want, err) } }) } - responseType, successes, failures, err := emitter.prepareResponses(&sdk.Operation{ID: "read", Responses: []*sdk.Response{ + responseType, successes, decodedSuccesses, failures, err := emitter.prepareResponses(&sdk.Operation{ID: "read", Responses: []*sdk.Response{ nil, {Status: "200", Content: jsonContent(schema("string", ""))}, {Status: "default", Content: jsonContent(schema("object", ""))}, }}) - if err != nil || responseType != "string" || len(successes) != 1 || len(failures) != 1 || failures[0].Status != "default" { - t.Fatalf("unexpected prepared responses: %q %#v %#v %v", responseType, successes, failures, err) + if err != nil || responseType != "string" || len(successes) != 1 || len(decodedSuccesses) != 1 || decodedSuccesses[0] != "200" || len(failures) != 1 || failures[0].Status != "default" { + t.Fatalf("unexpected prepared responses: %q %#v %#v %#v %v", responseType, successes, decodedSuccesses, failures, err) } - _, successes, _, err = emitter.prepareResponses(&sdk.Operation{ID: "wildcard", Responses: []*sdk.Response{{Status: "2xx"}}}) + _, successes, _, _, err = emitter.prepareResponses(&sdk.Operation{ID: "wildcard", Responses: []*sdk.Response{{Status: "2xx"}}}) if err != nil || len(successes) != 1 || successes[0] != "2XX" { t.Fatalf("lowercase wildcard was not normalized: %#v, %v", successes, err) } } +func TestDecodeStatusConditionHonorsExactBodylessOverride(t *testing.T) { + tests := []struct { + name string + decoded []string + successes []string + want string + wantError bool + }{ + {name: "exact override", decoded: []string{"2XX"}, successes: []string{"2XX", "204"}, want: "(response.StatusCode/100 == 2) && response.StatusCode != 204"}, + {name: "decoded exact", decoded: []string{"200"}, successes: []string{"200", "204"}, want: "response.StatusCode == 200"}, + {name: "nonoverlapping exact", decoded: []string{"2XX"}, successes: []string{"2XX", "304"}, want: "response.StatusCode/100 == 2"}, + {name: "irrelevant malformed status", decoded: []string{"2XX"}, successes: []string{"2XX", "bad"}, want: "response.StatusCode/100 == 2"}, + {name: "invalid decoded status", decoded: []string{"bad"}, wantError: true}, + {name: "invalid overlapping exact", decoded: []string{"2XX"}, successes: []string{"2A4"}, wantError: true}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + condition, err := decodeStatusCondition(test.decoded, test.successes) + if test.wantError { + if err == nil { + t.Fatalf("condition = %q, error = nil", condition) + } + return + } + if err != nil || condition != test.want { + t.Fatalf("condition = %q, error = %v, want %q", condition, err, test.want) + } + }) + } +} + func TestPrepareWorkflowContracts(t *testing.T) { operation := operationView{ ID: "create", ResourceName: "Things", MethodName: "Create", ParamsType: "CreateParams", @@ -718,6 +749,7 @@ func TestPrepareOperationRejectsInvalidResponseStatuses(t *testing.T) { want string }{ {name: "success", responses: []*sdk.Response{{Status: "2a0"}}, want: "success responses"}, + {name: "JSON success", responses: []*sdk.Response{{Status: "2a0", Content: jsonContent(schema("string", ""))}}, want: "JSON success responses"}, {name: "error", responses: []*sdk.Response{{Status: "204"}, {Status: "oops"}}, want: "error response"}, } { t.Run(test.name, func(t *testing.T) { diff --git a/generator/sdk/gosdk/generator.go b/generator/sdk/gosdk/generator.go index 70aff48d..12fe2283 100644 --- a/generator/sdk/gosdk/generator.go +++ b/generator/sdk/gosdk/generator.go @@ -399,12 +399,18 @@ func (e *emitter) prepareOperation(operation *sdk.Operation, methodName string) view.Body = &bodyView{Type: typeName, FieldType: fieldType, Required: operation.RequestBody.Required, ContentType: mediaType, Description: operation.RequestBody.Description} } view.HasParameters = len(view.Parameters) > 0 || view.Body != nil - responseType, successes, errorResponses, err := e.prepareResponses(operation) + responseType, successes, decodedSuccesses, errorResponses, err := e.prepareResponses(operation) if err != nil { return view, err } view.ResponseType = responseType view.SuccessStatuses = successes + if len(decodedSuccesses) > 0 { + view.DecodeCondition, err = decodeStatusCondition(decodedSuccesses, successes) + if err != nil { + return view, fmt.Errorf("gosdk: operation %q JSON success responses: %w", operation.ID, err) + } + } view.SuccessCondition, err = statusCondition(successes) if err != nil { return view, fmt.Errorf("gosdk: operation %q success responses: %w", operation.ID, err) @@ -607,10 +613,11 @@ func generatedField(types map[string]*modelgen.GeneratedType, generatedType *mod return modelgen.GeneratedField{}, false } -func (e *emitter) prepareResponses(operation *sdk.Operation) (string, []string, []errorResponseView, error) { +func (e *emitter) prepareResponses(operation *sdk.Operation) (string, []string, []string, []errorResponseView, error) { responseType := "struct{}" var selectedType string var successes []string + var decodedSuccesses []string var errorResponses []errorResponseView for _, response := range operation.Responses { if response == nil { @@ -620,20 +627,21 @@ 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("gosdk: operation %q response %s: %w", operation.ID, status, err) + return "", nil, 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("gosdk: operation %q response %s: %w", operation.ID, status, err) + return "", nil, nil, nil, fmt.Errorf("gosdk: operation %q response %s: %w", operation.ID, status, err) } } if isSuccess { successes = append(successes, status) if typeName != "" { + decodedSuccesses = append(decodedSuccesses, status) if selectedType != "" && selectedType != typeName { - return "", nil, nil, fmt.Errorf("gosdk: operation %q has incompatible success response types %s and %s", operation.ID, selectedType, typeName) + return "", nil, nil, nil, fmt.Errorf("gosdk: operation %q has incompatible success response types %s and %s", operation.ID, selectedType, typeName) } selectedType = typeName } @@ -642,12 +650,12 @@ 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("gosdk: operation %q has no declared 2xx response", operation.ID) + return "", nil, nil, nil, fmt.Errorf("gosdk: operation %q has no declared 2xx response", operation.ID) } if selectedType != "" { responseType = selectedType } - return responseType, successes, errorResponses, nil + return responseType, successes, decodedSuccesses, errorResponses, nil } func (e *emitter) schemaType(proxy *highbase.SchemaProxy, hint string) (string, error) { @@ -932,6 +940,37 @@ func statusCondition(statuses []string) (string, error) { return strings.Join(conditions, " || "), nil } +func decodeStatusCondition(decoded, successes []string) (string, error) { + condition, err := statusCondition(decoded) + if err != nil { + return "", err + } + decodedStatuses := make(map[string]struct{}, len(decoded)) + decodedRanges := make(map[byte]struct{}) + for _, status := range decoded { + canonical := strings.ToUpper(status) + decodedStatuses[canonical] = struct{}{} + if len(canonical) == 3 && canonical[1:] == "XX" { + decodedRanges[canonical[0]] = struct{}{} + } + } + for _, status := range successes { + canonical := strings.ToUpper(status) + if _, ok := decodedStatuses[canonical]; ok || len(canonical) != 3 || canonical[1:] == "XX" { + continue + } + if _, overlapsRange := decodedRanges[canonical[0]]; !overlapsRange { + continue + } + code, parseErr := strconv.Atoi(canonical) + if parseErr != nil || code < 100 || code > 599 { + return "", fmt.Errorf("unsupported response status %q", status) + } + condition = fmt.Sprintf("(%s) && response.StatusCode != %d", condition, code) + } + return condition, nil +} + func statusIsSuccess(status string) bool { return len(status) == 3 && status[0] == '2' } diff --git a/generator/sdk/gosdk/generator_test.go b/generator/sdk/gosdk/generator_test.go index 0f140de4..9af522c4 100644 --- a/generator/sdk/gosdk/generator_test.go +++ b/generator/sdk/gosdk/generator_test.go @@ -405,10 +405,11 @@ paths: in: header schema: {type: integer, format: int32} responses: - "200": + "2XX": description: widgets content: application/json: {schema: {$ref: "#/components/schemas/WidgetPage"}} + "204": {description: no widgets} "400": description: bad request content: @@ -703,6 +704,14 @@ func TestGeneratedClient(t *testing.T) { _, _ = io.WriteString(w, "{") return } + if r.URL.Query().Get("limit") == "96" { + w.WriteHeader(http.StatusNoContent) + return + } + if r.URL.Query().Get("limit") == "95" { + w.Header().Set("X-Correlation-ID", "empty-success") + return + } w.Header().Set("X-Page", "one") _, _ = io.WriteString(w, ` + "`" + `{"items":[{"id":"a-1"}]}` + "`" + `) case r.Method == http.MethodGet && strings.Contains(r.URL.EscapedPath(), "/widgets/"): @@ -809,6 +818,16 @@ func TestGeneratedClient(t *testing.T) { if !errors.As(err, &decodeError) || decodeError.StatusCode != http.StatusOK || decodeError.Header.Get("X-Correlation-ID") != "broken-success" || string(decodeError.Body) != "{" || errors.Unwrap(decodeError) == nil { t.Fatalf("malformed success lost response context: %#v %v", decodeError, err) } + bodylessSuccess := 96 + response, err := client.Widgets.List(context.Background(), &ListWidgetsParams{TenantID: "tenant-1", Limit: &bodylessSuccess}) + if err != nil || response.StatusCode != http.StatusNoContent { + t.Fatalf("bodyless declared success failed: %#v %v", response, err) + } + emptySuccess := 95 + _, err = client.Widgets.List(context.Background(), &ListWidgetsParams{TenantID: "tenant-1", Limit: &emptySuccess}) + if !errors.As(err, &decodeError) || decodeError.StatusCode != http.StatusOK || decodeError.Header.Get("X-Correlation-ID") != "empty-success" || len(decodeError.Body) != 0 || errors.Unwrap(decodeError) == nil { + t.Fatalf("empty success was not rejected with response context: %#v %v", decodeError, err) + } for _, status := range []int{http.StatusTemporaryRedirect, http.StatusPermanentRedirect} { var followed atomic.Int32 diff --git a/generator/sdk/gosdk/templates/client.tmpl b/generator/sdk/gosdk/templates/client.tmpl index ff7db90b..3bc9a1b4 100644 --- a/generator/sdk/gosdk/templates/client.tmpl +++ b/generator/sdk/gosdk/templates/client.tmpl @@ -482,7 +482,7 @@ func NullableValue[T any](value T) **T { func decodeJSON(body []byte, target any) error { if len(body) == 0 { - return nil + return errors.New("decode response body: empty body") } if err := json.Unmarshal(body, target); err != nil { return fmt.Errorf("decode response body: %w", err) diff --git a/generator/sdk/gosdk/templates/resources.tmpl b/generator/sdk/gosdk/templates/resources.tmpl index 7d5ce68a..f9c83729 100644 --- a/generator/sdk/gosdk/templates/resources.tmpl +++ b/generator/sdk/gosdk/templates/resources.tmpl @@ -108,9 +108,11 @@ func (resource *{{$resource.TypeName}}) {{.MethodName}}(ctx context.Context{{if if {{.SuccessCondition}} { result := &Response[{{.ResponseType}}]{StatusCode: response.StatusCode, Header: header} {{- if ne .ResponseType "struct{}"}} + if {{.DecodeCondition}} { if err := decodeJSON(payload, &result.Value); err != nil { return nil, &ResponseDecodeError{StatusCode: response.StatusCode, Header: header, Body: payload, Cause: err} } + } {{- end}} return result, nil } diff --git a/generator/sdk/gosdk/view.go b/generator/sdk/gosdk/view.go index 5459e653..d3e81b1b 100644 --- a/generator/sdk/gosdk/view.go +++ b/generator/sdk/gosdk/view.go @@ -41,6 +41,7 @@ type operationView struct { ResponseType string SuccessStatuses []string SuccessCondition string + DecodeCondition string ErrorResponses []errorResponseView Security [][]securityRequirementView }