From 69f303ad4f9d704442965813d8af2357080c9f14 Mon Sep 17 00:00:00 2001 From: Mike Sawka Date: Fri, 25 Sep 2026 18:45:50 +0000 Subject: [PATCH] fix(web): stop handlers after writing error responses --- pkg/web/web.go | 7 +++- pkg/web/web_test.go | 81 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 87 insertions(+), 1 deletion(-) create mode 100644 pkg/web/web_test.go diff --git a/pkg/web/web.go b/pkg/web/web.go index 106db981e4..46dd866627 100644 --- a/pkg/web/web.go +++ b/pkg/web/web.go @@ -1,4 +1,4 @@ -// Copyright 2025, Command Line Inc. +// Copyright 2026, Command Line Inc. // SPDX-License-Identifier: Apache-2.0 package web @@ -119,12 +119,14 @@ func handleService(w http.ResponseWriter, r *http.Request) { err = json.Unmarshal(bodyData, &webCall) if err != nil { http.Error(w, fmt.Sprintf("invalid request body: %v", err), http.StatusBadRequest) + return } rtn := service.CallService(r.Context(), webCall) jsonRtn, err := json.Marshal(rtn) if err != nil { http.Error(w, fmt.Sprintf("error serializing response: %v", err), http.StatusInternalServerError) + return } w.Header().Set(ContentTypeHeaderKey, ContentTypeJson) w.Header().Set(ContentLengthHeaderKey, fmt.Sprintf("%d", len(jsonRtn))) @@ -157,6 +159,7 @@ func handleWaveFile(w http.ResponseWriter, r *http.Request) { offset, err = strconv.ParseInt(offsetStr, 10, 64) if err != nil { http.Error(w, fmt.Sprintf("invalid offset: %v", err), http.StatusBadRequest) + return } } if _, err := uuid.Parse(zoneId); err != nil { @@ -180,6 +183,7 @@ func handleWaveFile(w http.ResponseWriter, r *http.Request) { jsonFileBArr, err := json.Marshal(file) if err != nil { http.Error(w, fmt.Sprintf("error serializing file info: %v", err), http.StatusInternalServerError) + return } // can make more efficient by checking modtime + If-Modified-Since headers to allow caching dataStartIdx := file.DataStartIdx() @@ -239,6 +243,7 @@ func handleLocalStreamFile(w http.ResponseWriter, r *http.Request, path string, path, err := wavebase.ExpandHomeDir(path) if err != nil { http.Error(w, err.Error(), http.StatusBadRequest) + return } http.ServeFile(w, r, path) } diff --git a/pkg/web/web_test.go b/pkg/web/web_test.go new file mode 100644 index 0000000000..aa334c41ac --- /dev/null +++ b/pkg/web/web_test.go @@ -0,0 +1,81 @@ +// Copyright 2026, Command Line Inc. +// SPDX-License-Identifier: Apache-2.0 + +package web + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/wavetermdev/waveterm/pkg/service" +) + +type handlerTestService struct { + calls int +} + +func (svc *handlerTestService) Invoke() { + svc.calls++ +} + +func (svc *handlerTestService) Unserializable() chan int { + return make(chan int) +} + +func TestHandleServiceRejectsInvalidBodyWithoutCallingService(t *testing.T) { + svc := &handlerTestService{} + service.ServiceMap["handler-test"] = svc + t.Cleanup(func() { delete(service.ServiceMap, "handler-test") }) + + body := `{"service":"handler-test","method":"Invoke","args":[],"uicontext":123}` + var partialCall service.WebCallType + if err := json.Unmarshal([]byte(body), &partialCall); err == nil || partialCall.Service != "handler-test" || partialCall.Method != "Invoke" { + t.Fatalf("invalid test body did not partially decode a service call: %+v, %v", partialCall, err) + } + + req := httptest.NewRequest(http.MethodPost, "/wave/service", strings.NewReader(body)) + resp := httptest.NewRecorder() + handleService(resp, req) + + if resp.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want 400", resp.Code) + } + if svc.calls != 0 { + t.Fatalf("service invoked %d times for invalid body", svc.calls) + } + if !strings.HasPrefix(resp.Body.String(), "invalid request body:") || strings.Count(resp.Body.String(), "\n") != 1 { + t.Fatalf("unexpected response body: %q", resp.Body.String()) + } +} + +func TestHandleServiceStopsAfterMarshalError(t *testing.T) { + service.ServiceMap["handler-test"] = &handlerTestService{} + t.Cleanup(func() { delete(service.ServiceMap, "handler-test") }) + + req := httptest.NewRequest(http.MethodPost, "/wave/service", strings.NewReader(`{"service":"handler-test","method":"Unserializable","args":[]}`)) + resp := httptest.NewRecorder() + handleService(resp, req) + + if resp.Code != http.StatusInternalServerError { + t.Fatalf("status = %d, want 500", resp.Code) + } + if got := resp.Header().Get(ContentLengthHeaderKey); got != "" { + t.Fatalf("Content-Length = %q after error response, want unset", got) + } +} + +func TestHandleWaveFileRejectsInvalidOffsetFirst(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/wave/file?offset=invalid&zoneid=invalid&name=test", nil) + resp := httptest.NewRecorder() + handleWaveFile(resp, req) + + if resp.Code != http.StatusBadRequest { + t.Fatalf("status = %d, want 400", resp.Code) + } + if !strings.HasPrefix(resp.Body.String(), "invalid offset:") || strings.Count(resp.Body.String(), "\n") != 1 { + t.Fatalf("unexpected response body: %q", resp.Body.String()) + } +}