diff --git a/cmd/dashboard/dashboard_listener.go b/cmd/dashboard/dashboard_listener.go new file mode 100644 index 00000000..ca38c6a5 --- /dev/null +++ b/cmd/dashboard/dashboard_listener.go @@ -0,0 +1,15 @@ +package main + +import "fmt" + +type dashboardListenerKind string + +const ( + dashboardHTTPListener dashboardListenerKind = "http" + dashboardHTTPSListener dashboardListenerKind = "https" + dashboardReceiptListener dashboardListenerKind = "receipt" +) + +func dashboardListenerAddress(host string, port uint16) string { + return fmt.Sprintf("%s:%d", host, port) +} diff --git a/cmd/dashboard/dashboard_listener_default.go b/cmd/dashboard/dashboard_listener_default.go new file mode 100644 index 00000000..7a32947e --- /dev/null +++ b/cmd/dashboard/dashboard_listener_default.go @@ -0,0 +1,18 @@ +//go:build !agentcompat + +package main + +import ( + "net" + "net/http" +) + +func openReceiptGateListener() (net.Listener, error) { return nil, nil } + +func openDashboardListener(network, address string, _ dashboardListenerKind) (net.Listener, error) { + return net.Listen(network, address) +} + +func serveDashboardHTTPS(server *http.Server, certificatePath, keyPath string) error { + return server.ListenAndServeTLS(certificatePath, keyPath) +} diff --git a/cmd/dashboard/dashboard_listener_integration.go b/cmd/dashboard/dashboard_listener_integration.go new file mode 100644 index 00000000..006bb765 --- /dev/null +++ b/cmd/dashboard/dashboard_listener_integration.go @@ -0,0 +1,178 @@ +//go:build agentcompat + +package main + +import ( + "errors" + "fmt" + "net" + "net/http" + "os" + "strconv" +) + +const ( + dashboardHTTPListenerFDEnv = "NEZHA_AGENTCOMPAT_HTTP_LISTENER_FD" + dashboardHTTPSListenerFDEnv = "NEZHA_AGENTCOMPAT_HTTPS_LISTENER_FD" + dashboardReceiptListenerFDEnv = "NEZHA_AGENTCOMPAT_RECEIPT_LISTENER_FD" +) + +func openReceiptGateListener() (net.Listener, error) { + if os.Getenv(dashboardReceiptListenerFDEnv) == "" { + return nil, nil + } + return openDashboardListener("tcp", "127.0.0.1:0", dashboardReceiptListener) +} + +var ( + errDashboardListenerKind = errors.New("unsupported dashboard listener kind") + errDashboardListenerTCP = errors.New("inherited dashboard listener is not TCP") + errDashboardListenerLoopback = errors.New("inherited dashboard listener is not loopback") +) + +type dashboardListenerError struct { + Kind dashboardListenerKind + EnvironmentVariable string + Descriptor string + Err error +} + +func (e *dashboardListenerError) Error() string { + return fmt.Sprintf("adopt %s dashboard listener from %s=%q: %v", e.Kind, e.EnvironmentVariable, e.Descriptor, e.Err) +} + +func (e *dashboardListenerError) Unwrap() error { + return e.Err +} + +func openDashboardListener(network, address string, kind dashboardListenerKind) (net.Listener, error) { + environmentVariable, err := dashboardListenerEnvironmentVariable(kind) + if err != nil { + return nil, err + } + descriptor := os.Getenv(environmentVariable) + if descriptor == "" { + return net.Listen(network, address) + } + + fileDescriptor, err := strconv.Atoi(descriptor) + if err != nil || fileDescriptor < 0 { + if err == nil { + err = fmt.Errorf("negative file descriptor: %d", fileDescriptor) + } + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: fmt.Errorf("parse inherited file descriptor: %w", err), + } + } + + inheritedFile := os.NewFile(uintptr(fileDescriptor), environmentVariable) + if inheritedFile == nil { + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: errors.New("open inherited file descriptor"), + } + } + + listener, listenerErr := net.FileListener(inheritedFile) + closeErr := inheritedFile.Close() + if listenerErr != nil { + listenerErr = fmt.Errorf("create listener from inherited file descriptor: %w", listenerErr) + if closeErr != nil { + listenerErr = errors.Join(listenerErr, fmt.Errorf("close inherited file descriptor wrapper: %w", closeErr)) + } + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: errors.Join(errDashboardListenerTCP, listenerErr), + } + } + if closeErr != nil { + listenerCloseErr := listener.Close() + err = fmt.Errorf("close inherited file descriptor wrapper: %w", closeErr) + if listenerCloseErr != nil { + err = errors.Join(err, fmt.Errorf("close adopted listener: %w", listenerCloseErr)) + } + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: err, + } + } + + tcpListener, ok := listener.(*net.TCPListener) + if !ok { + if closeErr := listener.Close(); closeErr != nil { + err = errors.Join(errDashboardListenerTCP, fmt.Errorf("close rejected listener: %w", closeErr)) + } else { + err = errDashboardListenerTCP + } + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: err, + } + } + + tcpAddress, ok := tcpListener.Addr().(*net.TCPAddr) + if !ok { + if closeErr := tcpListener.Close(); closeErr != nil { + err = errors.Join(errDashboardListenerTCP, fmt.Errorf("close rejected TCP listener: %w", closeErr)) + } else { + err = errDashboardListenerTCP + } + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: err, + } + } + if !tcpAddress.IP.IsLoopback() { + if closeErr := tcpListener.Close(); closeErr != nil { + err = errors.Join(errDashboardListenerLoopback, fmt.Errorf("close rejected listener: %w", closeErr)) + } else { + err = errDashboardListenerLoopback + } + return nil, &dashboardListenerError{ + Kind: kind, + EnvironmentVariable: environmentVariable, + Descriptor: descriptor, + Err: err, + } + } + return tcpListener, nil +} + +func dashboardListenerEnvironmentVariable(kind dashboardListenerKind) (string, error) { + switch kind { + case dashboardHTTPListener: + return dashboardHTTPListenerFDEnv, nil + case dashboardHTTPSListener: + return dashboardHTTPSListenerFDEnv, nil + case dashboardReceiptListener: + return dashboardReceiptListenerFDEnv, nil + default: + return "", &dashboardListenerError{Kind: kind, Err: errDashboardListenerKind} + } +} + +func serveDashboardHTTPS(server *http.Server, certificatePath, keyPath string) (err error) { + listener, err := openDashboardListener("tcp", server.Addr, dashboardHTTPSListener) + if err != nil { + return err + } + defer func() { + if closeErr := listener.Close(); closeErr != nil && !errors.Is(closeErr, net.ErrClosed) { + err = errors.Join(err, fmt.Errorf("close HTTPS dashboard listener: %w", closeErr)) + } + }() + return server.ServeTLS(listener, certificatePath, keyPath) +} diff --git a/cmd/dashboard/dashboard_listener_integration_test.go b/cmd/dashboard/dashboard_listener_integration_test.go new file mode 100644 index 00000000..b3d1c5f0 --- /dev/null +++ b/cmd/dashboard/dashboard_listener_integration_test.go @@ -0,0 +1,236 @@ +//go:build agentcompat + +package main + +import ( + "errors" + "fmt" + "net" + "os" + "os/exec" + "strconv" + "syscall" + "testing" +) + +const ( + dashboardListenerHelperModeEnv = "NEZHA_AGENTCOMPAT_LISTENER_HELPER" + dashboardListenerHelperKindEnv = "NEZHA_AGENTCOMPAT_LISTENER_HELPER_KIND" +) + +func TestDashboardListener_UsesDefaultListen(t *testing.T) { + // Given + t.Setenv(dashboardHTTPListenerFDEnv, "") + t.Setenv(dashboardHTTPSListenerFDEnv, "") + + // When + listener, err := openDashboardListener("tcp", "127.0.0.1:0", dashboardHTTPListener) + + // Then + if err != nil { + t.Fatalf("open default dashboard listener: %v", err) + } + t.Cleanup(func() { + if err := listener.Close(); err != nil { + t.Errorf("close default dashboard listener: %v", err) + } + }) + address := listener.Addr().(*net.TCPAddr) + if !address.IP.IsLoopback() { + t.Fatalf("default listener address = %s, want loopback", address) + } +} + +func TestDashboardListener_AdoptsInheritedLoopback(t *testing.T) { + tests := []struct { + name string + environmentVariable string + kind dashboardListenerKind + }{ + {name: "HTTP", environmentVariable: dashboardHTTPListenerFDEnv, kind: dashboardHTTPListener}, + {name: "HTTPS", environmentVariable: dashboardHTTPSListenerFDEnv, kind: dashboardHTTPSListener}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + // Given + inherited := newDashboardTCPListener(t, "127.0.0.1:0") + inheritedFile, err := inherited.File() + if err != nil { + t.Fatalf("duplicate inherited listener for helper: %v", err) + } + t.Cleanup(func() { + if err := inheritedFile.Close(); err != nil { + t.Errorf("close helper listener file: %v", err) + } + }) + command := exec.Command(os.Args[0], "-test.run=^TestDashboardListenerAdoptionHelper$") + command.ExtraFiles = []*os.File{inheritedFile} + command.Env = append(os.Environ(), + dashboardListenerHelperModeEnv+"=1", + dashboardListenerHelperKindEnv+"="+string(test.kind), + test.environmentVariable+"=3", + ) + + // When + output, err := command.Output() + + // Then + if err != nil { + t.Fatalf("run inherited listener helper: %v", err) + } + var adoptedAddress string + var adoptedInode uint64 + if _, err := fmt.Sscanf(string(output), "%s %d", &adoptedAddress, &adoptedInode); err != nil { + t.Fatalf("parse helper output %q: %v", output, err) + } + if want := inherited.Addr().String(); adoptedAddress != want { + t.Fatalf("adopted listener address = %q, want %q", adoptedAddress, want) + } + if want := dashboardTCPListenerInode(t, inherited); adoptedInode != want { + t.Fatalf("adopted listener inode = %d, want %d", adoptedInode, want) + } + t.Logf("adopted address=%s inode=%d", adoptedAddress, adoptedInode) + }) + } +} + +func TestDashboardListenerAdoptionHelper(t *testing.T) { + if os.Getenv(dashboardListenerHelperModeEnv) != "1" { + return + } + kind := dashboardListenerKind(os.Getenv(dashboardListenerHelperKindEnv)) + listener, err := openDashboardListener("tcp", "127.0.0.1:1", kind) + if err != nil { + t.Fatalf("adopt inherited dashboard listener: %v", err) + } + t.Cleanup(func() { + if err := listener.Close(); err != nil { + t.Errorf("close adopted dashboard listener: %v", err) + } + }) + tcpListener, ok := listener.(*net.TCPListener) + if !ok { + t.Fatalf("adopted listener type = %T, want *net.TCPListener", listener) + } + fmt.Printf("%s %d\n", listener.Addr(), dashboardTCPListenerInode(t, tcpListener)) +} + +func TestDashboardListener_RejectsPipeFD(t *testing.T) { + // Given + reader, writer, err := os.Pipe() + if err != nil { + t.Fatalf("create pipe fixture: %v", err) + } + t.Cleanup(func() { + if err := reader.Close(); err != nil && !errors.Is(err, os.ErrClosed) && !errors.Is(err, syscall.EBADF) { + t.Errorf("close pipe reader: %v", err) + } + if err := writer.Close(); err != nil { + t.Errorf("close pipe writer: %v", err) + } + }) + t.Setenv(dashboardHTTPListenerFDEnv, strconv.FormatUint(uint64(reader.Fd()), 10)) + + // When + listener, err := openDashboardListener("tcp", "127.0.0.1:0", dashboardHTTPListener) + + // Then + if listener != nil { + t.Cleanup(func() { _ = listener.Close() }) + t.Fatal("pipe FD produced a dashboard listener") + } + requireDashboardListenerError(t, err, errDashboardListenerTCP) +} + +func TestDashboardListener_RejectsNonLoopback(t *testing.T) { + // Given + inherited := newDashboardTCPListener(t, "0.0.0.0:0") + setDashboardInheritedListenerFD(t, dashboardHTTPListenerFDEnv, inherited) + + // When + listener, err := openDashboardListener("tcp", "127.0.0.1:0", dashboardHTTPListener) + + // Then + if listener != nil { + t.Cleanup(func() { _ = listener.Close() }) + t.Fatal("non-loopback FD produced a dashboard listener") + } + requireDashboardListenerError(t, err, errDashboardListenerLoopback) +} + +func newDashboardTCPListener(t *testing.T, address string) *net.TCPListener { + t.Helper() + listener, err := net.Listen("tcp", address) + if err != nil { + t.Fatalf("listen on %s: %v", address, err) + } + tcpListener, ok := listener.(*net.TCPListener) + if !ok { + _ = listener.Close() + t.Fatalf("listener type = %T, want *net.TCPListener", listener) + } + t.Cleanup(func() { + if err := tcpListener.Close(); err != nil { + t.Errorf("close inherited listener fixture: %v", err) + } + }) + return tcpListener +} + +func setDashboardInheritedListenerFD(t *testing.T, environmentVariable string, listener *net.TCPListener) { + t.Helper() + rawConnection, err := listener.SyscallConn() + if err != nil { + t.Fatalf("get listener syscall connection: %v", err) + } + var inheritedFD int + var duplicateError error + if err := rawConnection.Control(func(fd uintptr) { + inheritedFD, duplicateError = syscall.Dup(int(fd)) + }); err != nil { + t.Fatalf("access listener FD: %v", err) + } + if duplicateError != nil { + t.Fatalf("duplicate listener FD: %v", duplicateError) + } + t.Cleanup(func() { + if err := syscall.Close(inheritedFD); err != nil && !errors.Is(err, syscall.EBADF) { + t.Errorf("close duplicated listener FD: %v", err) + } + }) + t.Setenv(environmentVariable, strconv.Itoa(inheritedFD)) +} + +func dashboardTCPListenerInode(t *testing.T, listener *net.TCPListener) uint64 { + t.Helper() + rawConnection, err := listener.SyscallConn() + if err != nil { + t.Fatalf("get listener syscall connection: %v", err) + } + var stat syscall.Stat_t + var statError error + if err := rawConnection.Control(func(fd uintptr) { + statError = syscall.Fstat(int(fd), &stat) + }); err != nil { + t.Fatalf("access listener FD: %v", err) + } + if statError != nil { + t.Fatalf("stat listener FD: %v", statError) + } + return stat.Ino +} + +func requireDashboardListenerError(t *testing.T, err, cause error) { + t.Helper() + if err == nil { + t.Fatal("openDashboardListener error = nil, want typed error") + } + if !errors.Is(err, cause) { + t.Fatalf("openDashboardListener error = %v, want cause %v", err, cause) + } + var listenerError *dashboardListenerError + if !errors.As(err, &listenerError) { + t.Fatalf("openDashboardListener error type = %T, want *dashboardListenerError", err) + } +} diff --git a/cmd/dashboard/dashboard_listener_test.go b/cmd/dashboard/dashboard_listener_test.go new file mode 100644 index 00000000..0b4af805 --- /dev/null +++ b/cmd/dashboard/dashboard_listener_test.go @@ -0,0 +1,26 @@ +package main + +import "testing" + +func TestDashboardListenerAddress_PreservesCurrentSemantics(t *testing.T) { + tests := []struct { + name string + host string + port uint16 + want string + }{ + {name: "empty host", host: "", port: 8008, want: ":8008"}, + {name: "loopback IPv4", host: "127.0.0.1", port: 8008, want: "127.0.0.1:8008"}, + {name: "loopback IPv6", host: "::1", port: 8008, want: "::1:8008"}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + got := dashboardListenerAddress(test.host, test.port) + + if got != test.want { + t.Fatalf("dashboardListenerAddress(%q, %d) = %q, want %q", test.host, test.port, got, test.want) + } + }) + } +} diff --git a/cmd/dashboard/main.go b/cmd/dashboard/main.go index 3f389161..928222ae 100644 --- a/cmd/dashboard/main.go +++ b/cmd/dashboard/main.go @@ -8,7 +8,6 @@ import ( "flag" "fmt" "log" - "net" "net/http" "os" "runtime/debug" @@ -140,10 +139,18 @@ func main() { log.Fatal(err) } - l, err := net.Listen("tcp", fmt.Sprintf("%s:%d", singleton.Conf.ListenHost, singleton.Conf.ListenPort)) + l, err := openDashboardListener("tcp", dashboardListenerAddress(singleton.Conf.ListenHost, singleton.Conf.ListenPort), dashboardHTTPListener) if err != nil { log.Fatal(err) } + receiptListener, err := openReceiptGateListener() + if err != nil { + log.Fatal(err) + } + if receiptListener != nil { + defer receiptListener.Close() + rpc.SetReceiptGateListener(receiptListener) + } singleton.CleanMonitorHistory() rpc.DispatchKeepalive() @@ -185,7 +192,7 @@ func main() { log.Printf("NEZHA>> Dashboard::START ON %s:%d", singleton.Conf.ListenHost, singleton.Conf.ListenPort) if singleton.Conf.HTTPS.ListenPort != 0 { go func() { - errChan <- muxServerHTTPS.ListenAndServeTLS(singleton.Conf.HTTPS.TLSCertPath, singleton.Conf.HTTPS.TLSKeyPath) + errChan <- serveDashboardHTTPS(muxServerHTTPS, singleton.Conf.HTTPS.TLSCertPath, singleton.Conf.HTTPS.TLSKeyPath) }() log.Printf("NEZHA>> Dashboard::START ON %s:%d", singleton.Conf.ListenHost, singleton.Conf.HTTPS.ListenPort) } @@ -195,6 +202,7 @@ func main() { return <-errChan }, func(c context.Context) error { log.Println("NEZHA>> Graceful::START") + rpc.CloseReceiptGate() singleton.RecordTransferHourlyUsage() singleton.CloseTSDB() log.Println("NEZHA>> Graceful::END")