240 lines
7.8 KiB
Go
240 lines
7.8 KiB
Go
package application
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
"net/http/httptest"
|
|
"strings"
|
|
"sync"
|
|
"testing"
|
|
"time"
|
|
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/jobs"
|
|
"git.nianxx.cn/wangxuming/NianAIGC/backend/internal/logging"
|
|
)
|
|
|
|
type recordingEventLogger struct {
|
|
mu sync.Mutex
|
|
inputs []logging.Input
|
|
}
|
|
|
|
type panickingEventLogger struct{}
|
|
|
|
func (panickingEventLogger) Append(context.Context, logging.Input) (logging.Entry, error) {
|
|
panic("log sink unavailable")
|
|
}
|
|
|
|
func (logger *recordingEventLogger) Append(_ context.Context, input logging.Input) (logging.Entry, error) {
|
|
logger.mu.Lock()
|
|
defer logger.mu.Unlock()
|
|
logger.inputs = append(logger.inputs, input)
|
|
return logging.Entry{}, nil
|
|
}
|
|
|
|
func (logger *recordingEventLogger) snapshot() []logging.Input {
|
|
logger.mu.Lock()
|
|
defer logger.mu.Unlock()
|
|
return append([]logging.Input(nil), logger.inputs...)
|
|
}
|
|
|
|
func TestHTTPEventLoggingRecordsServerErrorsButNotClientErrors(t *testing.T) {
|
|
logger := &recordingEventLogger{}
|
|
handler := WithHTTPEventLogging(http.HandlerFunc(func(w http.ResponseWriter, request *http.Request) {
|
|
if request.URL.Path == "/client-error" {
|
|
w.WriteHeader(http.StatusBadRequest)
|
|
return
|
|
}
|
|
w.WriteHeader(http.StatusBadGateway)
|
|
}), logger)
|
|
|
|
for _, target := range []string{"/client-error?secret=client", "/server-error?secret=server"} {
|
|
recorder := httptest.NewRecorder()
|
|
request := httptest.NewRequest(http.MethodPost, target, nil)
|
|
request.Header.Set("Cookie", "session=top-secret")
|
|
handler.ServeHTTP(recorder, request)
|
|
}
|
|
|
|
inputs := logger.snapshot()
|
|
if len(inputs) != 1 {
|
|
t.Fatalf("event count = %d, want 1: %#v", len(inputs), inputs)
|
|
}
|
|
entry := inputs[0]
|
|
if entry.Level != logging.Error || entry.Source != "http" || entry.Status != http.StatusBadGateway || entry.Method != http.MethodPost || entry.Path != "/server-error" {
|
|
t.Fatalf("event = %#v", entry)
|
|
}
|
|
assertEventContainsNoSensitiveData(t, entry)
|
|
}
|
|
|
|
func TestHTTPEventLoggingSinkFailureDoesNotChangeResponse(t *testing.T) {
|
|
handler := WithHTTPEventLogging(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
w.WriteHeader(http.StatusBadGateway)
|
|
_, _ = w.Write([]byte("upstream unavailable"))
|
|
}), panickingEventLogger{})
|
|
response := httptest.NewRecorder()
|
|
|
|
handler.ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/failure", nil))
|
|
|
|
if response.Code != http.StatusBadGateway || response.Body.String() != "upstream unavailable" {
|
|
t.Fatalf("response = %d %q", response.Code, response.Body.String())
|
|
}
|
|
}
|
|
|
|
func TestHTTPEventLoggingRecoversPanicBeforeCommitWithGenericResponseAndEvent(t *testing.T) {
|
|
logger := &recordingEventLogger{}
|
|
handler := WithHTTPEventLogging(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {
|
|
panic("password=hunter2")
|
|
}), logger)
|
|
recorder := httptest.NewRecorder()
|
|
request := httptest.NewRequest(http.MethodPut, "/panic?token=query-secret", strings.NewReader("body-secret"))
|
|
request.Header.Set("Cookie", "session=cookie-secret")
|
|
|
|
handler.ServeHTTP(recorder, request)
|
|
|
|
if recorder.Code != http.StatusInternalServerError {
|
|
t.Fatalf("status = %d, want 500", recorder.Code)
|
|
}
|
|
if recorder.Body.String() != "Internal Server Error\n" {
|
|
t.Fatalf("body = %q", recorder.Body.String())
|
|
}
|
|
inputs := logger.snapshot()
|
|
if len(inputs) != 1 || inputs[0].Status != http.StatusInternalServerError || inputs[0].Message != "HTTP handler panic" {
|
|
t.Fatalf("events = %#v", inputs)
|
|
}
|
|
assertEventContainsNoSensitiveData(t, inputs[0])
|
|
}
|
|
|
|
func TestHTTPEventLoggingStreamsResponseBeforeHandlerCompletes(t *testing.T) {
|
|
logger := &recordingEventLogger{}
|
|
release := make(chan struct{})
|
|
handler := WithHTTPEventLogging(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
_, _ = w.Write([]byte("first-chunk"))
|
|
<-release
|
|
_, _ = w.Write([]byte("second-chunk"))
|
|
}), logger)
|
|
recorder := &signalingResponseRecorder{
|
|
ResponseRecorder: httptest.NewRecorder(),
|
|
writeObserved: make(chan struct{}),
|
|
}
|
|
done := make(chan struct{})
|
|
go func() {
|
|
defer close(done)
|
|
handler.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/stream", nil))
|
|
}()
|
|
defer func() {
|
|
close(release)
|
|
<-done
|
|
}()
|
|
|
|
select {
|
|
case <-recorder.writeObserved:
|
|
if got := recorder.Body.String(); got != "first-chunk" {
|
|
t.Fatalf("body before handler completion = %q", got)
|
|
}
|
|
case <-time.After(500 * time.Millisecond):
|
|
t.Fatal("first response chunk was buffered until handler completion")
|
|
}
|
|
}
|
|
|
|
func TestHTTPEventLoggingCannotRewriteAlreadyCommittedResponseAfterPanic(t *testing.T) {
|
|
logger := &recordingEventLogger{}
|
|
handler := WithHTTPEventLogging(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
|
_, _ = w.Write([]byte("already-committed"))
|
|
panic("password=hunter2")
|
|
}), logger)
|
|
recorder := httptest.NewRecorder()
|
|
|
|
func() {
|
|
defer func() {
|
|
if recovered := recover(); recovered != http.ErrAbortHandler {
|
|
t.Fatalf("panic = %#v, want http.ErrAbortHandler", recovered)
|
|
}
|
|
}()
|
|
handler.ServeHTTP(recorder, httptest.NewRequest(http.MethodGet, "/stream-panic", nil))
|
|
}()
|
|
|
|
if recorder.Code != http.StatusOK || recorder.Body.String() != "already-committed" {
|
|
t.Fatalf("committed response = %d %q", recorder.Code, recorder.Body.String())
|
|
}
|
|
inputs := logger.snapshot()
|
|
if len(inputs) != 1 || inputs[0].Status != http.StatusInternalServerError || inputs[0].Message != "HTTP handler panic" {
|
|
t.Fatalf("events = %#v", inputs)
|
|
}
|
|
}
|
|
|
|
type signalingResponseRecorder struct {
|
|
*httptest.ResponseRecorder
|
|
writeObserved chan struct{}
|
|
once sync.Once
|
|
}
|
|
|
|
func (recorder *signalingResponseRecorder) Write(body []byte) (int, error) {
|
|
written, err := recorder.ResponseRecorder.Write(body)
|
|
recorder.once.Do(func() { close(recorder.writeObserved) })
|
|
return written, err
|
|
}
|
|
|
|
type stubTickRunner struct {
|
|
result jobs.TickResult
|
|
err error
|
|
}
|
|
|
|
func (runner stubTickRunner) Tick(context.Context, string) (jobs.TickResult, error) {
|
|
return runner.result, runner.err
|
|
}
|
|
|
|
func TestEventLoggingTickRunnerRecordsTickAndItemErrors(t *testing.T) {
|
|
t.Run("tick error", func(t *testing.T) {
|
|
logger := &recordingEventLogger{}
|
|
runner := WithTickEventLogging(stubTickRunner{err: errors.New("secret=provider-key")}, logger)
|
|
|
|
_, err := runner.Tick(context.Background(), "worker-secret")
|
|
if err == nil {
|
|
t.Fatal("Tick error = nil")
|
|
}
|
|
inputs := logger.snapshot()
|
|
if len(inputs) != 1 || inputs[0].Source != "worker" || inputs[0].Level != logging.Error || inputs[0].Message != "Worker tick failed" {
|
|
t.Fatalf("events = %#v", inputs)
|
|
}
|
|
assertEventContainsNoSensitiveData(t, inputs[0])
|
|
})
|
|
|
|
t.Run("item errors", func(t *testing.T) {
|
|
logger := &recordingEventLogger{}
|
|
result := jobs.TickResult{Jobs: []jobs.TickJob{
|
|
{ID: "job-safe", Error: "password=hunter2"},
|
|
{ID: "job-ok"},
|
|
{ID: "job-safe-2", Error: "token=provider-key"},
|
|
}}
|
|
runner := WithTickEventLogging(stubTickRunner{result: result}, logger)
|
|
|
|
got, err := runner.Tick(context.Background(), "embedded-worker")
|
|
if err != nil || len(got.Jobs) != 3 {
|
|
t.Fatalf("Tick() = %#v, %v", got, err)
|
|
}
|
|
inputs := logger.snapshot()
|
|
if len(inputs) != 2 {
|
|
t.Fatalf("event count = %d, want 2: %#v", len(inputs), inputs)
|
|
}
|
|
for _, input := range inputs {
|
|
if input.Source != "worker" || input.Level != logging.Error || input.Message != "Worker job failed" {
|
|
t.Fatalf("event = %#v", input)
|
|
}
|
|
assertEventContainsNoSensitiveData(t, input)
|
|
}
|
|
})
|
|
}
|
|
|
|
func assertEventContainsNoSensitiveData(t *testing.T, input logging.Input) {
|
|
t.Helper()
|
|
text := strings.ToLower(input.Message + input.Method + input.Path + input.Stack)
|
|
for _, sensitive := range []string{"hunter2", "provider-key", "query-secret", "cookie-secret", "body-secret", "top-secret"} {
|
|
if strings.Contains(text, sensitive) {
|
|
t.Fatalf("event contains sensitive data %q: %#v", sensitive, input)
|
|
}
|
|
}
|
|
if input.Error != nil || input.Details != nil {
|
|
t.Fatalf("event carries unsafe error/details: %#v", input)
|
|
}
|
|
}
|