288 строки
9.1 KiB
Go
288 строки
9.1 KiB
Go
package workercluster
|
|
|
|
import (
|
|
"bytes"
|
|
"crypto/subtle"
|
|
"encoding/base64"
|
|
"fmt"
|
|
"io"
|
|
"log"
|
|
"net"
|
|
"net/http"
|
|
"sync"
|
|
"time"
|
|
|
|
rafthttp "github.com/CanonicalLtd/raft-http"
|
|
)
|
|
|
|
// HTTPCreds is the basic-auth credential pair workers use to gate the
|
|
// rafthttp endpoint. The values come from distworker.HTTPConfig.Login
|
|
// and distworker.HTTPConfig.Password (env vars WORKER_LOGIN and
|
|
// WORKER_PASSWORD).
|
|
type HTTPCreds struct {
|
|
Login string
|
|
Password string
|
|
}
|
|
|
|
// IsConfigured returns true when both login and password are set.
|
|
// NewTransport and NewHandler refuse to start otherwise.
|
|
func (c HTTPCreds) IsConfigured() bool {
|
|
return c.Login != "" && c.Password != ""
|
|
}
|
|
|
|
// AuthDial returns a rafthttp.Dial function that wraps the inner
|
|
// connection so the Authorization header is injected on every HTTP
|
|
// request the rafthttp library writes over it.
|
|
//
|
|
// rafthttp builds its own http.Request without Authorization (GET for
|
|
// the stream upgrade, POST/DELETE for membership changes). It then
|
|
// writes the request to the net.Conn returned by Dial. We can't inject
|
|
// at the request layer; we have to do it at the connection layer.
|
|
//
|
|
// Implementation: the wrapper buffers the first Write, looks for the
|
|
// end-of-headers marker (\r\n\r\n), inserts an Authorization header
|
|
// just before it, then forwards the augmented buffer plus any further
|
|
// writes to the inner conn.
|
|
//
|
|
// TODO(security): replace with a rafthttp fork that supports an
|
|
// outbound Authorization header or use NewDialTLS with mTLS client
|
|
// certs once the worker identity model in plan section 7.1 lands.
|
|
func AuthDial(inner rafthttp.Dial, creds HTTPCreds) rafthttp.Dial {
|
|
if !creds.IsConfigured() {
|
|
panic("workercluster: AuthDial requires both login and password")
|
|
}
|
|
return func(addr string, timeout time.Duration) (net.Conn, error) {
|
|
conn, err := inner(addr, timeout)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &authInjectingConn{
|
|
Conn: conn,
|
|
auth: "Basic " + base64.StdEncoding.EncodeToString([]byte(creds.Login+":"+creds.Password)),
|
|
}, nil
|
|
}
|
|
}
|
|
|
|
// authInjectingConn wraps a net.Conn and rewrites the first HTTP
|
|
// request written to it so it carries an Authorization header.
|
|
//
|
|
// State machine:
|
|
//
|
|
// - injected = false: incoming bytes are appended to buf until we see
|
|
// the end-of-headers marker (\r\n\r\n).
|
|
// - once we see \r\n\r\n, we insert the Authorization header just
|
|
// before it, drain the buffer to the inner conn, and switch to
|
|
// passthrough.
|
|
// - if too much data arrives without a header terminator (e.g. a
|
|
// very large POST body), we forward as-is; the auth handler will
|
|
// reject the request.
|
|
// - subsequent writes pass through unchanged.
|
|
//
|
|
// This is sufficient for rafthttp: every HTTP request it writes is a
|
|
// self-contained, single-shot request over a fresh connection.
|
|
type authInjectingConn struct {
|
|
net.Conn
|
|
|
|
auth string
|
|
|
|
mu sync.Mutex
|
|
buf []byte
|
|
injected bool
|
|
}
|
|
|
|
func (a *authInjectingConn) Write(p []byte) (int, error) {
|
|
a.mu.Lock()
|
|
defer a.mu.Unlock()
|
|
|
|
if !a.injected {
|
|
a.buf = append(a.buf, p...)
|
|
if len(a.buf) > maxAuthHeaderBuffer {
|
|
// Too much data before we saw the header terminator;
|
|
// bail out and forward as-is. The request will be
|
|
// rejected by the auth handler on the other side,
|
|
// which is the correct failure mode.
|
|
a.injected = true
|
|
if _, err := a.Conn.Write(a.buf); err != nil {
|
|
return 0, err
|
|
}
|
|
a.buf = nil
|
|
return len(p), nil
|
|
}
|
|
if idx := bytes.Index(a.buf, []byte("\r\n\r\n")); idx >= 0 {
|
|
// Split around the header terminator.
|
|
head := a.buf[:idx]
|
|
rest := a.buf[idx:]
|
|
newBuf := make([]byte, 0, len(a.buf)+len(a.auth)+32)
|
|
newBuf = append(newBuf, head...)
|
|
newBuf = append(newBuf, []byte("\r\nAuthorization: ")...)
|
|
newBuf = append(newBuf, []byte(a.auth)...)
|
|
newBuf = append(newBuf, rest...)
|
|
a.buf = newBuf
|
|
a.injected = true
|
|
n, err := a.Conn.Write(a.buf)
|
|
a.buf = nil
|
|
return n, err
|
|
}
|
|
return len(p), nil
|
|
}
|
|
return a.Conn.Write(p)
|
|
}
|
|
|
|
// maxAuthHeaderBuffer caps the bytes we'll buffer waiting for the
|
|
// header terminator. 64 KiB is well past any reasonable rafthttp
|
|
// request and large enough to absorb the headers + a small body.
|
|
const maxAuthHeaderBuffer = 64 * 1024
|
|
|
|
// NewAuthHandler wraps an inner rafthttp.Handler with an HTTP basic-auth
|
|
// check. Requests without matching credentials are rejected with 401
|
|
// before the rafthttp path runs.
|
|
//
|
|
// The auth check uses crypto/subtle.ConstantTimeCompare to avoid timing
|
|
// leaks on the credential comparison.
|
|
func NewAuthHandler(inner *rafthttp.Handler, creds HTTPCreds, logger *log.Logger) http.Handler {
|
|
if logger == nil {
|
|
logger = log.Default()
|
|
}
|
|
if !creds.IsConfigured() {
|
|
// We panic on construction rather than at request time so a
|
|
// misconfigured worker fails fast at startup.
|
|
panic("workercluster: NewAuthHandler requires both login and password")
|
|
}
|
|
|
|
expectedUser := []byte(creds.Login)
|
|
expectedPass := []byte(creds.Password)
|
|
|
|
return &authHandler{inner: inner, expectedUser: expectedUser, expectedPass: expectedPass, logger: logger}
|
|
}
|
|
|
|
type authHandler struct {
|
|
inner *rafthttp.Handler
|
|
expectedUser []byte
|
|
expectedPass []byte
|
|
logger *log.Logger
|
|
}
|
|
|
|
func (a *authHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
user, pass, ok := r.BasicAuth()
|
|
if !ok ||
|
|
subtle.ConstantTimeCompare([]byte(user), a.expectedUser) != 1 ||
|
|
subtle.ConstantTimeCompare([]byte(pass), a.expectedPass) != 1 {
|
|
a.logger.Printf("[WARN] raft-http: rejected %s %s from %s: bad credentials",
|
|
r.Method, r.URL.Path, r.RemoteAddr)
|
|
w.Header().Set("WWW-Authenticate", `Basic realm="raft"`)
|
|
http.Error(w, "unauthorized", http.StatusUnauthorized)
|
|
return
|
|
}
|
|
a.inner.ServeHTTP(w, r)
|
|
}
|
|
|
|
// Unwrap exposes the inner rafthttp.Handler so callers like
|
|
// unwrapRafthttpHandler can find it.
|
|
func (a *authHandler) Unwrap() http.Handler { return a.inner }
|
|
|
|
// NewTransport builds a rafthttp Layer + NetworkTransport pair bound to
|
|
// the given listener. The returned Layer is ready to hand to
|
|
// raft.NewNetworkTransport. Close must be called on shutdown to drain
|
|
// the Layer's HTTP handler.
|
|
//
|
|
// dial is the rafthttp.Dial used to connect to peers; if nil, the
|
|
// rafthttp.NewDialTCP default is used. handler is an http.Handler that
|
|
// owns the rafthttp endpoint; production callers pass the auth wrapper
|
|
// from NewAuthHandler. The inbound listener is started by the caller
|
|
// because the listener needs to be running before peers can dial in.
|
|
func NewTransport(
|
|
raftPath string,
|
|
listener net.Listener,
|
|
handler http.Handler,
|
|
dial rafthttp.Dial,
|
|
logOutput io.Writer,
|
|
) (*rafthttp.Layer, *http.Server, error) {
|
|
if raftPath == "" {
|
|
raftPath = "/raft"
|
|
}
|
|
if listener == nil {
|
|
return nil, nil, fmt.Errorf("workercluster: listener is required")
|
|
}
|
|
if handler == nil {
|
|
return nil, nil, fmt.Errorf("workercluster: handler is required")
|
|
}
|
|
if dial == nil {
|
|
dial = rafthttp.NewDialTCP()
|
|
}
|
|
logger := log.New(logOutput, "[raft-http] ", log.LstdFlags)
|
|
|
|
realHandler, ok := unwrapRafthttpHandler(handler)
|
|
if !ok {
|
|
return nil, nil, fmt.Errorf("workercluster: handler must wrap a *rafthttp.Handler")
|
|
}
|
|
layer := rafthttp.NewLayerWithLogger(raftPath, listener.Addr(), realHandler, dial, logger)
|
|
|
|
server := &http.Server{
|
|
Handler: handler,
|
|
ReadHeaderTimeout: 10 * time.Second,
|
|
}
|
|
|
|
return layer, server, nil
|
|
}
|
|
|
|
// unwrapRafthttpHandler walks a chain of http.Handler wrappers and
|
|
// returns the *rafthttp.Handler at the bottom. Production wraps it in
|
|
// NewAuthHandler; tests can wrap it in additional middleware.
|
|
func unwrapRafthttpHandler(h http.Handler) (*rafthttp.Handler, bool) {
|
|
for {
|
|
switch v := h.(type) {
|
|
case *rafthttp.Handler:
|
|
return v, true
|
|
case interface{ Unwrap() http.Handler }:
|
|
h = v.Unwrap()
|
|
default:
|
|
return nil, false
|
|
}
|
|
}
|
|
}
|
|
|
|
// CheckBasicAuth is a small helper used by transport_test.go to confirm
|
|
// the auth wrapper rejects bad creds and accepts good ones without
|
|
// needing the rafthttp Library state.
|
|
func CheckBasicAuth(h http.Handler, r *http.Request, login, password string) bool {
|
|
r.Header.Set("Authorization", basicAuthHeader(login, password))
|
|
rec := &recordingResponseWriter{header: http.Header{}}
|
|
h.ServeHTTP(rec, r)
|
|
return rec.status == http.StatusOK
|
|
}
|
|
|
|
// basicAuthHeader returns the value of an HTTP Basic Authorization
|
|
// header for the given user/password pair. Exported for tests; production
|
|
// code uses Go's r.BasicAuth() helper.
|
|
func basicAuthHeader(user, pass string) string {
|
|
const prefix = "Basic "
|
|
value := user + ":" + pass
|
|
return prefix + base64Encode(value)
|
|
}
|
|
|
|
// base64Encode is a tiny indirection so tests do not import encoding/base64
|
|
// directly; the production call sites use Go's standard library.
|
|
func base64Encode(s string) string {
|
|
return base64.StdEncoding.EncodeToString([]byte(s))
|
|
}
|
|
|
|
// recordingResponseWriter is a minimal http.ResponseWriter for tests.
|
|
type recordingResponseWriter struct {
|
|
header http.Header
|
|
body []byte
|
|
status int
|
|
}
|
|
|
|
func (w *recordingResponseWriter) Header() http.Header { return w.header }
|
|
func (w *recordingResponseWriter) Write(b []byte) (int, error) {
|
|
w.body = append(w.body, b...)
|
|
if w.status == 0 {
|
|
w.status = http.StatusOK
|
|
}
|
|
return len(b), nil
|
|
}
|
|
|
|
func (w *recordingResponseWriter) WriteHeader(status int) {
|
|
w.status = status
|
|
}
|