Files

324 lines
9.7 KiB
Go

package forwarder
import (
"context"
"crypto/hmac"
"crypto/sha256"
"crypto/subtle"
"crypto/tls"
"crypto/x509"
"encoding/hex"
"errors"
"fmt"
"io"
"net/http"
"net/url"
"os"
"strconv"
"strings"
"time"
"github.com/gorilla/websocket"
)
const (
headerVersion = "X-Neko-Relay-Forwarded"
headerSource = "X-Neko-Relay-Source"
headerTarget = "X-Neko-Relay-Target"
headerTimestamp = "X-Neko-Relay-Timestamp"
headerSignature = "X-Neko-Relay-Signature"
maxHTTPResponse = 16 << 20
maxWSMessage = 4 << 20
)
type Config struct {
NodeID string
Endpoints map[string]string
Secret []byte
ClientCertFile string
ClientKeyFile string
CAFile string
HTTPClient *http.Client
WebSocketDialer *websocket.Dialer
Now func() time.Time
}
type Forwarder struct {
nodeID string
endpoints map[string]*url.URL
secret []byte
http *http.Client
websocket *websocket.Dialer
now func() time.Time
}
func exactOrigin(raw string) (*url.URL, error) {
parsed, err := url.Parse(strings.TrimSpace(raw))
if err != nil || parsed.Host == "" || parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
return nil, errors.New("invalid internal Relay origin")
}
loopback := parsed.Hostname() == "localhost" || parsed.Hostname() == "127.0.0.1" || parsed.Hostname() == "::1"
if parsed.Scheme != "https" && !(parsed.Scheme == "http" && loopback) {
return nil, errors.New("internal Relay origin must use HTTPS")
}
if parsed.Path != "" && parsed.Path != "/" {
return nil, errors.New("internal Relay origin must not contain a path")
}
parsed.Path = ""
return parsed, nil
}
func tlsConfig(certFile, keyFile, caFile string) (*tls.Config, error) {
certificate, err := tls.LoadX509KeyPair(certFile, keyFile)
if err != nil {
return nil, fmt.Errorf("load internal Relay mTLS identity: %w", err)
}
caPEM, err := os.ReadFile(caFile)
if err != nil {
return nil, fmt.Errorf("load internal Relay CA: %w", err)
}
roots := x509.NewCertPool()
if !roots.AppendCertsFromPEM(caPEM) {
return nil, errors.New("internal Relay CA contains no certificates")
}
return &tls.Config{
MinVersion: tls.VersionTLS13, RootCAs: roots,
Certificates: []tls.Certificate{certificate},
}, nil
}
func New(config Config) (*Forwarder, error) {
if strings.TrimSpace(config.NodeID) == "" || len(config.Secret) < 32 {
return nil, errors.New("forwarder requires node identity and a 32-byte secret")
}
endpoints := make(map[string]*url.URL, len(config.Endpoints))
for reference, raw := range config.Endpoints {
if reference == "" {
return nil, errors.New("internal endpoint reference is empty")
}
origin, err := exactOrigin(raw)
if err != nil {
return nil, fmt.Errorf("internal endpoint %q: %w", reference, err)
}
endpoints[reference] = origin
}
client := config.HTTPClient
dialer := config.WebSocketDialer
if client == nil || dialer == nil {
tlsConfiguration, err := tlsConfig(config.ClientCertFile, config.ClientKeyFile, config.CAFile)
if err != nil {
return nil, err
}
if client == nil {
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.Proxy = nil
transport.TLSClientConfig = tlsConfiguration.Clone()
client = &http.Client{Transport: transport, Timeout: 75 * time.Second}
}
if dialer == nil {
dialer = &websocket.Dialer{
TLSClientConfig: tlsConfiguration.Clone(), HandshakeTimeout: 10 * time.Second,
ReadBufferSize: 4096, WriteBufferSize: 4096, EnableCompression: false,
}
}
}
clone := *client
clone.CheckRedirect = func(_ *http.Request, _ []*http.Request) error { return http.ErrUseLastResponse }
if config.Now == nil {
config.Now = time.Now
}
return &Forwarder{
nodeID: strings.TrimSpace(config.NodeID), endpoints: endpoints,
secret: append([]byte(nil), config.Secret...), http: &clone,
websocket: dialer, now: config.Now,
}, nil
}
func requestTarget(origin *url.URL, request *http.Request) *url.URL {
target := *origin
target.Path = request.URL.Path
target.RawPath = request.URL.RawPath
target.RawQuery = request.URL.RawQuery
return &target
}
func transcript(method, requestURI, source, target, timestamp string) string {
return "nekonest-cloud/internal-forward/v1\x00" + strings.ToUpper(method) + "\x00" +
requestURI + "\x00" + source + "\x00" + target + "\x00" + timestamp
}
func signature(secret []byte, method, requestURI, source, target, timestamp string) string {
mac := hmac.New(sha256.New, secret)
_, _ = mac.Write([]byte(transcript(method, requestURI, source, target, timestamp)))
return hex.EncodeToString(mac.Sum(nil))
}
func stripForwardHeaders(header http.Header) {
for _, name := range []string{headerVersion, headerSource, headerTarget, headerTimestamp, headerSignature} {
header.Del(name)
}
}
func (forwarder *Forwarder) sign(header http.Header, method, requestURI, targetNodeID string) {
stripForwardHeaders(header)
timestamp := strconv.FormatInt(forwarder.now().Unix(), 10)
header.Set(headerVersion, "v1")
header.Set(headerSource, forwarder.nodeID)
header.Set(headerTarget, targetNodeID)
header.Set(headerTimestamp, timestamp)
header.Set(headerSignature, signature(
forwarder.secret, method, requestURI, forwarder.nodeID, targetNodeID, timestamp,
))
}
func VerifyIncoming(request *http.Request, targetNodeID string, secret []byte, now time.Time) error {
values := []string{
request.Header.Get(headerVersion), request.Header.Get(headerSource),
request.Header.Get(headerTarget), request.Header.Get(headerTimestamp),
request.Header.Get(headerSignature),
}
present := false
for _, value := range values {
present = present || strings.TrimSpace(value) != ""
}
if !present {
return nil
}
if values[0] != "v1" || values[1] == "" || values[2] != targetNodeID || len(secret) < 32 {
return errors.New("invalid internal Relay forwarding identity")
}
timestamp, err := strconv.ParseInt(values[3], 10, 64)
if err != nil || timestamp < now.Unix()-30 || timestamp > now.Unix()+30 {
return errors.New("internal Relay forwarding assertion expired")
}
expected := signature(secret, request.Method, request.URL.RequestURI(), values[1], values[2], values[3])
if len(values[4]) != len(expected) || subtle.ConstantTimeCompare([]byte(values[4]), []byte(expected)) != 1 {
return errors.New("invalid internal Relay forwarding signature")
}
return nil
}
func (forwarder *Forwarder) endpoint(reference string) (*url.URL, error) {
origin := forwarder.endpoints[reference]
if origin == nil {
return nil, errors.New("internal Relay endpoint reference is not configured")
}
clone := *origin
return &clone, nil
}
var hopHeaders = []string{
"Connection", "Proxy-Connection", "Keep-Alive", "Proxy-Authenticate",
"Proxy-Authorization", "Te", "Trailer", "Transfer-Encoding", "Upgrade",
}
func stripHopHeaders(header http.Header) {
for _, value := range header.Values("Connection") {
for _, token := range strings.Split(value, ",") {
header.Del(strings.TrimSpace(token))
}
}
for _, name := range hopHeaders {
header.Del(name)
}
}
func (forwarder *Forwarder) ForwardHTTP(w http.ResponseWriter, request *http.Request, endpointRef, targetNodeID string) error {
origin, err := forwarder.endpoint(endpointRef)
if err != nil {
return err
}
target := requestTarget(origin, request)
forwarded := request.Clone(request.Context())
forwarded.URL = target
forwarded.Host = target.Host
forwarded.RequestURI = ""
forwarded.Header = request.Header.Clone()
stripHopHeaders(forwarded.Header)
forwarder.sign(forwarded.Header, forwarded.Method, target.RequestURI(), targetNodeID)
response, err := forwarder.http.Do(forwarded)
if err != nil {
return err
}
defer response.Body.Close()
if response.StatusCode >= 300 && response.StatusCode < 400 {
return errors.New("internal Relay attempted to redirect a client")
}
if response.ContentLength > maxHTTPResponse {
return errors.New("internal Relay response exceeds limit")
}
body, err := io.ReadAll(io.LimitReader(response.Body, maxHTTPResponse+1))
if err != nil {
return err
}
if len(body) > maxHTTPResponse {
return errors.New("internal Relay response exceeds limit")
}
for key, values := range response.Header {
for _, value := range values {
w.Header().Add(key, value)
}
}
stripHopHeaders(w.Header())
w.WriteHeader(response.StatusCode)
_, err = w.Write(body)
return err
}
func (forwarder *Forwarder) DialWebSocket(
ctx context.Context, request *http.Request, endpointRef, targetNodeID string,
) (*websocket.Conn, *http.Response, error) {
origin, err := forwarder.endpoint(endpointRef)
if err != nil {
return nil, nil, err
}
target := requestTarget(origin, request)
if target.Scheme == "https" {
target.Scheme = "wss"
} else {
target.Scheme = "ws"
}
header := request.Header.Clone()
stripHopHeaders(header)
header.Del("Sec-WebSocket-Key")
header.Del("Sec-WebSocket-Version")
header.Del("Sec-WebSocket-Extensions")
header.Del("Sec-WebSocket-Protocol")
forwarder.sign(header, request.Method, target.RequestURI(), targetNodeID)
return forwarder.websocket.DialContext(ctx, target.String(), header)
}
func Tunnel(ctx context.Context, client, target *websocket.Conn, firstType int, firstFrame []byte) error {
client.SetReadLimit(maxWSMessage)
target.SetReadLimit(maxWSMessage)
if err := target.WriteMessage(firstType, firstFrame); err != nil {
return err
}
done := make(chan error, 2)
pump := func(destination, source *websocket.Conn) {
for {
messageType, message, err := source.ReadMessage()
if err != nil {
done <- err
return
}
if err := destination.WriteMessage(messageType, message); err != nil {
done <- err
return
}
}
}
go pump(target, client)
go pump(client, target)
select {
case <-ctx.Done():
_ = client.Close()
_ = target.Close()
return ctx.Err()
case err := <-done:
_ = client.Close()
_ = target.Close()
return err
}
}