mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
feat(agentcompat): add dashboard listener hooks
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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")
|
||||
|
||||
Reference in New Issue
Block a user