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 } }