feat: establish NekoNest Cloud control and relay
This commit is contained in:
@@ -0,0 +1,323 @@
|
||||
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
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user