diff options
| author | Mel <einebeere@gmail.com> | 2023-08-24 23:06:39 +0200 |
|---|---|---|
| committer | Mel <einebeere@gmail.com> | 2023-08-24 23:06:39 +0200 |
| commit | 27c39fce776c395711a9a70ca6a195515ec5ba5e (patch) | |
| tree | 7dd65575ca120378925101f058bf7f85eed1e261 /main.go | |
| parent | 50ae1c517fd5cd1e65afa7d1b9908fe3f46d2e63 (diff) | |
| download | cthcous-27c39fce776c395711a9a70ca6a195515ec5ba5e.tar.zst cthcous-27c39fce776c395711a9a70ca6a195515ec5ba5e.zip | |
Better log messages and more robust closing
Diffstat (limited to 'main.go')
| -rw-r--r-- | main.go | 101 |
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[:]) +} |
