mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 17:50:12 +00:00
feat(agentcompat): expose dashboard capability routes
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -0,0 +1,94 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/netip"
|
||||
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/peer"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/utils"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
func ctxWithRealIP(ctx context.Context) (context.Context, error) {
|
||||
var ip, connectingIp string
|
||||
p, ok := peer.FromContext(ctx)
|
||||
if ok {
|
||||
addrPort, err := netip.ParseAddrPort(p.Addr.String())
|
||||
if err == nil {
|
||||
connectingIp = addrPort.Addr().String()
|
||||
}
|
||||
}
|
||||
ctx = context.WithValue(ctx, model.CtxKeyConnectingIP{}, connectingIp)
|
||||
|
||||
if singleton.Conf.AgentRealIPHeader == "" {
|
||||
return ctx, nil
|
||||
}
|
||||
|
||||
if singleton.Conf.AgentRealIPHeader == model.ConfigUsePeerIP {
|
||||
if connectingIp == "" {
|
||||
return ctx, fmt.Errorf("connecting ip not found")
|
||||
}
|
||||
ip = connectingIp
|
||||
} else {
|
||||
vals := metadata.ValueFromIncomingContext(ctx, singleton.Conf.AgentRealIPHeader)
|
||||
if len(vals) == 0 {
|
||||
return ctx, fmt.Errorf("real ip header not found")
|
||||
}
|
||||
var err error
|
||||
ip, err = utils.GetIPFromHeader(vals[0])
|
||||
if err != nil {
|
||||
return ctx, err
|
||||
}
|
||||
}
|
||||
|
||||
if singleton.Conf.Debug {
|
||||
log.Printf("NEZHA>> gRPC Agent Real IP: %s, connecting IP: %s\n", ip, connectingIp)
|
||||
}
|
||||
|
||||
return context.WithValue(ctx, model.CtxKeyRealIP{}, ip), nil
|
||||
}
|
||||
|
||||
func waf(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
realip, _ := ctx.Value(model.CtxKeyRealIP{}).(string)
|
||||
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return handler(ctx, req)
|
||||
}
|
||||
|
||||
func getRealIp(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
ctx, err := ctxWithRealIP(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return handler(ctx, req)
|
||||
}
|
||||
|
||||
type realIPServerStream struct {
|
||||
grpc.ServerStream
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func (s *realIPServerStream) Context() context.Context { return s.ctx }
|
||||
|
||||
func getRealIpStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||
ctx, err := ctxWithRealIP(ss.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return handler(srv, &realIPServerStream{ServerStream: ss, ctx: ctx})
|
||||
}
|
||||
|
||||
func wafStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||
realip, _ := ss.Context().Value(model.CtxKeyRealIP{}).(string)
|
||||
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
||||
return err
|
||||
}
|
||||
return handler(srv, ss)
|
||||
}
|
||||
@@ -0,0 +1,115 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/goccy/go-json"
|
||||
|
||||
"github.com/hashicorp/go-uuid"
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/utils"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) {
|
||||
capabilityLease, capabilityErr := prepareNATCapability(r, natConfig)
|
||||
if capabilityErr != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write([]byte("NAT capability unavailable"))
|
||||
return
|
||||
}
|
||||
handler := serviceRPC.NezhaHandlerSingleton
|
||||
streamId := ""
|
||||
legacyStreamOwned := false
|
||||
relayCompleted := false
|
||||
defer func() {
|
||||
if capabilityLease.active {
|
||||
if !relayCompleted {
|
||||
capabilityLease.cleanup(handler)
|
||||
}
|
||||
return
|
||||
}
|
||||
if legacyStreamOwned {
|
||||
_ = handler.CloseStream(streamId)
|
||||
}
|
||||
}()
|
||||
server, _ := singleton.ServerShared.Get(natConfig.ServerID)
|
||||
if server == nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write([]byte("server not found or not connected"))
|
||||
return
|
||||
}
|
||||
if server.GetTaskStream() == nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write([]byte("server not found or not connected"))
|
||||
return
|
||||
}
|
||||
|
||||
streamId, err := uuid.GenerateUUID()
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "stream id error: %v", err))
|
||||
return
|
||||
}
|
||||
if capabilityLease.active {
|
||||
capabilityLease.streamLease, err = handler.CreateAgentCompatNATStream(capabilityLease.handle, streamId)
|
||||
} else {
|
||||
err = handler.CreateStream(streamId, 0, server.ID)
|
||||
}
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
w.Write(fmt.Appendf(nil, "stream limit: %v", err))
|
||||
return
|
||||
}
|
||||
legacyStreamOwned = !capabilityLease.active
|
||||
taskData, err := json.Marshal(model.TaskNAT{StreamID: streamId, Host: natConfig.Host})
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "task data error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
if err := server.SendTask(&proto.Task{Type: model.TaskTypeNAT, Data: string(taskData)}); err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "send task error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
// Authorization authenticates Dashboard access before NAT ingress; it must not become an origin credential for the configured NAT backend.
|
||||
wWrapped, err := utils.NewRequestWrapper(r, w)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := handler.UserConnected(streamId, wWrapped); err != nil {
|
||||
if closeErr := wWrapped.Close(); closeErr != nil {
|
||||
log.Printf("NEZHA>> NAT request wrapper close error after user connection failure: %v", closeErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
if capabilityLease.active {
|
||||
if err := handler.PublishAgentCompatNATStream(capabilityLease.handle, serviceRPC.AgentCompatNATPublication{
|
||||
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||
TargetServerID: server.ID,
|
||||
ResourceID: natConfig.ID,
|
||||
StreamID: streamId,
|
||||
}); err != nil {
|
||||
return
|
||||
}
|
||||
}
|
||||
if capabilityLease.active {
|
||||
capabilityLease.publicationOwned, err = handler.StartAgentCompatNATStream(capabilityLease.handle, time.Second*10)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
} else if err := handler.StartStream(streamId, time.Second*10); err != nil {
|
||||
return
|
||||
}
|
||||
relayCompleted = true
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
//go:build agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||
)
|
||||
|
||||
type natCapabilityLease struct {
|
||||
active bool
|
||||
publicationOwned bool
|
||||
streamLease *serviceRPC.AgentCompatNATStreamLease
|
||||
access serviceRPC.AgentCompatCapabilityAccess
|
||||
handle serviceRPC.AgentCompatNATPublishHandle
|
||||
}
|
||||
|
||||
func prepareNATCapability(request *http.Request, natConfig *model.NAT) (natCapabilityLease, error) {
|
||||
values := request.Header.Values(agentcompatcontract.IOStreamCapabilityHeader)
|
||||
if len(values) == 0 {
|
||||
request.Header.Del("Authorization")
|
||||
return natCapabilityLease{}, nil
|
||||
}
|
||||
request.Header.Del(agentcompatcontract.IOStreamCapabilityHeader)
|
||||
if len(values) != 1 || values[0] == "" {
|
||||
request.Header.Del("Authorization")
|
||||
return natCapabilityLease{}, errors.New("invalid NAT capability")
|
||||
}
|
||||
access, handle, err := serviceRPC.NezhaHandlerSingleton.ConsumeAgentCompatNATCapabilityForProfile(values[0], natConfig.ServerID, natConfig.ID)
|
||||
request.Header.Del("Authorization")
|
||||
if err != nil {
|
||||
return natCapabilityLease{}, errors.New("invalid NAT capability")
|
||||
}
|
||||
return natCapabilityLease{active: true, access: access, handle: handle}, nil
|
||||
}
|
||||
|
||||
func (lease natCapabilityLease) cleanup(handler *serviceRPC.NezhaHandler) {
|
||||
if !lease.active {
|
||||
return
|
||||
}
|
||||
_ = handler.CancelAgentCompatIOStreamCapability(lease.access)
|
||||
_ = handler.CloseAgentCompatNATStreamLease(lease.streamLease)
|
||||
_ = handler.UnregisterAgentCompatIOStreamCapability(lease.access)
|
||||
}
|
||||
@@ -0,0 +1,82 @@
|
||||
//go:build agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"log"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||
)
|
||||
|
||||
func TestPrepareNATCapabilityConsumesAndRemovesHeaderBeforeNATWork(t *testing.T) {
|
||||
// Given
|
||||
handler := serviceRPC.NewNezhaHandler()
|
||||
original := serviceRPC.NezhaHandlerSingleton
|
||||
serviceRPC.NezhaHandlerSingleton = handler
|
||||
t.Cleanup(func() { serviceRPC.NezhaHandlerSingleton = original })
|
||||
request := &http.Request{Header: make(http.Header)}
|
||||
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "malformed")
|
||||
|
||||
// When
|
||||
lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81})
|
||||
|
||||
// Then
|
||||
if err == nil {
|
||||
t.Fatal("malformed capability unexpectedly accepted")
|
||||
}
|
||||
if lease.active {
|
||||
t.Fatal("malformed capability unexpectedly activated")
|
||||
}
|
||||
if _, present := request.Header[agentcompatcontract.IOStreamCapabilityHeader]; present {
|
||||
t.Fatal("capability header remained after hook")
|
||||
}
|
||||
}
|
||||
|
||||
func TestPrepareNATCapabilityRejectsDuplicateHeaderAfterRemovingAllValues(t *testing.T) {
|
||||
// Given
|
||||
request := &http.Request{Header: make(http.Header)}
|
||||
request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, "one")
|
||||
request.Header.Add(agentcompatcontract.IOStreamCapabilityHeader, "two")
|
||||
|
||||
// When
|
||||
lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 92}, ServerID: 82})
|
||||
|
||||
// Then
|
||||
if err == nil {
|
||||
t.Fatal("duplicate capability header unexpectedly accepted")
|
||||
}
|
||||
if lease.active {
|
||||
t.Fatal("duplicate capability header unexpectedly activated")
|
||||
}
|
||||
if _, present := request.Header[agentcompatcontract.IOStreamCapabilityHeader]; present {
|
||||
t.Fatal("duplicate capability header remained after hook")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatSensitiveHeadersStayOutOfErrorsAndLogs(t *testing.T) {
|
||||
request := &http.Request{Header: make(http.Header)}
|
||||
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "capability-secret")
|
||||
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||
writer := &serveNATResponseWriter{}
|
||||
var logs bytes.Buffer
|
||||
originalOutput := log.Writer()
|
||||
log.SetOutput(&logs)
|
||||
t.Cleanup(func() { log.SetOutput(originalOutput) })
|
||||
|
||||
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81})
|
||||
|
||||
if writer.status != http.StatusServiceUnavailable {
|
||||
t.Fatalf("status = %d, want %d", writer.status, http.StatusServiceUnavailable)
|
||||
}
|
||||
for _, output := range []string{writer.body, logs.String()} {
|
||||
if strings.Contains(output, "capability-secret") || strings.Contains(output, "deterministic-secret") {
|
||||
t.Fatalf("sensitive value leaked in %q", output)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
//go:build !agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||
)
|
||||
|
||||
type natCapabilityLease struct {
|
||||
active bool
|
||||
publicationOwned bool
|
||||
streamLease *serviceRPC.AgentCompatNATStreamLease
|
||||
access serviceRPC.AgentCompatCapabilityAccess
|
||||
handle serviceRPC.AgentCompatNATPublishHandle
|
||||
}
|
||||
|
||||
func prepareNATCapability(request *http.Request, _ *model.NAT) (natCapabilityLease, error) {
|
||||
request.Header.Del("Authorization")
|
||||
return natCapabilityLease{}, nil
|
||||
}
|
||||
|
||||
func (lease natCapabilityLease) cleanup(*serviceRPC.NezhaHandler) {}
|
||||
@@ -0,0 +1,64 @@
|
||||
//go:build !agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/goccy/go-json"
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
)
|
||||
|
||||
func TestServeNATDefaultForwardsCapabilityHeaderAsOrdinaryData(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
connection := newServeNATConn()
|
||||
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||
request := &http.Request{
|
||||
Method: http.MethodPost,
|
||||
URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"},
|
||||
Header: make(http.Header),
|
||||
Body: &serveNATBody{},
|
||||
}
|
||||
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "ordinary-data")
|
||||
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||
var nat model.TaskNAT
|
||||
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||
return err
|
||||
}
|
||||
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||
}
|
||||
done := make(chan struct{})
|
||||
deadline, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
go func() {
|
||||
ServeNAT(&serveNATResponseWriter{conn: connection}, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"})
|
||||
close(done)
|
||||
}()
|
||||
select {
|
||||
case <-agent.writeDone:
|
||||
case <-deadline.Done():
|
||||
t.Fatal("agent did not receive ordinary NAT request")
|
||||
}
|
||||
_ = connection.Close()
|
||||
select {
|
||||
case <-done:
|
||||
case <-deadline.Done():
|
||||
t.Fatal("ServeNAT did not finish")
|
||||
}
|
||||
forwarded := strings.ToLower(string(agent.writtenBytes()))
|
||||
if !strings.Contains(forwarded, strings.ToLower(agentcompatcontract.IOStreamCapabilityHeader)+": ordinary-data") {
|
||||
t.Fatal("default NAT flow did not forward capability header as ordinary request data")
|
||||
}
|
||||
if strings.Contains(forwarded, "Authorization:") || strings.Contains(forwarded, "deterministic-secret") {
|
||||
t.Fatal("default NAT flow forwarded Authorization credentials")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
//go:build !agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
)
|
||||
|
||||
func TestPrepareNATCapabilityDefaultPreservesHeaderAsOrdinaryRequestData(t *testing.T) {
|
||||
// Given
|
||||
request := &http.Request{Header: make(http.Header)}
|
||||
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, "ordinary-data")
|
||||
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||
|
||||
// When
|
||||
lease, err := prepareNATCapability(request, &model.NAT{Common: model.Common{ID: 91}, ServerID: 81})
|
||||
|
||||
// Then
|
||||
if err != nil {
|
||||
t.Fatalf("default NAT capability hook returned error: %v", err)
|
||||
}
|
||||
if lease.active {
|
||||
t.Fatal("default NAT capability hook unexpectedly activated")
|
||||
}
|
||||
if got := request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader); got != "ordinary-data" {
|
||||
t.Fatalf("default NAT capability hook changed header to %q", got)
|
||||
}
|
||||
if got := request.Header.Get("Authorization"); got != "" {
|
||||
t.Fatalf("default NAT capability hook retained Authorization %q", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,245 @@
|
||||
//go:build agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/goccy/go-json"
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestServeNATAgentCompatReleasesCapabilityBeforeServerDispatch(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 101)
|
||||
originalServers := singleton.ServerShared
|
||||
singleton.ServerShared = singleton.NewServerClass()
|
||||
t.Cleanup(func() { singleton.ServerShared = originalServers })
|
||||
request := capabilityRequest(capability)
|
||||
|
||||
ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 101}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
assertNATCapabilitySlotReleased(t, fixture.handler, access)
|
||||
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||
require.Empty(t, fixture.taskStream.sent)
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatTaskStreamUnavailableReleasesCapabilityBeforeActivity(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 105)
|
||||
fixture.server.SetTaskStream(nil)
|
||||
request := capabilityRequest(capability)
|
||||
|
||||
ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 105}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
assertNATCapabilitySlotReleased(t, fixture.handler, access)
|
||||
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||
require.Empty(t, fixture.taskStream.sent)
|
||||
if request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader) != "" {
|
||||
t.Fatal("capability header remained after task stream unavailable")
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatCreateQuotaFailurePreservesExistingStreams(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 106)
|
||||
const existingStreamCount = 40
|
||||
for index := 0; index < existingStreamCount; index++ {
|
||||
require.NoError(t, fixture.handler.CreateStreamWithPurpose("quota-existing-"+string(rune(index)), 0, fixture.server.ID, serviceRPC.PurposeNAT))
|
||||
}
|
||||
t.Cleanup(func() {
|
||||
for index := 0; index < existingStreamCount; index++ {
|
||||
_ = fixture.handler.CloseStream("quota-existing-" + string(rune(index)))
|
||||
}
|
||||
})
|
||||
request := capabilityRequest(capability)
|
||||
|
||||
ServeNAT(&serveNATResponseWriter{}, request, &model.NAT{Common: model.Common{ID: 106}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
assertNATCapabilitySlotReleased(t, fixture.handler, access)
|
||||
require.Equal(t, existingStreamCount, fixture.handler.StreamCount())
|
||||
require.Empty(t, fixture.taskStream.sent)
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatPublishCancellationClosesRegistryOwnedEndpoints(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 107)
|
||||
connection := newServeNATConn()
|
||||
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||
observerEntered := make(chan struct{})
|
||||
observerRelease := make(chan struct{})
|
||||
var observerOnce sync.Once
|
||||
fixture.handler.SetAgentCompatCapabilityPublishObserverForTest(func() {
|
||||
observerOnce.Do(func() { close(observerEntered) })
|
||||
<-observerRelease
|
||||
})
|
||||
t.Cleanup(func() { fixture.handler.SetAgentCompatCapabilityPublishObserverForTest(nil) })
|
||||
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||
var nat model.TaskNAT
|
||||
require.NoError(t, json.Unmarshal([]byte(task.Data), &nat))
|
||||
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||
}
|
||||
request := capabilityRequest(capability)
|
||||
writer := &serveNATResponseWriter{conn: connection}
|
||||
serveDone := make(chan struct{})
|
||||
go func() {
|
||||
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 107}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
close(serveDone)
|
||||
}()
|
||||
select {
|
||||
case <-observerEntered:
|
||||
case <-serveDone:
|
||||
t.Fatal("ServeNAT returned before publish pause")
|
||||
}
|
||||
fixture.handler.CreateStreamWithPurpose("publish-replacement", 0, fixture.server.ID, serviceRPC.PurposeNAT)
|
||||
require.NoError(t, fixture.handler.CancelAgentCompatIOStreamCapability(access))
|
||||
close(observerRelease)
|
||||
completionContext, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
select {
|
||||
case <-serveDone:
|
||||
case <-completionContext.Done():
|
||||
t.Fatal("ServeNAT did not finish after publish cancellation")
|
||||
}
|
||||
assertCapabilityReleased(t, fixture.handler, access)
|
||||
require.Equal(t, int32(1), connection.closeCount.Load())
|
||||
require.Equal(t, int32(1), agent.closeCount.Load())
|
||||
require.Equal(t, 1, fixture.handler.StreamCount())
|
||||
select {
|
||||
case <-agent.writeDone:
|
||||
t.Fatal("agent endpoint was used after publish cancellation")
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatSendTaskFailureReleasesStreamAndCapabilityBeforeWrapper(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 102)
|
||||
fixture.taskStream.sendErr = io.ErrClosedPipe
|
||||
connection := newServeNATConn()
|
||||
request := capabilityRequest(capability)
|
||||
request.Body = &serveNATBody{}
|
||||
writer := &serveNATResponseWriter{conn: connection}
|
||||
|
||||
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 102}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
assertCapabilityReleased(t, fixture.handler, access)
|
||||
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||
require.Equal(t, int32(0), connection.closeCount.Load())
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatWrapperFailureReleasesCapabilityWithoutSecondResponse(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 103)
|
||||
request := capabilityRequest(capability)
|
||||
writer := &nonHijackingResponseWriter{}
|
||||
|
||||
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 103}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
assertCapabilityReleased(t, fixture.handler, access)
|
||||
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||
require.Equal(t, 0, writer.writes)
|
||||
}
|
||||
|
||||
func TestServeNATAgentCompatStartFailureCancelsPublishedCapabilityAndClosesEndpoints(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, access := registerNATCapabilityForTest(t, fixture.handler, fixture.server.ID, 104)
|
||||
connection := newServeNATConn()
|
||||
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||
var nat model.TaskNAT
|
||||
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||
return err
|
||||
}
|
||||
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||
}
|
||||
request := capabilityRequest(capability)
|
||||
writer := &serveNATResponseWriter{conn: connection}
|
||||
|
||||
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 104}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
assertCapabilityReleased(t, fixture.handler, access)
|
||||
require.Equal(t, 0, fixture.handler.StreamCount())
|
||||
require.Equal(t, int32(1), connection.closeCount.Load())
|
||||
require.Equal(t, int32(1), agent.closeCount.Load())
|
||||
}
|
||||
|
||||
func registerNATCapabilityForTest(t *testing.T, handler *serviceRPC.NezhaHandler, serverID, resourceID uint64) (string, serviceRPC.AgentCompatCapabilityAccess) {
|
||||
t.Helper()
|
||||
capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: resourceID, UserID: resourceID + 1}, Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||
TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
parsed, err := serviceRPC.ParseAgentCompatIOStreamCapability(capability.String())
|
||||
require.NoError(t, err)
|
||||
return capability.String(), serviceRPC.AgentCompatCapabilityAccess{
|
||||
Capability: parsed, Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: resourceID, UserID: resourceID + 1},
|
||||
Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: serverID, ResourceID: resourceID, ServerAccessAllowed: true,
|
||||
}
|
||||
}
|
||||
|
||||
func capabilityRequest(value string) *http.Request {
|
||||
request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: &serveNATBody{}}
|
||||
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, value)
|
||||
return request
|
||||
}
|
||||
|
||||
func assertCapabilityReleased(t *testing.T, handler *serviceRPC.NezhaHandler, access serviceRPC.AgentCompatCapabilityAccess) {
|
||||
t.Helper()
|
||||
_, _, err := handler.ConsumeAgentCompatNATCapabilityForProfile(access.Capability.String(), access.TargetServerID, access.ResourceID)
|
||||
require.ErrorIs(t, err, serviceRPC.ErrAgentCompatCapabilityHidden)
|
||||
}
|
||||
|
||||
func assertNATCapabilitySlotReleased(t *testing.T, handler *serviceRPC.NezhaHandler, access serviceRPC.AgentCompatCapabilityAccess) {
|
||||
t.Helper()
|
||||
assertCapabilityReleased(t, handler, access)
|
||||
type activeCapability struct {
|
||||
capability serviceRPC.AgentCompatIOStreamCapability
|
||||
resourceID uint64
|
||||
}
|
||||
active := make([]activeCapability, 0, 16)
|
||||
for index := uint64(0); index < 16; index++ {
|
||||
capability, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||
Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: access.TargetServerID,
|
||||
ResourceID: access.ResourceID + index + 1000, ServerAccessAllowed: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
active = append(active, activeCapability{capability: capability, resourceID: access.ResourceID + index + 1000})
|
||||
}
|
||||
_, err := handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||
Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT, TargetServerID: access.TargetServerID,
|
||||
ResourceID: access.ResourceID + 2000, ServerAccessAllowed: true,
|
||||
})
|
||||
require.ErrorIs(t, err, serviceRPC.ErrAgentCompatCapabilityUnavailable)
|
||||
for _, item := range active {
|
||||
parsed, parseErr := serviceRPC.ParseAgentCompatIOStreamCapability(item.capability.String())
|
||||
require.NoError(t, parseErr)
|
||||
require.NoError(t, handler.CancelAgentCompatIOStreamCapability(serviceRPC.AgentCompatCapabilityAccess{
|
||||
Capability: parsed, Owner: access.Owner, Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||
TargetServerID: access.TargetServerID, ResourceID: item.resourceID, ServerAccessAllowed: true,
|
||||
}))
|
||||
}
|
||||
}
|
||||
|
||||
type nonHijackingResponseWriter struct {
|
||||
writes int
|
||||
}
|
||||
|
||||
func (writer *nonHijackingResponseWriter) Header() http.Header { return make(http.Header) }
|
||||
func (writer *nonHijackingResponseWriter) Write(data []byte) (int, error) {
|
||||
writer.writes += len(data)
|
||||
return len(data), nil
|
||||
}
|
||||
func (*nonHijackingResponseWriter) WriteHeader(int) {}
|
||||
@@ -0,0 +1,151 @@
|
||||
//go:build agentcompat
|
||||
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/goccy/go-json"
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
serviceRPC "github.com/nezhahq/nezha/service/rpc"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestServeNATAgentCompatPublishesExactStreamAfterRequestTransfer(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
capability, err := fixture.handler.RegisterAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityRegistration{
|
||||
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2},
|
||||
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||
TargetServerID: fixture.server.ID,
|
||||
ResourceID: 91,
|
||||
ServerAccessAllowed: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
access, err := serviceRPC.ParseAgentCompatIOStreamCapability(capability.String())
|
||||
require.NoError(t, err)
|
||||
connection := newServeNATConn()
|
||||
writer := &serveNATResponseWriter{conn: connection}
|
||||
agent := &orderedNATAgent{readStarted: make(chan struct{}), readRelease: make(chan struct{}), writeStarted: make(chan struct{})}
|
||||
taskStreamIDReady := make(chan string, 1)
|
||||
request := &http.Request{
|
||||
Method: http.MethodPost,
|
||||
URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"},
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader("payload")),
|
||||
}
|
||||
request.Header.Set(agentcompatcontract.IOStreamCapabilityHeader, capability.String())
|
||||
request.Header.Set("Authorization", "Bearer deterministic-secret")
|
||||
request.Header.Set("X-Ordinary-NAT", "ordinary")
|
||||
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||
var nat model.TaskNAT
|
||||
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := fixture.handler.AgentConnected(nat.StreamID, agent); err != nil {
|
||||
return err
|
||||
}
|
||||
taskStreamIDReady <- nat.StreamID
|
||||
return nil
|
||||
}
|
||||
serveDone := make(chan struct{})
|
||||
go func() {
|
||||
ServeNAT(writer, request, &model.NAT{Common: model.Common{ID: 91}, ServerID: fixture.server.ID, Host: "target.example"})
|
||||
close(serveDone)
|
||||
}()
|
||||
streamID, err := fixture.handler.WaitAgentCompatIOStreamCapability(context.Background(), serviceRPC.AgentCompatCapabilityAccess{
|
||||
Capability: access,
|
||||
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2},
|
||||
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||
TargetServerID: fixture.server.ID,
|
||||
ResourceID: 91,
|
||||
ServerAccessAllowed: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
taskStreamID := <-taskStreamIDReady
|
||||
require.Equal(t, taskStreamID, streamID)
|
||||
select {
|
||||
case <-agent.readStarted:
|
||||
case <-serveDone:
|
||||
t.Fatal("ServeNAT returned before StartStream read")
|
||||
}
|
||||
select {
|
||||
case <-agent.writeStarted:
|
||||
case <-serveDone:
|
||||
t.Fatal("ServeNAT returned before request reached agent")
|
||||
}
|
||||
close(agent.readRelease)
|
||||
require.NoError(t, connection.Close())
|
||||
completionContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
select {
|
||||
case <-serveDone:
|
||||
case <-completionContext.Done():
|
||||
t.Fatal("ServeNAT did not complete after StartStream")
|
||||
}
|
||||
require.NoError(t, fixture.handler.CancelAgentCompatIOStreamCapability(serviceRPC.AgentCompatCapabilityAccess{
|
||||
Capability: access,
|
||||
Owner: serviceRPC.AgentCompatCapabilityOwner{PATID: 1, UserID: 2},
|
||||
Purpose: serviceRPC.AgentCompatCapabilityNAT,
|
||||
TargetServerID: fixture.server.ID,
|
||||
ResourceID: 91,
|
||||
ServerAccessAllowed: true,
|
||||
}))
|
||||
require.NotContains(t, string(agent.writtenBytes()), capability.String())
|
||||
forwarded := string(agent.writtenBytes())
|
||||
for _, expected := range []string{"POST /nat HTTP/1.1", "Host: example.test", "X-Ordinary-Nat: ordinary", "payload"} {
|
||||
require.Contains(t, forwarded, expected)
|
||||
}
|
||||
require.NotContains(t, forwarded, agentcompatcontract.IOStreamCapabilityHeader)
|
||||
require.NotContains(t, forwarded, "Authorization:")
|
||||
require.NotContains(t, forwarded, "deterministic-secret")
|
||||
}
|
||||
|
||||
type orderedNATAgent struct {
|
||||
mu sync.Mutex
|
||||
written bytes.Buffer
|
||||
readStarted chan struct{}
|
||||
readRelease chan struct{}
|
||||
writeStarted chan struct{}
|
||||
readOnce sync.Once
|
||||
writeOnce sync.Once
|
||||
closeCount int
|
||||
}
|
||||
|
||||
func (agent *orderedNATAgent) Read([]byte) (int, error) {
|
||||
agent.readOnce.Do(func() { close(agent.readStarted) })
|
||||
<-agent.readRelease
|
||||
return 0, io.EOF
|
||||
}
|
||||
|
||||
func (agent *orderedNATAgent) Write(data []byte) (int, error) {
|
||||
agent.mu.Lock()
|
||||
count, err := agent.written.Write(data)
|
||||
agent.mu.Unlock()
|
||||
agent.writeOnce.Do(func() { close(agent.writeStarted) })
|
||||
return count, err
|
||||
}
|
||||
|
||||
func (agent *orderedNATAgent) Close() error {
|
||||
agent.mu.Lock()
|
||||
defer agent.mu.Unlock()
|
||||
agent.closeCount++
|
||||
return nil
|
||||
}
|
||||
|
||||
func (agent *orderedNATAgent) writtenBytes() []byte {
|
||||
agent.mu.Lock()
|
||||
defer agent.mu.Unlock()
|
||||
return append([]byte(nil), agent.written.Bytes()...)
|
||||
}
|
||||
|
||||
var _ io.ReadWriteCloser = (*orderedNATAgent)(nil)
|
||||
@@ -0,0 +1,93 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"io"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/goccy/go-json"
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
)
|
||||
|
||||
func TestServeNATClosesHijackedResourcesWhenStreamDisappearsBeforeUserTransfer(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
conn := newServeNATConn()
|
||||
body := &serveNATBody{}
|
||||
writer := &serveNATResponseWriter{conn: conn}
|
||||
request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: body, ContentLength: 0}
|
||||
request.Header.Set("X-NAT-Test", "failure")
|
||||
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||
var nat model.TaskNAT
|
||||
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||
return err
|
||||
}
|
||||
return fixture.handler.CloseStream(nat.StreamID)
|
||||
}
|
||||
|
||||
ServeNAT(writer, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"})
|
||||
|
||||
if got := body.closeCount.Load(); got == 0 {
|
||||
t.Fatal("request body was not closed after failed transfer")
|
||||
}
|
||||
if got := conn.closeCount.Load(); got != 1 {
|
||||
t.Fatalf("hijacked connection close count = %d, want 1", got)
|
||||
}
|
||||
select {
|
||||
case <-conn.readDone:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("hijacked connection read remained blocked after failed transfer")
|
||||
}
|
||||
if writer.status != 0 || writer.writes != 0 {
|
||||
t.Fatalf("hijacked failure wrote HTTP response: status=%d writes=%d", writer.status, writer.writes)
|
||||
}
|
||||
}
|
||||
|
||||
func TestServeNATTransfersRequestAndRegistryOwnsSuccessfulCleanup(t *testing.T) {
|
||||
fixture := newServeNATFixture(t)
|
||||
conn := newServeNATConn()
|
||||
body := &serveNATBody{}
|
||||
agent := &serveNATAgent{readErr: io.EOF, writeDone: make(chan struct{})}
|
||||
writer := &serveNATResponseWriter{conn: conn}
|
||||
request := &http.Request{Method: http.MethodPost, URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"}, Header: make(http.Header), Body: body, ContentLength: 0}
|
||||
request.Header.Set("X-NAT-Test", "success")
|
||||
fixture.taskStream.onSend = func(task *proto.Task) error {
|
||||
var nat model.TaskNAT
|
||||
if err := json.Unmarshal([]byte(task.Data), &nat); err != nil {
|
||||
return err
|
||||
}
|
||||
return fixture.handler.AgentConnected(nat.StreamID, agent)
|
||||
}
|
||||
|
||||
serveDone := make(chan struct{})
|
||||
go func() {
|
||||
ServeNAT(writer, request, &model.NAT{ServerID: fixture.server.ID, Host: "target.example"})
|
||||
close(serveDone)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-agent.writeDone:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("agent did not receive the transferred NAT request")
|
||||
}
|
||||
select {
|
||||
case <-serveDone:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("ServeNAT did not finish after the successful stream was closed")
|
||||
}
|
||||
if got := body.closeCount.Load(); got == 0 {
|
||||
t.Fatal("registry-owned request body was not closed")
|
||||
}
|
||||
if got := conn.closeCount.Load(); got != 1 {
|
||||
t.Fatalf("registry-owned hijacked connection close count = %d, want 1", got)
|
||||
}
|
||||
if got := agent.closeCount.Load(); got != 1 {
|
||||
t.Fatalf("registry-owned agent close count = %d, want 1", got)
|
||||
}
|
||||
if got := strings.ToLower(string(agent.writtenBytes())); !strings.Contains(got, "post /nat http/1.1") || !strings.Contains(got, "x-nat-test: success") {
|
||||
t.Fatalf("agent received incomplete NAT request: %q", got)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,171 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"bytes"
|
||||
"context"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
rpcService "github.com/nezhahq/nezha/service/rpc"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"gorm.io/driver/sqlite"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
type serveNATFixture struct {
|
||||
handler *rpcService.NezhaHandler
|
||||
server *model.Server
|
||||
taskStream *serveNATTaskStream
|
||||
}
|
||||
|
||||
func newServeNATFixture(t *testing.T) serveNATFixture {
|
||||
t.Helper()
|
||||
originalDB, originalServerShared, originalHandler := singleton.DB, singleton.ServerShared, rpcService.NezhaHandlerSingleton
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||
require.NoError(t, err)
|
||||
require.NoError(t, db.AutoMigrate(model.Server{}))
|
||||
server := &model.Server{Common: model.Common{ID: 7}, UUID: "serve-nat-test", Name: "serve-nat-test"}
|
||||
require.NoError(t, db.Create(server).Error)
|
||||
singleton.DB = db
|
||||
singleton.ServerShared = singleton.NewServerClass()
|
||||
handler := rpcService.NewNezhaHandler()
|
||||
rpcService.NezhaHandlerSingleton = handler
|
||||
taskStream := &serveNATTaskStream{}
|
||||
server, ok := singleton.ServerShared.Get(server.ID)
|
||||
require.True(t, ok)
|
||||
server.SetTaskStream(taskStream)
|
||||
taskStream.server = server
|
||||
t.Cleanup(func() {
|
||||
rpcService.NezhaHandlerSingleton, singleton.ServerShared, singleton.DB = originalHandler, originalServerShared, originalDB
|
||||
if dbSQL, dbErr := db.DB(); dbErr == nil {
|
||||
_ = dbSQL.Close()
|
||||
}
|
||||
})
|
||||
return serveNATFixture{handler: handler, server: server, taskStream: taskStream}
|
||||
}
|
||||
|
||||
type serveNATTaskStream struct {
|
||||
server *model.Server
|
||||
onSend func(*proto.Task) error
|
||||
sendErr error
|
||||
sent []*proto.Task
|
||||
}
|
||||
|
||||
func (stream *serveNATTaskStream) Send(task *proto.Task) error {
|
||||
stream.sent = append(stream.sent, task)
|
||||
if stream.sendErr != nil {
|
||||
return stream.sendErr
|
||||
}
|
||||
if stream.onSend != nil {
|
||||
return stream.onSend(task)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
func (*serveNATTaskStream) Recv() (*proto.TaskResult, error) { return nil, io.EOF }
|
||||
func (*serveNATTaskStream) SetHeader(metadata.MD) error { return nil }
|
||||
func (*serveNATTaskStream) SendHeader(metadata.MD) error { return nil }
|
||||
func (*serveNATTaskStream) SetTrailer(metadata.MD) {}
|
||||
func (*serveNATTaskStream) Context() context.Context { return context.Background() }
|
||||
func (*serveNATTaskStream) SendMsg(any) error { return nil }
|
||||
func (*serveNATTaskStream) RecvMsg(any) error { return io.EOF }
|
||||
|
||||
type serveNATResponseWriter struct {
|
||||
conn *serveNATConn
|
||||
header http.Header
|
||||
status, writes int
|
||||
body string
|
||||
}
|
||||
|
||||
func (writer *serveNATResponseWriter) Header() http.Header {
|
||||
if writer.header == nil {
|
||||
writer.header = make(http.Header)
|
||||
}
|
||||
return writer.header
|
||||
}
|
||||
func (writer *serveNATResponseWriter) Write(data []byte) (int, error) {
|
||||
writer.writes += len(data)
|
||||
writer.body += string(data)
|
||||
return len(data), nil
|
||||
}
|
||||
func (writer *serveNATResponseWriter) WriteHeader(status int) { writer.status = status }
|
||||
func (writer *serveNATResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
|
||||
go func() { _, _ = writer.conn.Read(make([]byte, 1)) }()
|
||||
return writer.conn, bufio.NewReadWriter(bufio.NewReader(writer.conn), bufio.NewWriter(writer.conn)), nil
|
||||
}
|
||||
|
||||
type serveNATBody struct{ closeCount atomic.Int32 }
|
||||
|
||||
func (*serveNATBody) Read([]byte) (int, error) { return 0, io.EOF }
|
||||
func (body *serveNATBody) Close() error { body.closeCount.Add(1); return nil }
|
||||
|
||||
type serveNATConn struct {
|
||||
closed, readDone chan struct{}
|
||||
readDoneOnce, closeOnce sync.Once
|
||||
closeCount atomic.Int32
|
||||
}
|
||||
|
||||
func newServeNATConn() *serveNATConn {
|
||||
return &serveNATConn{closed: make(chan struct{}), readDone: make(chan struct{})}
|
||||
}
|
||||
func (conn *serveNATConn) Read([]byte) (int, error) {
|
||||
<-conn.closed
|
||||
conn.readDoneOnce.Do(func() { close(conn.readDone) })
|
||||
return 0, io.EOF
|
||||
}
|
||||
func (*serveNATConn) Write(data []byte) (int, error) { return len(data), nil }
|
||||
func (conn *serveNATConn) Close() error {
|
||||
conn.closeCount.Add(1)
|
||||
conn.closeOnce.Do(func() { close(conn.closed) })
|
||||
return nil
|
||||
}
|
||||
func (*serveNATConn) LocalAddr() net.Addr { return serveNATAddr("local") }
|
||||
func (*serveNATConn) RemoteAddr() net.Addr { return serveNATAddr("remote") }
|
||||
func (*serveNATConn) SetDeadline(time.Time) error { return nil }
|
||||
func (*serveNATConn) SetReadDeadline(time.Time) error { return nil }
|
||||
func (*serveNATConn) SetWriteDeadline(time.Time) error { return nil }
|
||||
|
||||
type serveNATAddr string
|
||||
|
||||
func (addr serveNATAddr) Network() string { return "test" }
|
||||
func (addr serveNATAddr) String() string { return string(addr) }
|
||||
|
||||
type serveNATAgent struct {
|
||||
mu sync.Mutex
|
||||
written bytes.Buffer
|
||||
readErr error
|
||||
writeDone chan struct{}
|
||||
writeOnce sync.Once
|
||||
closeCount atomic.Int32
|
||||
}
|
||||
|
||||
func (agent *serveNATAgent) Read([]byte) (int, error) { return 0, agent.readErr }
|
||||
func (agent *serveNATAgent) Write(data []byte) (int, error) {
|
||||
agent.mu.Lock()
|
||||
defer agent.mu.Unlock()
|
||||
count, err := agent.written.Write(data)
|
||||
if agent.writeDone != nil {
|
||||
agent.writeOnce.Do(func() { close(agent.writeDone) })
|
||||
}
|
||||
return count, err
|
||||
}
|
||||
func (agent *serveNATAgent) Close() error { agent.closeCount.Add(1); return nil }
|
||||
func (agent *serveNATAgent) writtenBytes() []byte {
|
||||
agent.mu.Lock()
|
||||
defer agent.mu.Unlock()
|
||||
return append([]byte(nil), agent.written.Bytes()...)
|
||||
}
|
||||
|
||||
var _ proto.NezhaService_RequestTaskServer = (*serveNATTaskStream)(nil)
|
||||
var _ http.Hijacker = (*serveNATResponseWriter)(nil)
|
||||
var _ net.Conn = (*serveNATConn)(nil)
|
||||
var _ io.ReadWriteCloser = (*serveNATAgent)(nil)
|
||||
+9
-238
@@ -1,27 +1,23 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"time"
|
||||
"net"
|
||||
|
||||
"github.com/goccy/go-json"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/peer"
|
||||
|
||||
"github.com/hashicorp/go-uuid"
|
||||
"github.com/nezhahq/nezha/model"
|
||||
"github.com/nezhahq/nezha/pkg/utils"
|
||||
"github.com/nezhahq/nezha/proto"
|
||||
rpcService "github.com/nezhahq/nezha/service/rpc"
|
||||
"github.com/nezhahq/nezha/service/singleton"
|
||||
)
|
||||
|
||||
func SetReceiptGateListener(listener net.Listener) {
|
||||
rpcService.SetReceiptGateListener(listener)
|
||||
}
|
||||
|
||||
func CloseReceiptGate() {
|
||||
rpcService.CloseReceiptGate()
|
||||
}
|
||||
|
||||
// SetMCPKillSwitchObserver re-exports the service/rpc hook so cmd/dashboard
|
||||
// can wire singleton.Conf.EnableMCP without importing the inner rpc package
|
||||
// (cmd/dashboard already imports cmd/dashboard/rpc for ServeRPC).
|
||||
@@ -46,228 +42,3 @@ func ServeRPC() *grpc.Server {
|
||||
proto.RegisterNezhaServiceServer(server, rpcService.NezhaHandlerSingleton)
|
||||
return server
|
||||
}
|
||||
|
||||
func ctxWithRealIP(ctx context.Context) (context.Context, error) {
|
||||
var ip, connectingIp string
|
||||
p, ok := peer.FromContext(ctx)
|
||||
if ok {
|
||||
addrPort, err := netip.ParseAddrPort(p.Addr.String())
|
||||
if err == nil {
|
||||
connectingIp = addrPort.Addr().String()
|
||||
}
|
||||
}
|
||||
ctx = context.WithValue(ctx, model.CtxKeyConnectingIP{}, connectingIp)
|
||||
|
||||
if singleton.Conf.AgentRealIPHeader == "" {
|
||||
return ctx, nil
|
||||
}
|
||||
|
||||
if singleton.Conf.AgentRealIPHeader == model.ConfigUsePeerIP {
|
||||
if connectingIp == "" {
|
||||
return ctx, fmt.Errorf("connecting ip not found")
|
||||
}
|
||||
// Peer-IP mode: peer IP is the real IP. Leaving ip="" makes
|
||||
// CheckIP/BlockIP short-circuit on empty IP, disabling the WAF.
|
||||
ip = connectingIp
|
||||
} else {
|
||||
vals := metadata.ValueFromIncomingContext(ctx, singleton.Conf.AgentRealIPHeader)
|
||||
if len(vals) == 0 {
|
||||
return ctx, fmt.Errorf("real ip header not found")
|
||||
}
|
||||
var err error
|
||||
ip, err = utils.GetIPFromHeader(vals[0])
|
||||
if err != nil {
|
||||
return ctx, err
|
||||
}
|
||||
}
|
||||
|
||||
if singleton.Conf.Debug {
|
||||
log.Printf("NEZHA>> gRPC Agent Real IP: %s, connecting IP: %s\n", ip, connectingIp)
|
||||
}
|
||||
|
||||
return context.WithValue(ctx, model.CtxKeyRealIP{}, ip), nil
|
||||
}
|
||||
|
||||
func waf(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
realip, _ := ctx.Value(model.CtxKeyRealIP{}).(string)
|
||||
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return handler(ctx, req)
|
||||
}
|
||||
|
||||
func getRealIp(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
ctx, err := ctxWithRealIP(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return handler(ctx, req)
|
||||
}
|
||||
|
||||
// realIPServerStream overrides Context() so stream handlers and
|
||||
// authHandler.check observe the resolved real IP, like the unary path.
|
||||
type realIPServerStream struct {
|
||||
grpc.ServerStream
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
func (s *realIPServerStream) Context() context.Context { return s.ctx }
|
||||
|
||||
func getRealIpStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||
ctx, err := ctxWithRealIP(ss.Context())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return handler(srv, &realIPServerStream{ServerStream: ss, ctx: ctx})
|
||||
}
|
||||
|
||||
func wafStream(srv any, ss grpc.ServerStream, info *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
|
||||
realip, _ := ss.Context().Value(model.CtxKeyRealIP{}).(string)
|
||||
if err := model.CheckIP(singleton.DB, realip); err != nil {
|
||||
return err
|
||||
}
|
||||
return handler(srv, ss)
|
||||
}
|
||||
|
||||
func DispatchTask(serviceSentinelDispatchBus <-chan *model.Service) {
|
||||
for task := range serviceSentinelDispatchBus {
|
||||
if task == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
switch task.Cover {
|
||||
case model.ServiceCoverIgnoreAll:
|
||||
for id, enabled := range task.SkipServers {
|
||||
if !enabled {
|
||||
continue
|
||||
}
|
||||
|
||||
server, _ := singleton.ServerShared.Get(id)
|
||||
if server == nil {
|
||||
continue
|
||||
}
|
||||
if !canSendTaskToServer(task, server) {
|
||||
continue
|
||||
}
|
||||
// SendTask 走 holder-scoped send mutex,避免与 cron /
|
||||
// server-transfer / MCP CallAgent / fs.transfer 等并发
|
||||
// SendMsg 同一 RequestTask stream。
|
||||
if err := server.SendTask(task.PB()); err != nil &&
|
||||
!errors.Is(err, model.ErrTaskStreamOffline) {
|
||||
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
|
||||
}
|
||||
}
|
||||
case model.ServiceCoverAll:
|
||||
// 快照后逐个 SendTask,不在 ServerShared 的 listMu.RLock 内做阻塞
|
||||
// gRPC:否则一个卡死 agent 会拖死需要写锁的 server 生命周期操作。
|
||||
for id, server := range singleton.ServerShared.GetList() {
|
||||
if server == nil || task.SkipServers[id] {
|
||||
continue
|
||||
}
|
||||
if !canSendTaskToServer(task, server) {
|
||||
continue
|
||||
}
|
||||
if err := server.SendTask(task.PB()); err != nil &&
|
||||
!errors.Is(err, model.ErrTaskStreamOffline) {
|
||||
log.Printf("NEZHA>> DispatchTask send error (server=%d): %v", id, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func DispatchKeepalive() {
|
||||
singleton.CronShared.AddFunc("@every 20s", func() {
|
||||
list := singleton.ServerShared.GetSortedList()
|
||||
for _, s := range list {
|
||||
if s == nil {
|
||||
continue
|
||||
}
|
||||
if err := s.SendTask(&proto.Task{Type: model.TaskTypeKeepalive}); err != nil &&
|
||||
!errors.Is(err, model.ErrTaskStreamOffline) {
|
||||
log.Printf("NEZHA>> Keepalive send error (server=%d): %v", s.ID, err)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func ServeNAT(w http.ResponseWriter, r *http.Request, natConfig *model.NAT) {
|
||||
server, _ := singleton.ServerShared.Get(natConfig.ServerID)
|
||||
if server == nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write([]byte("server not found or not connected"))
|
||||
return
|
||||
}
|
||||
if server.GetTaskStream() == nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write([]byte("server not found or not connected"))
|
||||
return
|
||||
}
|
||||
|
||||
streamId, err := uuid.GenerateUUID()
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "stream id error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
// NAT streams are anonymous HTTP-facing tunnels; they are NOT reachable
|
||||
// via /ws/terminal or /ws/file (which check stream ownership), so the
|
||||
// creator user ID does not need to identify a real user. The targetServerID
|
||||
// IS required though — the receiving agent must prove it is the server the
|
||||
// NAT config addressed, otherwise any agent that snoops the streamId can
|
||||
// answer NAT traffic on behalf of an unrelated host.
|
||||
if err := rpcService.NezhaHandlerSingleton.CreateStream(streamId, 0, server.ID); err != nil {
|
||||
w.WriteHeader(http.StatusTooManyRequests)
|
||||
w.Write(fmt.Appendf(nil, "stream limit: %v", err))
|
||||
return
|
||||
}
|
||||
defer rpcService.NezhaHandlerSingleton.CloseStream(streamId)
|
||||
|
||||
taskData, err := json.Marshal(model.TaskNAT{
|
||||
StreamID: streamId,
|
||||
Host: natConfig.Host,
|
||||
})
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "task data error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
if err := server.SendTask(&proto.Task{
|
||||
Type: model.TaskTypeNAT,
|
||||
Data: string(taskData),
|
||||
}); err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "send task error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
wWrapped, err := utils.NewRequestWrapper(r, w)
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "request wrapper error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
if err := rpcService.NezhaHandlerSingleton.UserConnected(streamId, wWrapped); err != nil {
|
||||
w.WriteHeader(http.StatusServiceUnavailable)
|
||||
w.Write(fmt.Appendf(nil, "user connected error: %v", err))
|
||||
return
|
||||
}
|
||||
|
||||
rpcService.NezhaHandlerSingleton.StartStream(streamId, time.Second*10)
|
||||
}
|
||||
|
||||
func canSendTaskToServer(task *model.Service, server *model.Server) bool {
|
||||
var role model.Role
|
||||
singleton.UserLock.RLock()
|
||||
if u, ok := singleton.UserInfoMap[task.UserID]; !ok {
|
||||
role = model.RoleMember
|
||||
} else {
|
||||
role = u.Role
|
||||
}
|
||||
singleton.UserLock.RUnlock()
|
||||
|
||||
return task.UserID == server.GetUserID() || role.IsAdmin()
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user