package handlers_test import ( "compress/gzip" "crypto/rand" "encoding/json" "errors" "io" "net/http" "net/http/httptest" "net/url" "sync" "testing" "time" chimw "github.com/go-chi/chi/middleware" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "sneak.berlin/go/webhooker/internal/database" "sneak.berlin/go/webhooker/internal/delivery" ) // errClientGone is the write failure of a client that has gone away. var errClientGone = errors.New("client gone") // downloadPath is the archive download route of a target. func downloadPath(webhookID, targetID string) string { return "/hook/" + webhookID + "/targets/" + targetID + "/download" } // renameTarget submits the edit form renaming a target to Renamed. func renameTarget( env *sourceTestEnv, webhookID, targetID string, ) *httptest.ResponseRecorder { form := url.Values{} form.Set("name", "Renamed") return submitTargetEdit(env, webhookID, targetID, form) } // TestHandleTargetDownload proves a database target's archive // downloads as a gzipped JSON attachment named for the webhook, the // target and the time, here with no archive file yet, so with no // rows; and that a target of another type has no download. func TestHandleTargetDownload(t *testing.T) { t.Parallel() env := setupSourceTest(t) wh := seedWebhookWithRetention(t, env.db, 7) archive := seedTarget(t, env.db, wh.ID, database.TargetTypeDatabase) logTarget := seedTarget(t, env.db, wh.ID, database.TargetTypeLog) w := serveTarget( env, http.MethodGet, downloadPath(wh.ID, archive.ID), nil, ) require.Equal(t, http.StatusOK, w.Code, w.Body.String()) assert.Equal(t, "application/gzip", w.Header().Get("Content-Type")) assert.Regexp(t, `^attachment; filename="archive-seeded-t-database-`+ `\d{8}T\d{6}Z\.json\.gz"$`, w.Header().Get("Content-Disposition"), ) zr, err := gzip.NewReader(w.Body) require.NoError(t, err) var got map[string]json.RawMessage require.NoError(t, json.NewDecoder(zr).Decode(&got)) assert.JSONEq(t, `{"id":"`+archive.ID+`","name":"t-database"}`, string(got["target"]), ) assert.JSONEq(t, `[]`, string(got["archived_events"])) w = serveTarget( env, http.MethodGet, downloadPath(wh.ID, logTarget.ID), nil, ) assert.Equal(t, http.StatusNotFound, w.Code) } // TestHandleTargetDownload_WaitsForRename proves a download reads the // target's names and opens its archive under the lock a rename holds: // started while an edit is renaming the archive, it waits, and is // named for the target's new name. func TestHandleTargetDownload_WaitsForRename(t *testing.T) { t.Parallel() env := setupSourceTest(t) wh := seedWebhookWithRetention(t, env.db, 7) archive := seedTarget(t, env.db, wh.ID, database.TargetTypeDatabase) renaming, release := env.archives.BlockNextRename() edited := make(chan *httptest.ResponseRecorder, 1) go func() { edited <- renameTarget(env, wh.ID, archive.ID) }() <-renaming downloaded := make(chan *httptest.ResponseRecorder, 1) go func() { downloaded <- serveTarget( env, http.MethodGet, downloadPath(wh.ID, archive.ID), nil, ) }() select { case <-downloaded: release() t.Fatal("the download did not wait for the rename") case <-time.After(100 * time.Millisecond): } release() require.Equal(t, http.StatusSeeOther, (<-edited).Code) w := <-downloaded require.Equal(t, http.StatusOK, w.Code) assert.Contains(t, w.Header().Get("Content-Disposition"), "archive-seeded-renamed-", ) } // stalledWriter is a response writer whose first write waits until // resume is closed, closing writing when it starts to wait. type stalledWriter struct { *httptest.ResponseRecorder once sync.Once writing chan struct{} resume chan struct{} } func (s *stalledWriter) Write(b []byte) (int, error) { s.once.Do(func() { close(s.writing) <-s.resume }) return s.ResponseRecorder.Write(b) } // TestHandleTargetDownload_StreamsWithoutTheLock proves a download // lets go of the rename lock once its archive is open: while the // download is stalled writing, an edit can still rename the target. func TestHandleTargetDownload_StreamsWithoutTheLock(t *testing.T) { t.Parallel() env := setupSourceTest(t) wh := seedWebhookWithRetention(t, env.db, 7) archive := seedTarget(t, env.db, wh.ID, database.TargetTypeDatabase) req := httptest.NewRequestWithContext( t.Context(), http.MethodGet, downloadPath(wh.ID, archive.ID), nil, ) for _, c := range env.cookies { req.AddCookie(c) } sw := &stalledWriter{ ResponseRecorder: httptest.NewRecorder(), writing: make(chan struct{}), resume: make(chan struct{}), } downloaded := make(chan struct{}) go func() { targetRouter(env).ServeHTTP(sw, req) close(downloaded) }() <-sw.writing edited := make(chan *httptest.ResponseRecorder, 1) go func() { edited <- renameTarget(env, wh.ID, archive.ID) }() select { case w := <-edited: assert.Equal(t, http.StatusSeeOther, w.Code) case <-time.After(10 * time.Second): t.Error("the rename waited for the download") } close(sw.resume) <-downloaded assert.Equal(t, http.StatusOK, sw.Code) } // seedArchive writes rows to the archive file at path, each with a // body of bodySize random bytes, which do not compress. Its table has // only the columns the test fills; an export writes the others empty. func seedArchive(t *testing.T, path string, rows, bodySize int) { t.Helper() db, err := database.OpenSQLite(path, database.SQLiteModeCreate) require.NoError(t, err) defer func() { require.NoError(t, db.Close()) }() _, err = db.ExecContext(t.Context(), "CREATE TABLE archived_events (id INTEGER PRIMARY KEY, body TEXT)", ) require.NoError(t, err) body := make([]byte, bodySize) for range rows { _, _ = rand.Read(body) _, err = db.ExecContext(t.Context(), "INSERT INTO archived_events (body) VALUES (?)", string(body), ) require.NoError(t, err) } } // TestHandleTargetDownload_OutlastsTheRequestLimit proves a download // runs for as long as the client keeps reading. Behind a request limit // and a server write timeout of a tenth of a second, the client stops // reading once the response has started, waits three times as long, // and still gets the whole file. The archive is larger than a // connection holds, so the download is still being written while the // client waits. func TestHandleTargetDownload_OutlastsTheRequestLimit(t *testing.T) { t.Parallel() const ( limit = 100 * time.Millisecond rows = 12 bodySize = 1 << 20 ) env := setupSourceTest(t) wh := seedWebhookWithRetention(t, env.db, 7) archive := seedTarget(t, env.db, wh.ID, database.TargetTypeDatabase) seedArchive( t, delivery.ArchivePath(env.dbMgr, &wh, archive), rows, bodySize, ) srv := httptest.NewUnstartedServer( chimw.Timeout(limit)(targetRouter(env)), ) srv.Config.WriteTimeout = limit srv.Start() t.Cleanup(srv.Close) req, err := http.NewRequestWithContext( t.Context(), http.MethodGet, srv.URL+downloadPath(wh.ID, archive.ID), nil, ) require.NoError(t, err) for _, c := range env.cookies { req.AddCookie(c) } resp, err := srv.Client().Do(req) require.NoError(t, err) defer func() { _ = resp.Body.Close() }() require.Equal(t, http.StatusOK, resp.StatusCode) time.Sleep(3 * limit) zr, err := gzip.NewReader(resp.Body) require.NoError(t, err) var ( got map[string]json.RawMessage events []json.RawMessage ) require.NoError(t, json.NewDecoder(zr).Decode(&got)) require.NoError(t, json.Unmarshal(got["archived_events"], &events)) assert.Len(t, events, rows) // Reading to the end makes the gzip reader check that the file was // finished. _, err = io.ReadAll(zr) require.NoError(t, err) } // brokenWriter is a response writer whose writes fail once the // response has started, as they do when the client goes away. type brokenWriter struct { *httptest.ResponseRecorder } func (b brokenWriter) Write(p []byte) (int, error) { if b.Body.Len() > 0 { return 0, errClientGone } return b.ResponseRecorder.Write(p) } // TestHandleTargetDownload_AbortsWhenItFails proves a download that // fails after its response has started aborts the connection, so the // client sees a failed download rather than a file that looks // complete and does not decompress. func TestHandleTargetDownload_AbortsWhenItFails(t *testing.T) { t.Parallel() env := setupSourceTest(t) wh := seedWebhookWithRetention(t, env.db, 7) archive := seedTarget(t, env.db, wh.ID, database.TargetTypeDatabase) req := httptest.NewRequestWithContext( t.Context(), http.MethodGet, downloadPath(wh.ID, archive.ID), nil, ) for _, c := range env.cookies { req.AddCookie(c) } w := brokenWriter{ResponseRecorder: httptest.NewRecorder()} assert.PanicsWithValue(t, http.ErrAbortHandler, func() { targetRouter(env).ServeHTTP(w, req) }) assert.Equal(t, http.StatusOK, w.Code) }