package proxy import ( "context" "errors" "net/http" "net/http/httptrace" "net/http/httputil" "net/netip" "os" "sync" "sync/atomic" "time" "sneak.berlin/go/smallwebwaf/internal/ratelimit" "sneak.berlin/go/smallwebwaf/internal/requestlog" ) // flushAfterEachWrite has ReverseProxy pass on each part of the app's // answer as soon as it arrives. const flushAfterEachWrite time.Duration = -1 // refusal is smallwebwaf refusing a request, or refusing to go on with it: // the status the client is answered if the response has not started yet, // 0 to close the connection without an answer, and the action the log // line names. type refusal struct { status int action string } // request is one request on its way through smallwebwaf, from the moment // its headers have been read to its log line. type request struct { h *handler in *http.Request // rc sets the deadlines of the connection to the client. rc *http.ResponseController out *responseWriter body *requestBody // nil for a request without a body line requestlog.Line client netip.Addr peer netip.Addr peerTrusted bool start time.Time // upstreamStart is when the request was handed to the app. upstreamStart time.Time // cancel ends the request to the app. cancel context.CancelFunc // refused is the first refusal, from whichever goroutine meets it. refused atomic.Pointer[refusal] // complete is true once the app's whole answer has been passed on. complete bool // mu guards what follows. The timeouts run on goroutines of their // own, and the transport starts and stops them from its own; once // timersStopped is set, none of them acts any more. mu sync.Mutex timersStopped bool clientRequestTimer *time.Timer upstreamRequestTimer *time.Timer upstreamResponseTimer *time.Timer // requestSent is when the app had been sent the whole request. requestSent time.Time } // newRequest starts handling r: it notes the time and works out the // client. func (h *handler) newRequest(w http.ResponseWriter, r *http.Request) *request { start := time.Now() peer := peerAddress(r) trusted := h.config.TrustedProxies client := clientAddress(peer, r.Header.Values("X-Forwarded-For"), trusted) rq := &request{ h: h, in: r, rc: http.NewResponseController(w), out: &responseWriter{ResponseWriter: w}, client: client, peer: peer, peerTrusted: isInside(peer, trusted), start: start, line: requestlog.Line{ Time: requestlog.FormatTime(start), ClientIP: client.String(), PeerIP: peer.String(), Method: r.Method, Host: r.Host, Path: r.URL.EscapedPath(), Query: r.URL.RawQuery, Protocol: r.Proto, Referer: r.Referer(), UserAgent: r.UserAgent(), Action: requestlog.ActionForward, }, } if r.Body != http.NoBody { rq.body = &requestBody{body: limitBody(r.Body, h.config.RequestMaxBytes), rq: rq} } return rq } // check is the one place where a request can be refused once its client // is known, before its body is read or anything reaches the app. It // returns nil to let the request through. A client in SWWAF_ALLOW_NETS // skips every check but the size limit. For any other client, // SWWAF_DENY_NETS comes first, then a ban on its netblock, so that a // client either refuses is not looked up, and then the country lists; a // request any of them refuses is not counted for the rate limits. Then // come the rate limits, unless the client is in // SWWAF_RATE_LIMIT_EXEMPT_NETS, so that every other request is counted, // one refused for its size too. Every refusal but the size limit's is // answered with SWWAF_BAN_RESPONSE. ctx is the request's own context. func (rq *request) check(ctx context.Context) *refusal { cfg := rq.h.config allowed := isInside(rq.client, cfg.AllowNets) exempt := isInside(rq.client, cfg.RateLimitExemptNets) now := rq.h.now() if !allowed && isInside(rq.client, cfg.DenyNets) { return rq.banResponse(requestlog.ActionDenied) } if !allowed && rq.banned(now) { return rq.banResponse(requestlog.ActionBanned) } if !allowed && rq.countryDenied(ctx) { return rq.banResponse(requestlog.ActionCountryDenied) } if !allowed && !exempt && rq.limitBroken(now) { return rq.banResponse(requestlog.ActionRateLimited) } maxBytes := cfg.RequestMaxBytes if maxBytes > 0 && rq.in.ContentLength > maxBytes { return &refusal{ status: http.StatusRequestEntityTooLarge, action: requestlog.ActionTooLarge, } } return nil } // forward passes the request to the app and the app's answer back. ctx // is the request's own context. func (rq *request) forward(ctx context.Context) { ctx, cancel := context.WithCancel(ctx) defer cancel() rq.cancel = cancel ctx = httptrace.WithClientTrace(ctx, &httptrace.ClientTrace{ WroteRequest: rq.wroteRequest, }) out := rq.in.WithContext(ctx) if rq.body != nil { out.Body = rq.body } reverseProxy := &httputil.ReverseProxy{ Rewrite: rq.rewrite, Transport: rq.h.transport, FlushInterval: flushAfterEachWrite, ErrorLog: rq.h.errorLog, ModifyResponse: rq.modifyResponse, ErrorHandler: rq.answerError, } rq.startRequestTimers() rq.upstreamStart = time.Now() reverseProxy.ServeHTTP(rq.out, out) } // rewrite makes the request the app receives: the client's request, // unchanged, sent to SWWAF_UPSTREAM_URL, with the forwarded headers set. func (rq *request) rewrite(pr *httputil.ProxyRequest) { upstream := rq.h.config.UpstreamURL pr.Out.URL.Scheme = upstream.Scheme pr.Out.URL.Host = upstream.Host // ReverseProxy drops query parameters it cannot parse; the app gets // the query as the client sent it. pr.Out.URL.RawQuery = pr.In.URL.RawQuery setForwardedHeaders(pr.In, pr.Out, rq.peer, rq.peerTrusted) } // modifyResponse looks at the app's answer before ReverseProxy passes it // on. func (rq *request) modifyResponse(res *http.Response) error { rq.line.UpstreamStatus = res.StatusCode if res.StatusCode == http.StatusSwitchingProtocols { // An upgraded connection, such as a WebSocket, is not cut by the // timeouts. ReverseProxy writes this answer straight to the // connection it takes over, not through rq.out. rq.stopTimers() rq.out.status = res.StatusCode return nil } maxBytes := rq.h.config.ResponseMaxBytes if maxBytes > 0 && res.Body != http.NoBody && res.ContentLength > maxBytes { rq.refuse(refusal{status: http.StatusBadGateway, action: requestlog.ActionTooLarge}) return errResponseTooLarge } res.Body = &responseBody{body: limitBody(res.Body, maxBytes), rq: rq} rq.startClientResponseTimeout() return nil } // answerError is ReverseProxy's ErrorHandler: the request could not be // passed to the app, or the app's answer cannot be passed on. func (rq *request) answerError(_ http.ResponseWriter, _ *http.Request, err error) { refused := rq.refused.Load() if refused == nil { if rq.in.Context().Err() != nil { return // the client has gone, and there is no one to answer } rq.h.processLog.Warn("request to the app failed", "error", err.Error()) refused = &refusal{ status: http.StatusBadGateway, action: requestlog.ActionUpstreamError, } } rq.answer(*refused) } // answer sends smallwebwaf's own answer, unless the response has already // started, and records the refusal for the log line. func (rq *request) answer(r refusal) { rq.refused.CompareAndSwap(nil, &r) if rq.out.status != 0 { return // too late to answer: the connection can only be cut } if r.status == 0 { // SWWAF_BAN_RESPONSE is close. This panic has Go's server close // the connection without an answer, and log nothing; the log line // is still written as the handler returns. panic(http.ErrAbortHandler) } // A client found too slow is read no more; any other may go on // sending until its time is up, so that Go's server can read the // rest of the body and end the request cleanly. deadline := rq.clientRequestDeadline() if r.status == http.StatusRequestTimeout { deadline = time.Now() } rq.stopReadingBody(deadline) timeout := rq.h.config.ClientResponseTimeout if timeout > 0 { _ = rq.rc.SetWriteDeadline(time.Now().Add(timeout)) } http.Error(rq.out, http.StatusText(r.status), r.status) } // refuse records r, unless an earlier refusal was, and ends the request // to the app. func (rq *request) refuse(r refusal) { rq.refused.CompareAndSwap(nil, &r) rq.cancel() } // finish ends the request's timeouts and writes its log line. func (rq *request) finish() { rq.stopTimers() refused := rq.refused.Load() if refused == nil { rq.stopReadingBody(rq.clientRequestDeadline()) } line := &rq.line line.Status = rq.out.status line.ResponseBytes = rq.out.bytes if rq.body != nil { line.RequestBytes = rq.body.bytes.Load() } switch { case refused != nil: line.Action = refused.action case errors.Is(rq.out.err, os.ErrDeadlineExceeded): // The client took longer than SWWAF_CLIENT_RESPONSE_TIMEOUT to // take the response. line.Action = requestlog.ActionTimedOut case !rq.complete && (rq.out.err != nil || rq.in.Context().Err() != nil): line.Aborted = true } now := time.Now() line.DurationTotal = requestlog.Milliseconds(now.Sub(rq.start)) if !rq.upstreamStart.IsZero() { line.DurationUpstreamTotal = requestlog.Milliseconds(now.Sub(rq.upstreamStart)) } err := requestlog.Write(rq.h.requestLog, line) if err != nil { rq.h.processLog.Error("writing the request log failed", "error", err.Error()) } } // addToHistory adds the request, which has ended, to its client's // history. func (rq *request) addToHistory() { var requestBytes int64 if rq.body != nil { requestBytes = rq.body.bytes.Load() } rq.h.limiter.AddToHistory(clientGroup(rq.client), rq.h.now(), ratelimit.Request{ Country: rq.line.Country, Forwarded: !rq.upstreamStart.IsZero(), Status: rq.out.status, RequestBytes: requestBytes, ResponseBytes: rq.out.bytes, BrokeLimit: rq.line.Offence == requestlog.OffenceLimit, }) } // clientRequestDeadline is when the client must have sent its whole // request, or zero when SWWAF_CLIENT_REQUEST_TIMEOUT is off. func (rq *request) clientRequestDeadline() time.Time { timeout := rq.h.config.ClientRequestTimeout if timeout == 0 { return time.Time{} } return rq.start.Add(timeout) } // stopReadingBody ends, at deadline, the reading of a client body that has // not arrived whole: Go's server then reads no more of it, and closes the // connection after the answer. func (rq *request) stopReadingBody(deadline time.Time) { if rq.body == nil || rq.body.received.Load() { return } _ = rq.rc.SetReadDeadline(deadline) } // startRequestTimers starts the timeouts that run while the request goes // to the app: SWWAF_CLIENT_REQUEST_TIMEOUT until the client has sent its // whole body, and SWWAF_UPSTREAM_REQUEST_TIMEOUT until the app has been // sent the whole request. func (rq *request) startRequestTimers() { rq.mu.Lock() defer rq.mu.Unlock() if rq.body != nil && rq.h.config.ClientRequestTimeout > 0 { rq.clientRequestTimer = time.AfterFunc( time.Until(rq.clientRequestDeadline()), rq.requestTimedOut) } timeout := rq.h.config.UpstreamRequestTimeout if timeout > 0 { rq.upstreamRequestTimer = time.AfterFunc(timeout, rq.requestTimedOut) } } // requestTimedOut is called when a request timeout runs out while the // request is still on its way to the app. The answer names the side // smallwebwaf was waiting on at that moment: 408 when it was waiting for // the client to send more of its body, 504 when it was waiting for the // app to be reached or to take what it had. func (rq *request) requestTimedOut() { rq.mu.Lock() defer rq.mu.Unlock() if rq.timersStopped { return } if rq.body == nil || !rq.body.waiting.Load() { rq.refuse(refusal{ status: http.StatusGatewayTimeout, action: requestlog.ActionTimedOut, }) return } rq.refuse(refusal{ status: http.StatusRequestTimeout, action: requestlog.ActionTimedOut, }) // The transport gives up on the app only once its Read of the // client's body returns, so that Read is ended now. The lock keeps // this from reaching the connection after the request is handled. _ = rq.rc.SetReadDeadline(time.Now()) } // bodyReceived is called once the client has sent its whole body. func (rq *request) bodyReceived() { rq.mu.Lock() defer rq.mu.Unlock() stopTimer(rq.clientRequestTimer) } // wroteRequest is called once the app has been sent the whole request: // the request timeouts end and SWWAF_UPSTREAM_RESPONSE_TIMEOUT starts. func (rq *request) wroteRequest(info httptrace.WroteRequestInfo) { if info.Err != nil { return // the transport gives up, or tries again } rq.mu.Lock() defer rq.mu.Unlock() if rq.timersStopped { return } stopTimer(rq.clientRequestTimer) stopTimer(rq.upstreamRequestTimer) rq.requestSent = time.Now() timeout := rq.h.config.UpstreamResponseTimeout if timeout > 0 { rq.upstreamResponseTimer = time.AfterFunc(timeout, rq.responseTimedOut) } } // responseTimedOut is called when SWWAF_UPSTREAM_RESPONSE_TIMEOUT runs out // before the app has sent its whole answer. func (rq *request) responseTimedOut() { rq.mu.Lock() defer rq.mu.Unlock() if !rq.timersStopped { rq.refuse(refusal{ status: http.StatusGatewayTimeout, action: requestlog.ActionTimedOut, }) } } // responseReceived is called once the app has sent its whole answer. func (rq *request) responseReceived() { rq.complete = true rq.stopTimers() } // startClientResponseTimeout sets SWWAF_CLIENT_RESPONSE_TIMEOUT on the // connection to the client: the response must reach the client within it // of the end of the request, or of now if the app answers before it has // the whole request. func (rq *request) startClientResponseTimeout() { timeout := rq.h.config.ClientResponseTimeout if timeout == 0 { return } from := rq.sentAt() if from.IsZero() { from = time.Now() } _ = rq.rc.SetWriteDeadline(from.Add(timeout)) } // sentAt is when the app had been sent the whole request, or zero. func (rq *request) sentAt() time.Time { rq.mu.Lock() defer rq.mu.Unlock() return rq.requestSent } // stopTimers stops the request's timeouts and keeps any from starting // later: the app's answer is complete, the connection upgraded, or the // request handled. func (rq *request) stopTimers() { rq.mu.Lock() defer rq.mu.Unlock() rq.timersStopped = true stopTimer(rq.clientRequestTimer) stopTimer(rq.upstreamRequestTimer) stopTimer(rq.upstreamResponseTimer) } // stopTimer stops t, which is nil when its timeout is off. func stopTimer(t *time.Timer) { if t != nil { t.Stop() } }