feat(agentcompat): expose dashboard capability routes

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:30:50 +00:00
co-authored by naiba/CloudCode
parent 0c6eb416a5
commit 2640b86d3a
42 changed files with 4083 additions and 271 deletions
+94
View File
@@ -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)
}
+115
View File
@@ -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)
+93
View File
@@ -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)
}
}
+171
View File
@@ -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
View File
@@ -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()
}