package network import ( "context" "crypto/rand" "crypto/rsa" "crypto/tls" "crypto/x509" "crypto/x509/pkix" "fmt" "log" "math/big" "net" "os/exec" "sync" "time" "github.com/apernet/quic-go" ) type NetworkServer struct { config Config tun *TUNInterface pool *IPPool ln *quic.Listener clients map[string]*networkClient clientsMu sync.Mutex ctx context.Context cancel context.CancelFunc wg sync.WaitGroup } type networkClient struct { ip net.IP stream *quic.Stream } func NewNetworkServer(config Config) (*NetworkServer, error) { ip, ipnet, err := net.ParseCIDR(config.Pool) if err != nil { return nil, fmt.Errorf("parse pool: %w", err) } pool, err := NewIPPool(ip, ipnet) if err != nil { return nil, fmt.Errorf("create pool: %w", err) } tunName := config.TUN.Name if tunName == "" { tunName = defaultTUNName } mtu := config.TUN.MTU if mtu <= 0 { mtu = defaultTUNMTU } tunDev, err := OpenTUN(tunName, pool.Gateway().String()+"/24", mtu) if err != nil { return nil, fmt.Errorf("open TUN: %w", err) } ctx, cancel := context.WithCancel(context.Background()) return &NetworkServer{ config: config, tun: tunDev, pool: pool, clients: make(map[string]*networkClient), ctx: ctx, cancel: cancel, }, nil } func (s *NetworkServer) Serve() error { listenAddr := s.config.Listen if listenAddr == "" { listenAddr = defaultNetworkListen } tlsCfg := s.config.TLS if tlsCfg == nil { var err error tlsCfg, err = GenerateNetworkTLSConfig() if err != nil { return fmt.Errorf("generate TLS: %w", err) } } ln, err := quic.ListenAddr(listenAddr, tlsCfg, nil) if err != nil { return fmt.Errorf("listen: %w", err) } s.ln = ln log.Printf("network server listening on %s", listenAddr) log.Printf("TUN %s: %s", s.tun.Name, s.pool.Gateway()) log.Printf("IP pool: %s", s.config.Pool) s.wg.Add(1) go s.tunToClients() for { conn, err := ln.Accept(s.ctx) if err != nil { if s.ctx.Err() != nil { return nil } return fmt.Errorf("accept: %w", err) } go s.handleConn(conn) } } func (s *NetworkServer) handleConn(conn *quic.Conn) { stream, err := conn.AcceptStream(s.ctx) if err != nil { log.Printf("accept stream: %v", err) return } assignedIP, err := s.handleAuth(stream) if err != nil { log.Printf("auth error from %s: %v", conn.RemoteAddr(), err) stream.Close() return } s.tun.Route(assignedIP.String() + "/32") log.Printf("client %s authenticated, IP: %s", conn.RemoteAddr(), assignedIP) nc := &networkClient{ ip: assignedIP, stream: stream, } s.clientsMu.Lock() s.clients[assignedIP.String()] = nc s.clientsMu.Unlock() stopCh := make(chan struct{}) go func() { ticker := time.NewTicker(defaultKeepalive) defer ticker.Stop() for { select { case <-stopCh: return case <-ticker.C: SendKeepalive(stream) } } }() s.clientToTun(nc) close(stopCh) s.clientsMu.Lock() delete(s.clients, assignedIP.String()) s.clientsMu.Unlock() s.pool.Release(assignedIP) exec.Command("ip", "route", "del", assignedIP.String()+"/32", "dev", s.tun.Name).Run() log.Printf("client %s (%s) disconnected", conn.RemoteAddr(), assignedIP) } func (s *NetworkServer) handleAuth(stream *quic.Stream) (net.IP, error) { f, err := ReceiveFrame(stream) if err != nil { return nil, fmt.Errorf("recv auth: %w", err) } if f.Type != FrameAuth { return nil, fmt.Errorf("expected auth, got %d", f.Type) } if string(f.Payload) != s.config.Token { SendAuthErr(stream, "invalid token") return nil, fmt.Errorf("invalid token") } assignedIP := s.pool.Allocate() if assignedIP == nil { SendAuthErr(stream, "no IP available") return nil, fmt.Errorf("no IP available") } if err := SendAuthOK(stream, assignedIP); err != nil { s.pool.Release(assignedIP) return nil, fmt.Errorf("send auth ok: %w", err) } return assignedIP, nil } func (s *NetworkServer) clientToTun(nc *networkClient) { for { f, err := ReceiveFrame(nc.stream) if err != nil { return } switch f.Type { case FrameData: if err := s.tun.Write(f.Payload); err != nil { log.Printf("write TUN: %v", err) return } case FrameKeepalive: default: } } } func (s *NetworkServer) tunToClients() { defer s.wg.Done() for { pkt, err := s.tun.Read() if err != nil { if s.ctx.Err() != nil { return } log.Printf("read TUN: %v", err) return } if len(pkt) < 20 { continue } dstIP := net.IP(pkt[16:20]).String() s.clientsMu.Lock() nc, ok := s.clients[dstIP] s.clientsMu.Unlock() if !ok { continue } if err := SendData(nc.stream, pkt); err != nil { log.Printf("send to %s: %v", dstIP, err) return } } } func (s *NetworkServer) Close() error { s.cancel() if s.ln != nil { s.ln.Close() } s.tun.Close() s.wg.Wait() return nil } func GenerateNetworkTLSConfig() (*tls.Config, error) { key, err := rsa.GenerateKey(rand.Reader, 2048) if err != nil { return nil, fmt.Errorf("generate key: %w", err) } template := x509.Certificate{ SerialNumber: big.NewInt(1), Subject: pkix.Name{ Organization: []string{"Hysteria Network"}, }, NotBefore: time.Now(), NotAfter: time.Now().Add(10 * 365 * 24 * time.Hour), KeyUsage: x509.KeyUsageKeyEncipherment | x509.KeyUsageDigitalSignature, ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth}, BasicConstraintsValid: true, } template.DNSNames = append(template.DNSNames, "localhost") certDER, err := x509.CreateCertificate(rand.Reader, &template, &template, &key.PublicKey, key) if err != nil { return nil, fmt.Errorf("create cert: %w", err) } cert := tls.Certificate{ Certificate: [][]byte{certDER}, PrivateKey: key, } return &tls.Config{ Certificates: []tls.Certificate{cert}, NextProtos: []string{"hysteria-network"}, }, nil }