feat(agentcompat): add dashboard listener hooks

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:33:38 +00:00
co-authored by naiba/CloudCode
parent 3f6de265a0
commit aee586baad
6 changed files with 484 additions and 3 deletions
+15
View File
@@ -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)
}
@@ -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)
}
@@ -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)
}
@@ -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)
}
}
+26
View File
@@ -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)
}
})
}
}
+11 -3
View File
@@ -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")