diff --git a/coderd/cachecompress/compress.go b/coderd/cachecompress/compress.go index 4f3ba5314d..9adff6a4de 100644 --- a/coderd/cachecompress/compress.go +++ b/coderd/cachecompress/compress.go @@ -240,9 +240,7 @@ func (c *Compressor) serveRef(w http.ResponseWriter, r *http.Request, headers ht } for key, values := range headers { - for _, value := range values { - w.Header().Add(key, value) - } + w.Header()[key] = values } w.Header().Set("Content-Encoding", cref.key.encoding) w.Header().Add("Vary", "Accept-Encoding") diff --git a/coderd/cachecompress/compress_internal_test.go b/coderd/cachecompress/compress_internal_test.go index 2f90ea8d12..b4756614ba 100644 --- a/coderd/cachecompress/compress_internal_test.go +++ b/coderd/cachecompress/compress_internal_test.go @@ -155,6 +155,41 @@ type nopEncoder struct { func (nopEncoder) Close() error { return nil } +func TestCompressorPresetHeaders(t *testing.T) { + t.Parallel() + + logger := testutil.Logger(t) + tempDir := t.TempDir() + cacheDir := filepath.Join(tempDir, "cache") + err := os.MkdirAll(cacheDir, 0o700) + require.NoError(t, err) + srcDir := filepath.Join(tempDir, "src") + err = os.MkdirAll(srcDir, 0o700) + require.NoError(t, err) + err = os.WriteFile(filepath.Join(srcDir, "file.html"), []byte("textstring"), 0o600) + require.NoError(t, err) + + compressor := NewCompressor(logger, 5, cacheDir, http.FS(os.DirFS(srcDir))) + + for range 2 { + ctx := testutil.Context(t, testutil.WaitShort) + req := httptest.NewRequestWithContext(ctx, "GET", "/file.html", nil) + req.Header.Set("Accept-Encoding", "gzip") + + respRec := httptest.NewRecorder() + respRec.Header().Set("X-Original-Content-Length", "10") + respRec.Header().Set("ETag", `"abc123"`) + + compressor.ServeHTTP(respRec, req) + resp := respRec.Result() + + require.Equal(t, http.StatusOK, resp.StatusCode) + require.Equal(t, []string{"10"}, resp.Header.Values("X-Original-Content-Length")) + require.Equal(t, []string{`"abc123"`}, resp.Header.Values("ETag")) + require.NoError(t, resp.Body.Close()) + } +} + // nolint: tparallel // we want to assert the state of the cache, so run synchronously func TestCompressorHeadings(t *testing.T) { t.Parallel() diff --git a/site/site_test.go b/site/site_test.go index 3427a70129..3527f31106 100644 --- a/site/site_test.go +++ b/site/site_test.go @@ -562,7 +562,7 @@ func TestServingBin(t *testing.T) { } if tr.wantEtag != "" { - assert.NotEmpty(t, resp.Header.Get("ETag"), "etag header is empty") + assert.Equal(t, []string{tr.wantEtag}, resp.Header.Values("ETag"), "etag header values did not match") assert.Equal(t, tr.wantEtag, resp.Header.Get("ETag"), "etag did not match") } @@ -570,6 +570,8 @@ func TestServingBin(t *testing.T) { // This is a custom header that we set to help the // client know the size of the decompressed data. See // the comment in site.go. + headerValues := resp.Header.Values("X-Original-Content-Length") + assert.Len(t, headerValues, 1, "X-Original-Content-Length should have exactly one value") headerStr := resp.Header.Get("X-Original-Content-Length") assert.NotEmpty(t, headerStr, "X-Original-Content-Length header is empty") originalSize, err := strconv.Atoi(headerStr)