package proxy import ( "errors" "io" "net/http" "sync/atomic" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) // errResponseTooLarge ends an app's response body that is longer than // SWWAF_RESPONSE_MAX_BYTES. var errResponseTooLarge = errors.New( "the response body is over SWWAF_RESPONSE_MAX_BYTES") // requestBody is the client's request body on its way to the app. The // transport reads it on a goroutine of its own. type requestBody struct { // body is the client's body, ending in an *http.MaxBytesError past // SWWAF_REQUEST_MAX_BYTES. body io.ReadCloser rq *request // waiting is true while a Read waits for the client to send more. waiting atomic.Bool // received is true once the client has sent the whole body. received atomic.Bool // bytes is how much of the body has been read. bytes atomic.Int64 } // Read reads from the client's body. func (b *requestBody) Read(p []byte) (int, error) { b.waiting.Store(true) n, err := b.body.Read(p) b.waiting.Store(false) b.bytes.Add(int64(n)) var tooLarge *http.MaxBytesError switch { case errors.Is(err, io.EOF): b.received.Store(true) b.rq.bodyReceived() case errors.As(err, &tooLarge): b.rq.refuse(refusal{ status: http.StatusRequestEntityTooLarge, action: requestlog.ActionTooLarge, }) } return n, err } // Close closes the client's body. func (b *requestBody) Close() error { return b.body.Close() } // responseBody is the app's response body on its way to the client. type responseBody struct { // body is the app's body, ending in an *http.MaxBytesError past // SWWAF_RESPONSE_MAX_BYTES. body io.ReadCloser rq *request } // Read reads from the app's body. func (b *responseBody) Read(p []byte) (int, error) { n, err := b.body.Read(p) if err == nil { return n, nil } var tooLarge *http.MaxBytesError switch { case errors.Is(err, io.EOF): b.rq.responseReceived() case errors.As(err, &tooLarge): b.rq.refuse(refusal{ status: http.StatusBadGateway, action: requestlog.ActionTooLarge, }) return n, errResponseTooLarge case b.rq.in.Context().Err() == nil: // The answer broke off, not because the client went away. If a // timeout cut it, that refusal came first and is the one kept. b.rq.refuse(refusal{ status: http.StatusBadGateway, action: requestlog.ActionUpstreamError, }) } return n, err } // Close closes the app's body. func (b *responseBody) Close() error { return b.body.Close() } // limitBody returns body, cut off with an *http.MaxBytesError after // maxBytes, or unchanged if maxBytes is zero, which is off. func limitBody(body io.ReadCloser, maxBytes int64) io.ReadCloser { if maxBytes == 0 { return body } // Without a ResponseWriter, MaxBytesReader only counts and cuts off. return http.MaxBytesReader(nil, body, maxBytes) } // responseWriter is the response to the client. It notes the status and // size for the log line, and the first error writing to the client. type responseWriter struct { http.ResponseWriter // status is the final status sent, or zero before one is. status int bytes int64 err error } // WriteHeader sends the status and headers. An informational 1xx status // is passed on and the final status still comes later. func (w *responseWriter) WriteHeader(status int) { if status >= http.StatusOK && w.status == 0 { w.status = status } w.ResponseWriter.WriteHeader(status) } // Write sends part of the body. func (w *responseWriter) Write(p []byte) (int, error) { if w.status == 0 { w.status = http.StatusOK } n, err := w.ResponseWriter.Write(p) w.bytes += int64(n) w.noteError(err) return n, err } // FlushError sends what has been written so far. // http.ResponseController calls it, as ReverseProxy does after each // write. func (w *responseWriter) FlushError() error { err := http.NewResponseController(w.ResponseWriter).Flush() w.noteError(err) return err } // Unwrap lets http.ResponseController reach net/http's own // ResponseWriter, which is how ReverseProxy takes over the connection of // an upgraded request. func (w *responseWriter) Unwrap() http.ResponseWriter { return w.ResponseWriter } // noteError keeps the first error writing to the client. func (w *responseWriter) noteError(err error) { if w.err == nil { w.err = err } }