about summary refs log tree commit diff
path: root/main.go
blob: 1d3a483bf6f07c03bda2639944358c8069d4ad6f (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
package main

import (
	"bytes"
	"context"
	"errors"
	"fmt"
	"io"
	"log/slog"
	"net"
	"os"
	"os/signal"
	"sync"
	"time"
)

const defaultReadTimeout = 2 * time.Second

func main() {
	slog.Info("starting cthcous...")

	ports := []uint{1234, 4321}
	socketDir := "./"

	wg := &sync.WaitGroup{}
	socketCtx, _ := 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))

		wg.Add(1)
		go socket(wg, socketCtx, socketPath, port)
	}

	wg.Wait()
}

func socket(wg *sync.WaitGroup, ctx context.Context, path string, port uint) {
	defer wg.Done()

	source, err := net.Listen("unix", path)
	if err != nil {
		slog.Error("could not listen on socket.", slog.Any("error", err))
		return
	}
	source.(*net.UnixListener).SetUnlinkOnClose(true)
	closeOnCtx(ctx, source)

	for {
		conn, err := source.Accept()
		if err != nil {
			if errors.Is(err, net.ErrClosed) {
				slog.Info("socket closed.", slog.String("path", path))
				break
			}

			slog.Error("could not accept connection.", slog.Any("error", err))
			continue
		}
		conn.SetReadDeadline(time.Now().Add(defaultReadTimeout))
		closeOnCtx(ctx, conn)

		slog.Info("accepted connection.", slog.String("remoteAddr", conn.RemoteAddr().String()))

		go bridge(ctx, conn, port)
	}
}

func bridge(ctx context.Context, source net.Conn, port uint) {
	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()))
				break
			}
			if errors.Is(err, os.ErrDeadlineExceeded) {
				slog.Warn("read deadline exceeded. closing connection.", slog.String("remoteAddr", source.RemoteAddr().String()))
				break
			}
			slog.Error("could not read from 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))
			break
		}
	}
	source.Close()
	slog.Info("connection completed. attempting write...", slog.String("remoteAddr", source.RemoteAddr().String()))

	destination, err := net.Dial("tcp", fmt.Sprintf(":%d", port))
	if err != nil {
		slog.Error("could not connect to destination socket.", 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))
		return
	}

	slog.Info("write to destination socket successful.", slog.String("remoteAddr", source.RemoteAddr().String()))
}

func closeOnCtx(ctx context.Context, 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))
		}
	}()
}