about summary refs log tree commit diff
diff options
context:
space:
mode:
-rw-r--r--main.go101
1 files changed, 73 insertions, 28 deletions
diff --git a/main.go b/main.go
index 1d3a483..79bb0d4 100644
--- a/main.go
+++ b/main.go
@@ -3,119 +3,164 @@ package main
 import (
 	"bytes"
 	"context"
+	"encoding/base64"
+	"encoding/binary"
 	"errors"
 	"fmt"
 	"io"
 	"log/slog"
+	"math/rand"
 	"net"
 	"os"
 	"os/signal"
+	"path"
 	"sync"
 	"time"
 )
 
 const defaultReadTimeout = 2 * time.Second
 
+type BridgeInfo struct {
+	ID                 uint
+	OriginSocketPath   string
+	DestinationAddress string
+}
+
 func main() {
 	slog.Info("starting cthcous...")
 
 	ports := []uint{1234, 4321}
 	socketDir := "./"
 
+	bridgeInfos := make([]BridgeInfo, len(ports))
+	for i, port := range ports {
+		socketPath := fmt.Sprintf("%s/%d.socket", socketDir, port)
+		socketPath = path.Clean(socketPath)
+
+		bridgeInfos[i] = BridgeInfo{
+			ID:                 uint(i),
+			OriginSocketPath:   socketPath,
+			DestinationAddress: fmt.Sprintf(":%d", port),
+		}
+	}
+
 	wg := &sync.WaitGroup{}
-	socketCtx, _ := signal.NotifyContext(context.Background(), os.Interrupt, os.Kill)
+	ctx, _ := signal.NotifyContext(context.Background(), os.Interrupt, os.Kill)
 
-	slog.Info("will open sockets.", slog.Any("ports", ports), slog.String("socketDir", socketDir))
-	for _, port := range ports {
-		socketPath := fmt.Sprintf("%s/%d.socket", socketDir, port)
-		slog.Info("opening socket...", slog.Uint64("port", uint64(port)), slog.String("socketPath", socketPath))
+	slog.Info("will open socket bridges.", slog.Any("bridges", bridgeInfos))
+	for _, bridgeInfo := range bridgeInfos {
+		bridgeLogger := slog.Default().With(slog.Group(
+			"bridge",
+			slog.Any("id", bridgeInfo.ID),
+			slog.String("origin", bridgeInfo.OriginSocketPath),
+			slog.String("destination", bridgeInfo.DestinationAddress),
+		))
+
+		bridgeLogger.Info("opening bridge socket...")
 
 		wg.Add(1)
-		go socket(wg, socketCtx, socketPath, port)
+		go openBridgeSocket(ctx, wg, bridgeLogger, bridgeInfo)
 	}
 
 	wg.Wait()
 }
 
-func socket(wg *sync.WaitGroup, ctx context.Context, path string, port uint) {
+func openBridgeSocket(ctx context.Context, wg *sync.WaitGroup, logger *slog.Logger, info BridgeInfo) {
 	defer wg.Done()
 
-	source, err := net.Listen("unix", path)
+	source, err := net.Listen("unix", info.OriginSocketPath)
 	if err != nil {
-		slog.Error("could not listen on socket.", slog.Any("error", err))
+		logger.Error("could not listen on socket.", slog.Any("error", err))
 		return
 	}
 	source.(*net.UnixListener).SetUnlinkOnClose(true)
-	closeOnCtx(ctx, source)
+	closeOnCtx(ctx, logger, source)
+
+	logger.Info("bridge socket opened. listening...")
 
 	for {
 		conn, err := source.Accept()
 		if err != nil {
 			if errors.Is(err, net.ErrClosed) {
-				slog.Info("socket closed.", slog.String("path", path))
+				logger.Info("bridge socket closed.")
 				break
 			}
 
-			slog.Error("could not accept connection.", slog.Any("error", err))
+			logger.Error("could not accept connection.", slog.Any("error", err))
 			continue
 		}
 		conn.SetReadDeadline(time.Now().Add(defaultReadTimeout))
-		closeOnCtx(ctx, conn)
+		closeOnCtx(ctx, logger, conn)
 
-		slog.Info("accepted connection.", slog.String("remoteAddr", conn.RemoteAddr().String()))
+		// Connections to UNIX sockets don't have remote addresses.
+		// To make it easier to identify connections, we generate a random ID.
+		connID := makeID()
+		connLogger := logger.With(slog.Group("connection", slog.String("id", connID)))
 
-		go bridge(ctx, conn, port)
+		connLogger.Info("accepted connection. beginning bridge.")
+		go bridgeDataFromOriginToDestination(ctx, logger, conn, info)
 	}
 }
 
-func bridge(ctx context.Context, source net.Conn, port uint) {
+func bridgeDataFromOriginToDestination(ctx context.Context, logger *slog.Logger, source net.Conn, info BridgeInfo) {
 	data := bytes.NewBuffer([]byte{})
 	for {
 		buffer := make([]byte, 1024)
 		n, err := source.Read(buffer)
 		if err != nil {
 			if errors.Is(err, io.EOF) {
-				slog.Info("got EOF...", slog.String("remoteAddr", source.RemoteAddr().String()))
+				logger.Info("received EOF. closing connection...")
 				break
 			}
 			if errors.Is(err, os.ErrDeadlineExceeded) {
-				slog.Warn("read deadline exceeded. closing connection.", slog.String("remoteAddr", source.RemoteAddr().String()))
+				logger.Warn("read deadline exceeded. closing origin connection...")
 				break
 			}
-			slog.Error("could not read from connection.", slog.Any("error", err))
+			logger.Error("could not read from connection due to unknown error. closing origin connection...", slog.Any("error", err))
 			break
 		}
 
 		_, err = data.Write(buffer[:n])
 		if err != nil {
-			slog.Error("could not write to buffer.", slog.Any("error", err))
+			logger.Error("could not write to buffer due to unknown error. closing origin connection...", slog.Any("error", err))
 			break
 		}
 	}
 	source.Close()
-	slog.Info("connection completed. attempting write...", slog.String("remoteAddr", source.RemoteAddr().String()))
+	logger.Info("origin connection completed. attempting write to destination...")
 
-	destination, err := net.Dial("tcp", fmt.Sprintf(":%d", port))
+	destination, err := net.Dial("tcp", info.DestinationAddress)
 	if err != nil {
-		slog.Error("could not connect to destination socket.", slog.Any("error", err))
+		logger.Error("could not connect to destination socket. bridge aborted.", slog.Any("error", err))
 		return
 	}
 	defer destination.Close()
 
 	_, err = destination.Write(data.Bytes())
 	if err != nil {
-		slog.Error("could not write to destination socket.", slog.Any("error", err))
+		logger.Error("could not write to destination socket. bridge aborted.", slog.Any("error", err))
 		return
 	}
-
-	slog.Info("write to destination socket successful.", slog.String("remoteAddr", source.RemoteAddr().String()))
+	logger.Info("write to destination socket successful.")
+	logger.Info("bridge completed.")
 }
 
-func closeOnCtx(ctx context.Context, closer io.Closer) {
+func closeOnCtx(ctx context.Context, logger *slog.Logger, closer io.Closer) {
 	go func() {
 		<-ctx.Done()
 		if err := closer.Close(); err != nil {
-			slog.Error("could not close after finished context.", slog.Any("error", err))
+			if errors.Is(err, net.ErrClosed) {
+				logger.Warn("socket or connection already closed. harmless race condition.")
+				return
+			}
+			logger.Error("could not close socket or connection after finished context.", slog.Any("error", err))
 		}
 	}()
 }
+
+func makeID() string {
+	r := rand.Uint64()
+	b := [8]byte{}
+	binary.LittleEndian.PutUint64(b[:], r)
+	return base64.StdEncoding.EncodeToString(b[:])
+}