Initial commit

This commit is contained in:
nullcat
2026-05-17 18:56:39 +08:00
commit f0f9a23813
9 changed files with 858 additions and 0 deletions
+81
View File
@@ -0,0 +1,81 @@
package core
import (
"bytes"
"os"
"github.com/BurntSushi/toml"
)
type ForwardRule struct {
Protocol string `toml:"protocol"`
TailscalePort int `toml:"tailscale_port"`
LocalAddr string `toml:"local_addr"`
}
type ConnectRule struct {
Protocol string `toml:"protocol"`
LocalPort int `toml:"local_port"`
LocalAddr string `toml:"local_addr"`
DstAddr string `toml:"dst_addr"`
LanEnable *bool `toml:"lan_enable"`
LanMotd string `toml:"lan_motd"`
}
type Core struct {
AuthKey string `toml:"auth_key"`
ControlURL string `toml:"control_url"`
Hostname string `toml:"hostname"`
Ephemeral bool `toml:"ephemeral"`
AcceptRoutes bool `toml:"accept_routes"`
}
func (r ConnectRule) LANEnabled() bool {
if r.LanEnable != nil {
return *r.LanEnable
}
return r.Protocol == "minecraft"
}
func (r ConnectRule) LANMotdOr(def string) string {
if r.LanMotd != "" {
return r.LanMotd
}
return def
}
type Config struct {
Core Core `toml:"core"`
Forward map[string][]ForwardRule `toml:"forward"`
Connect map[string][]ConnectRule `toml:"connect"`
}
func LoadConfig(path string) (*Config, error) {
hostname, err := os.Hostname()
if err != nil {
hostname = "unknown"
}
cfg := &Config{
Core: Core{
Hostname: hostname,
Ephemeral: true,
AcceptRoutes: true,
},
Forward: make(map[string][]ForwardRule),
Connect: make(map[string][]ConnectRule),
}
if _, err = os.Stat(path); err != nil {
// write a basic config
buf := new(bytes.Buffer)
err = toml.NewEncoder(buf).Encode(cfg)
if err != nil {
return nil, err
}
err = os.WriteFile(path, buf.Bytes(), 0644)
}
_, err = toml.DecodeFile(path, cfg)
if err != nil {
return nil, err
}
return cfg, nil
}
+429
View File
@@ -0,0 +1,429 @@
package core
import (
"context"
"fmt"
"io"
"log/slog"
"net"
"sync"
"time"
"tailscale.com/tsnet"
)
func StartForwarders(ctx context.Context, srv *tsnet.Server, rules map[string][]ForwardRule) {
for tag, rrs := range rules {
for _, rule := range rrs {
rule := rule
tag := tag
slog.Debug("starting forwarder",
slog.String("tag", tag),
slog.String("protocol", rule.Protocol),
slog.Int("tailscale_port", rule.TailscalePort),
slog.String("local_addr", rule.LocalAddr),
)
go runForwarder(ctx, srv, rule, tag)
}
}
}
func StartConnectors(ctx context.Context, srv *tsnet.Server, rules map[string][]ConnectRule) {
for tag, rrs := range rules {
for _, rule := range rrs {
rule := rule
tag := tag
args := []any{
slog.String("tag", tag),
slog.String("protocol", rule.Protocol),
slog.Int("local_port", rule.LocalPort),
slog.String("dst_addr", rule.DstAddr),
}
if rule.LocalAddr != "" {
args = append(args, slog.String("local_addr", rule.LocalAddr))
}
slog.Debug("starting connector", args...)
go runConnector(ctx, srv, rule, tag)
}
}
}
func RuleLogger(rule any, tag string) *slog.Logger {
var args []any
switch r := rule.(type) {
case ForwardRule:
args = []any{
slog.String("protocol", r.Protocol),
slog.Int("tailscale_port", r.TailscalePort),
slog.String("local_addr", r.LocalAddr),
}
case ConnectRule:
args = []any{
slog.String("protocol", r.Protocol),
slog.Int("local_port", r.LocalPort),
slog.String("dst_addr", r.DstAddr),
}
if r.LocalAddr != "" {
args = append(args, slog.String("local_addr", r.LocalAddr))
}
}
if tag != "" {
args = append(args, slog.String("tag", tag))
}
return slog.With(args...)
}
func runForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, tag string) {
logger := RuleLogger(rule, tag)
switch rule.Protocol {
case "tcp":
runTCPForwarder(ctx, srv, rule, logger)
case "udp":
runUDPForwarder(ctx, srv, rule, logger)
default:
logger.Error("unsupported protocol, expected tcp or udp")
}
}
func runTCPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, logger *slog.Logger) {
ip4, ip6 := srv.TailscaleIPs()
ip := ip4
if !ip.IsValid() {
ip = ip6
}
ln, err := srv.Listen("tcp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort))
if err != nil {
logger.Error("failed to listen", "error", err)
return
}
logger.Info("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort)))
go func() {
<-ctx.Done()
ln.Close()
}()
for {
conn, err := ln.Accept()
if err != nil {
if ctx.Err() != nil {
return
}
logger.Error("accept error", "error", err)
continue
}
go handleTCPForward(ctx, srv, conn, rule, logger)
}
}
func handleTCPForward(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ForwardRule, logger *slog.Logger) {
remoteAddrStr := conn.RemoteAddr().String()
clog := logger.With(slog.String("remote", remoteAddrStr))
lc, err := srv.LocalClient()
if err == nil {
who, err := lc.WhoIs(ctx, remoteAddrStr)
if err == nil {
clog = clog.With(slog.String("user", who.UserProfile.LoginName))
}
}
connType := getConnType(ctx, srv, remoteAddrStr)
clog.Info("accepted connection",
slog.String("conn_type", connType),
slog.String("local_addr", rule.LocalAddr),
)
defer conn.Close()
localConn, err := net.Dial("tcp", rule.LocalAddr)
if err != nil {
clog.Error("failed to dial local", "error", err)
return
}
defer localConn.Close()
toTs, toLocal := pipeConns(conn, localConn)
clog.Info("connection closed", slog.Int64("ts_rx_bytes", toLocal), slog.Int64("ts_tx_bytes", toTs))
}
func getConnType(ctx context.Context, srv *tsnet.Server, remoteAddrStr string) string {
lc, err := srv.LocalClient()
if err != nil {
return "unknown"
}
status, err := lc.Status(ctx)
if err != nil {
return "unknown"
}
remoteHost, _, err := net.SplitHostPort(remoteAddrStr)
if err != nil {
return "unknown"
}
for _, peer := range status.Peer {
for _, addr := range peer.TailscaleIPs {
if addr.String() == remoteHost {
if peer.Relay != "" {
return fmt.Sprintf("derp(%s)", peer.Relay)
}
return "direct"
}
}
}
return "unknown"
}
func runUDPForwarder(ctx context.Context, srv *tsnet.Server, rule ForwardRule, logger *slog.Logger) {
ip4, ip6 := srv.TailscaleIPs()
ip := ip4
if !ip.IsValid() {
ip = ip6
}
pc, err := srv.ListenPacket("udp", fmt.Sprintf("%s:%d", ip.String(), rule.TailscalePort))
if err != nil {
logger.Error("failed to listen", "error", err)
return
}
logger.Info("listening", slog.String("on", fmt.Sprintf("tailscale:%s:%d", ip.String(), rule.TailscalePort)))
relay := &udpRelay{
listenConn: pc,
dialAddr: rule.LocalAddr,
logger: logger,
direction: "local",
sessions: make(map[string]*udpSession),
}
relay.run(ctx)
}
type udpSession struct {
conn net.Conn
remote net.Addr
lastUse time.Time
}
type udpRelay struct {
listenConn net.PacketConn
dialAddr string
logger *slog.Logger
direction string
mu sync.Mutex
sessions map[string]*udpSession
}
func (r *udpRelay) run(ctx context.Context) {
go func() {
<-ctx.Done()
r.listenConn.Close()
}()
go func() {
ticker := time.NewTicker(2 * time.Minute)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
r.cleanup()
}
}
}()
buf := make([]byte, 65535)
for {
select {
case <-ctx.Done():
return
default:
}
n, from, err := r.listenConn.ReadFrom(buf)
if err != nil {
if ctx.Err() != nil {
return
}
r.logger.Error("udp read error", "error", err)
return
}
key := from.String()
r.mu.Lock()
sess, exists := r.sessions[key]
if !exists {
dialed, err := net.Dial("udp", r.dialAddr)
if err != nil {
r.mu.Unlock()
r.logger.Error("failed to dial", "error", err)
continue
}
sess = &udpSession{conn: dialed, remote: from, lastUse: time.Now()}
r.sessions[key] = sess
r.mu.Unlock()
r.logger.Debug("new udp session", slog.String("remote", key), slog.String("direction", r.direction))
go r.readSession(key, sess)
} else {
sess.lastUse = time.Now()
r.mu.Unlock()
}
if _, err := sess.conn.Write(buf[:n]); err != nil {
r.logger.Error("failed to write", "error", err)
r.removeSession(key)
}
}
}
func (r *udpRelay) readSession(key string, sess *udpSession) {
buf := make([]byte, 65535)
for {
n, err := sess.conn.Read(buf)
if err != nil {
r.removeSession(key)
return
}
if _, err := r.listenConn.WriteTo(buf[:n], sess.remote); err != nil {
r.logger.Error("failed to write back", "error", err)
r.removeSession(key)
return
}
}
}
func (r *udpRelay) removeSession(key string) {
r.mu.Lock()
defer r.mu.Unlock()
if s, ok := r.sessions[key]; ok {
s.conn.Close()
delete(r.sessions, key)
r.logger.Debug("udp session closed", slog.String("remote", key))
}
}
func (r *udpRelay) cleanup() {
r.mu.Lock()
defer r.mu.Unlock()
threshold := time.Now().Add(-5 * time.Minute)
for key, s := range r.sessions {
if s.lastUse.Before(threshold) {
s.conn.Close()
delete(r.sessions, key)
r.logger.Debug("udp session cleaned up", slog.String("remote", key))
}
}
}
func runConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, tag string) {
logger := RuleLogger(rule, tag)
switch rule.Protocol {
case "tcp", "minecraft":
runTCPConnector(ctx, srv, rule, logger)
case "udp":
runUDPConnector(ctx, srv, rule, logger)
default:
logger.Error("unsupported protocol, expected tcp or udp")
}
}
func runTCPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, logger *slog.Logger) {
bindIP := rule.LocalAddr
if bindIP == "" {
bindIP = "0.0.0.0"
}
if rule.LANEnabled() && bindIP != "0.0.0.0" {
logger.Warn("lan_enable forces local_addr to 0.0.0.0, overriding")
bindIP = "0.0.0.0"
}
addr := fmt.Sprintf("%s:%d", bindIP, rule.LocalPort)
ln, err := net.Listen("tcp", addr)
if err != nil {
logger.Error("failed to listen locally", "error", err)
return
}
logger.Info("listening", slog.String("on", addr))
go func() {
<-ctx.Done()
ln.Close()
}()
for {
conn, err := ln.Accept()
if err != nil {
if ctx.Err() != nil {
return
}
logger.Error("accept error", "error", err)
continue
}
go handleTCPConnect(ctx, srv, conn, rule, logger)
}
}
func handleTCPConnect(ctx context.Context, srv *tsnet.Server, conn net.Conn, rule ConnectRule, logger *slog.Logger) {
clog := logger.With(slog.String("local_client", conn.RemoteAddr().String()))
defer conn.Close()
tsConn, err := srv.Dial(ctx, "tcp", rule.DstAddr)
if err != nil {
clog.Error("failed to dial tailscale", "error", err)
return
}
defer tsConn.Close()
clog.Info("accepted connection", slog.String("dst_addr", rule.DstAddr))
toConn, toTs := pipeConns(conn, tsConn)
clog.Info("connection closed", slog.Int64("ts_rx_bytes", toTs), slog.Int64("ts_tx_bytes", toConn))
}
func runUDPConnector(ctx context.Context, srv *tsnet.Server, rule ConnectRule, logger *slog.Logger) {
bindIP := rule.LocalAddr
if bindIP == "" {
bindIP = "0.0.0.0"
}
addr := fmt.Sprintf("%s:%d", bindIP, rule.LocalPort)
addrUDP, err := net.ResolveUDPAddr("udp", addr)
if err != nil {
logger.Error("failed to resolve local addr", "error", err)
return
}
pc, err := net.ListenUDP("udp", addrUDP)
if err != nil {
logger.Error("failed to listen locally", "error", err)
return
}
logger.Info("listening", slog.String("on", addr))
relay := &udpRelay{
listenConn: pc,
dialAddr: rule.DstAddr,
logger: logger,
direction: "tailscale",
sessions: make(map[string]*udpSession),
}
relay.run(ctx)
}
func pipeConns(a, b net.Conn) (toA, toB int64) {
var wg sync.WaitGroup
wg.Add(2)
go func() {
defer wg.Done()
n, _ := io.Copy(a, b)
toA = n
}()
go func() {
defer wg.Done()
n, _ := io.Copy(b, a)
toB = n
}()
wg.Wait()
return
}
+94
View File
@@ -0,0 +1,94 @@
package core
import (
"context"
"fmt"
"net"
"time"
"log/slog"
)
type LanEntry struct {
Motd string
Port int
}
func LanDiscoverService(ctx context.Context, entryList []LanEntry, logger *slog.Logger) {
mcastAddrs := []string{
"224.0.2.60:4445",
"[ff75:230::60]:4445",
}
var fdList []*net.UDPConn
for _, addrStr := range mcastAddrs {
addr, err := net.ResolveUDPAddr("udp", addrStr)
if err != nil {
logger.With(slog.String("error", err.Error())).Error("failed to resolve udp address")
continue
}
fd, err := net.DialUDP("udp", nil, addr)
if err != nil {
logger.With(slog.String("error", err.Error())).Error("failed to dial udp server")
continue
}
fdList = append(fdList, fd)
}
if len(fdList) == 0 {
// init fail
logger.Warn("all multicast binding failed, service discovery is disabled")
return
}
for _, entry := range entryList {
// debug log to print entry detail
logger.With(
slog.Int("port", entry.Port),
slog.String("motd", entry.Motd),
).Debug("discover service: %s on %d", entry.Motd, entry.Port)
}
ticker := time.NewTicker(1500 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
logger.Debug("shutting down lan discovery service")
for _, fd := range fdList {
fd.Close()
}
return
case <-ticker.C:
for _, e := range entryList {
msg := fmt.Sprintf("[MOTD]%s[/MOTD][AD]%d[/AD]", e.Motd, e.Port)
for _, c := range fdList {
_, err := c.Write([]byte(msg))
if err != nil {
logger.With(slog.String("error", err.Error())).Error("failed to write to udp server")
return
}
}
}
}
}
}
func RunLanDiscoverService(ctx context.Context, rules map[string][]ConnectRule, logger *slog.Logger) {
var lanEntries []LanEntry
for tag, rs := range rules {
for _, rule := range rs {
if !rule.LANEnabled() {
continue
}
motd := rule.LANMotdOr(tag)
lanEntries = append(lanEntries, LanEntry{
Motd: motd,
Port: rule.LocalPort,
})
}
}
go LanDiscoverService(ctx, lanEntries, logger)
}
+44
View File
@@ -0,0 +1,44 @@
package core
import (
"os"
"strings"
"time"
"log/slog"
"github.com/lmittmann/tint"
"github.com/mattn/go-colorable"
)
func parseLevel(s string) slog.Level {
switch strings.ToLower(s) {
case "debug":
return slog.LevelDebug
case "warn", "warning":
return slog.LevelWarn
case "error":
return slog.LevelError
default:
return slog.LevelInfo
}
}
func NewLogger(level string, useJsonFormat bool) *slog.Logger {
w := os.Stdout
var logger *slog.Logger
if !useJsonFormat {
logger = slog.New(tint.NewHandler(colorable.NewColorable(w), &tint.Options{
Level: parseLevel(level),
TimeFormat: time.DateTime,
//NoColor: !isatty.IsTerminal(w.Fd()),
}))
} else {
logger = slog.New(slog.NewJSONHandler(w, &slog.HandlerOptions{
Level: parseLevel(level),
}))
}
slog.SetDefault(logger)
return logger
}
+65
View File
@@ -0,0 +1,65 @@
package core
import (
"context"
fmt2 "fmt"
"log/slog"
"tailscale.com/ipn"
"tailscale.com/tsnet"
)
func InitTsNet(ctx context.Context, cfg *Core, logger *slog.Logger) (*tsnet.Server, error) {
srv := &tsnet.Server{
Hostname: "tslink-" + cfg.Hostname,
AuthKey: cfg.AuthKey,
Ephemeral: cfg.Ephemeral,
Logf: func(fmt string, args ...interface{}) {
logger.With(slog.String("from", "tsnet")).Debug(fmt2.Sprintf(fmt, args...))
},
UserLogf: func(fmt string, args ...interface{}) {
logger.With(slog.String("from", "tsnet")).Info(fmt2.Sprintf(fmt, args...))
},
RunWebClient: true,
}
if cfg.ControlURL != "" {
srv.ControlURL = cfg.ControlURL
}
logger.Debug("starting tsnet server")
if err := srv.Start(); err != nil {
logger.With(slog.String("error", err.Error())).Error("starting tsnet server failed")
return nil, err
}
status, err := srv.Up(ctx)
if err != nil {
logger.With(slog.String("error", err.Error())).Error("bring up tsnet server failed")
return nil, err
}
for _, ip := range status.TailscaleIPs {
logger.With(slog.String("ip", ip.String())).Info("ip got from tsnet")
}
if cfg.AcceptRoutes {
lc, err := srv.LocalClient()
if err != nil {
logger.With(slog.String("error", err.Error())).Error("error from getting local client")
} else {
_, err = lc.EditPrefs(ctx, &ipn.MaskedPrefs{
Prefs: ipn.Prefs{RouteAll: true},
RouteAllSet: true,
})
if err != nil {
logger.With(slog.String("error", err.Error())).Error("error from editing prefs")
} else {
logger.Debug("subnet route accepted")
}
}
}
return srv, nil
}