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{} ctx, _ := signal.NotifyContext(context.Background(), os.Interrupt, os.Kill) 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 openBridgeSocket(ctx, wg, bridgeLogger, bridgeInfo) } wg.Wait() } func openBridgeSocket(ctx context.Context, wg *sync.WaitGroup, logger *slog.Logger, info BridgeInfo) { defer wg.Done() source, err := net.Listen("unix", info.OriginSocketPath) if err != nil { logger.Error("could not listen on socket.", slog.Any("error", err)) return } source.(*net.UnixListener).SetUnlinkOnClose(true) closeOnCtx(ctx, logger, source) logger.Info("bridge socket opened. listening...") for { conn, err := source.Accept() if err != nil { if errors.Is(err, net.ErrClosed) { logger.Info("bridge socket closed.") break } logger.Error("could not accept connection.", slog.Any("error", err)) continue } conn.SetReadDeadline(time.Now().Add(defaultReadTimeout)) closeOnCtx(ctx, logger, conn) // 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))) connLogger.Info("accepted connection. beginning bridge.") go bridgeDataFromOriginToDestination(ctx, logger, conn, info) } } 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) { logger.Info("received EOF. closing connection...") break } if errors.Is(err, os.ErrDeadlineExceeded) { logger.Warn("read deadline exceeded. closing origin connection...") break } 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 { logger.Error("could not write to buffer due to unknown error. closing origin connection...", slog.Any("error", err)) break } } source.Close() logger.Info("origin connection completed. attempting write to destination...") destination, err := net.Dial("tcp", info.DestinationAddress) if err != nil { 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 { logger.Error("could not write to destination socket. bridge aborted.", slog.Any("error", err)) return } logger.Info("write to destination socket successful.") logger.Info("bridge completed.") } func closeOnCtx(ctx context.Context, logger *slog.Logger, closer io.Closer) { go func() { <-ctx.Done() if err := closer.Close(); err != nil { 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[:]) }