324 lines
9.7 KiB
Go
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
|
|
}
|
|
}
|