package main import ( "bytes" "context" "encoding/base64" "encoding/binary" "encoding/json" "errors" "flag" "fmt" "io" "log/slog" "math/rand" "net" "os" "os/signal" "path" "sync" "time" ) const version = "0.0.1" const defaultConfigPath = "/etc/cthcous/config.json" const defaultReadTimeout = 2 * time.Second type BridgeInfo struct { ID uint OriginSocketPath string DestinationAddress string } func main() { slog.Info(fmt.Sprintf("starting cthcous v%s...", version)) flags := parseFlags() if flags.ConfigPath == "" { slog.Warn("no config path specified. using default.", slog.String("path", defaultConfigPath)) flags.ConfigPath = defaultConfigPath } slog.Info("loading config...", slog.String("path", flags.ConfigPath)) config, err := loadConfig(flags.ConfigPath) if err != nil { slog.Error("could not load config.", slog.Any("error", err)) return } if len(config.Ports) == 0 { slog.Error("no ports specified in config. nothing to do. exiting.") return } bridgeInfos := make([]BridgeInfo, len(config.Ports)) for i, port := range config.Ports { var socketPath string if override, ok := config.PathOverrides[port]; ok { slog.Info("using override for a socket path.", slog.Any("port", port), slog.String("path", override)) socketPath = override } else { socketPath = fmt.Sprintf("%s/%d.socket", config.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() slog.Info("cthcous wait group completed. cya!") } 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[:]) } type Flags struct { ConfigPath string } func parseFlags() (flags Flags) { flag.StringVar(&flags.ConfigPath, "c", "", "path to config file") flag.Parse() return } type Config struct { Ports []uint `json:"ports"` SocketDir string `json:"socketDir"` PathOverrides map[uint]string `json:"pathOverrides"` } var defaultConfig = Config{ Ports: []uint{}, SocketDir: "/run/cthcous/", PathOverrides: map[uint]string{}, } func loadConfig(path string) (Config, error) { configFile, err := os.Open(path) if err != nil { if errors.Is(err, os.ErrNotExist) { return Config{}, fmt.Errorf("config file does not exist: %w", err) } return Config{}, err } defer configFile.Close() config := defaultConfig err = json.NewDecoder(configFile).Decode(&config) if err != nil { return Config{}, fmt.Errorf("could not decode config file: %w", err) } return config, nil }