From fbfd9428d126b2c43d47e7c83b65d8849be93352 Mon Sep 17 00:00:00 2001 From: Sasha Mitchell Date: Sat, 8 Aug 2026 05:00:46 +0700 Subject: [PATCH] fix(remote-debug): gate registration messages on completed handshake The remote-debugging TCP server authenticates clients with a handshake message carrying a per-tenant key, but onMessage never enforced it: ASSET_CHUNK, the declaration messages and END were all processed while runtime.handshake was still false. The only enforcement was a 10 second timer closing connections that never handshake, and a full registration completes well inside that window, so an unauthenticated client could write assets to media storage, consume connection slots and register arbitrary plugin identities in the node registry. Reject every registration message other than HAND_SHAKE until the handshake completes, close the connection with a handshake failure and latch handshakeFailed so later messages are ignored. The 10 second timer stays as-is. --- internal/core/debugging_runtime/hooks.go | 8 + internal/core/debugging_runtime/hooks_test.go | 173 ++++++++++++++++++ 2 files changed, 181 insertions(+) create mode 100644 internal/core/debugging_runtime/hooks_test.go diff --git a/internal/core/debugging_runtime/hooks.go b/internal/core/debugging_runtime/hooks.go index 758e20026..a696bbd6d 100644 --- a/internal/core/debugging_runtime/hooks.go +++ b/internal/core/debugging_runtime/hooks.go @@ -178,6 +178,14 @@ func (s *DifyServer) onMessage(runtime *RemotePluginRuntime, message []byte) { return } + // only the handshake itself is accepted before the handshake completes, + // any other registration message is rejected and fails closed + if !runtime.handshake && registerPayload.Type != plugin_entities.REGISTER_EVENT_TYPE_HAND_SHAKE { + runtime.handshakeFailed = true + closeConn([]byte("handshake failed, registration message before handshake\n")) + return + } + switch registerPayload.Type { case plugin_entities.REGISTER_EVENT_TYPE_HAND_SHAKE: if connectionInfo, err := s.handleHandleShake(runtime, registerPayload); err != nil { diff --git a/internal/core/debugging_runtime/hooks_test.go b/internal/core/debugging_runtime/hooks_test.go new file mode 100644 index 000000000..a2f8f9a23 --- /dev/null +++ b/internal/core/debugging_runtime/hooks_test.go @@ -0,0 +1,173 @@ +package debugging_runtime + +import ( + "fmt" + "net" + "strings" + "testing" + "time" + + cloudoss "github.com/langgenius/dify-cloud-kit/oss" + "github.com/langgenius/dify-cloud-kit/oss/factory" + "github.com/langgenius/dify-plugin-daemon/internal/core/plugin_manager/media_transport" + "github.com/langgenius/dify-plugin-daemon/internal/types/app" + "github.com/langgenius/dify-plugin-daemon/pkg/entities/constants" + "github.com/langgenius/dify-plugin-daemon/pkg/entities/manifest_entities" + "github.com/langgenius/dify-plugin-daemon/pkg/entities/plugin_entities" + "github.com/langgenius/dify-plugin-daemon/pkg/utils/network" + "github.com/langgenius/dify-plugin-daemon/pkg/utils/parser" +) + +// Registration messages sent before a completed handshake must be rejected +// and the connection must fail closed without registering any runtime. +func TestRegistrationRejectedBeforeHandshake(t *testing.T) { + port, err := network.GetRandomPort() + if err != nil { + t.Fatalf("failed to get random port: %s", err.Error()) + } + + oss, err := factory.Load("local", cloudoss.OSSArgs{ + Local: &cloudoss.Local{ + Path: t.TempDir(), + }, + }) + if err != nil { + t.Fatalf("failed to load local storage: %s", err.Error()) + } + + server := NewDebuggingPluginServer(&app.Config{ + PluginRemoteInstallingHost: "127.0.0.1", + PluginRemoteInstallingPort: port, + PluginRemoteInstallingMaxConn: 10, + PluginRemoteInstallServerEventLoopNums: 1, + }, media_transport.NewAssetsBucket(oss, "assets", 10)) + defer server.Stop() + go server.Launch() + + registered := make(chan struct{}, 1) + server.AddNotifier(&TestPluginRuntimeNotifier{ + onConnected: func(runtime *RemotePluginRuntime) error { + registered <- struct{}{} + return nil + }, + }) + + // wait for the server to start + time.Sleep(time.Second * 2) + + manifest := parser.MarshalJsonBytes(&plugin_entities.PluginDeclaration{ + PluginDeclarationWithoutAdvancedFields: plugin_entities.PluginDeclarationWithoutAdvancedFields{ + Version: "1.0.0", + Type: manifest_entities.PluginType, + Description: plugin_entities.I18nObject{ + EnUS: "test", + }, + Author: "test", + Name: "pre_handshake_test", + Icon: "icon.svg", + Label: plugin_entities.I18nObject{EnUS: "test"}, + CreatedAt: time.Now(), + Resource: plugin_entities.PluginResourceRequirement{Memory: 1}, + Plugins: plugin_entities.PluginExtensions{ + Endpoints: []string{"test"}, + }, + Meta: plugin_entities.PluginMeta{ + Version: "0.0.1", + Arch: []constants.Arch{constants.AMD64}, + Runner: plugin_entities.PluginRunner{ + Language: constants.Python, + Version: "3.12", + Entrypoint: "main", + }, + }, + }, + }) + + messages := []struct { + name string + payload plugin_entities.RemotePluginRegisterPayload + }{ + { + name: "manifest declaration", + payload: plugin_entities.RemotePluginRegisterPayload{ + Type: plugin_entities.REGISTER_EVENT_TYPE_MANIFEST_DECLARATION, + Data: manifest, + }, + }, + { + name: "endpoint declaration", + payload: plugin_entities.RemotePluginRegisterPayload{ + Type: plugin_entities.REGISTER_EVENT_TYPE_ENDPOINT_DECLARATION, + Data: parser.MarshalJsonBytes([]plugin_entities.EndpointProviderDeclaration{ + { + Settings: []plugin_entities.ProviderConfig{}, + Endpoints: []plugin_entities.EndpointDeclaration{ + {Path: "/test", Method: "GET"}, + }, + }, + }), + }, + }, + { + name: "asset chunk", + payload: plugin_entities.RemotePluginRegisterPayload{ + Type: plugin_entities.REGISTER_EVENT_TYPE_ASSET_CHUNK, + Data: parser.MarshalJsonBytes(plugin_entities.RemotePluginRegisterAssetChunk{ + Filename: "icon.svg", + Data: "QUJD", + End: true, + }), + }, + }, + { + name: "initialization end", + payload: plugin_entities.RemotePluginRegisterPayload{ + Type: plugin_entities.REGISTER_EVENT_TYPE_END, + Data: []byte("{}"), + }, + }, + } + + for _, tt := range messages { + t.Run(tt.name, func(t *testing.T) { + conn, err := net.Dial("tcp", fmt.Sprintf("127.0.0.1:%d", port)) + if err != nil { + t.Fatalf("failed to connect to plugin server: %s", err.Error()) + } + defer conn.Close() + + // NOTE: no handshake message, straight to registration + if _, err := conn.Write(parser.MarshalJsonBytes(tt.payload)); err != nil { + t.Fatalf("failed to write payload: %s", err.Error()) + } + if _, err := conn.Write([]byte("\n\n")); err != nil { + t.Fatalf("failed to write delimiter: %s", err.Error()) + } + + // the server must reject the message and close the connection + conn.SetReadDeadline(time.Now().Add(5 * time.Second)) + var reply strings.Builder + buf := make([]byte, 1024) + for { + n, err := conn.Read(buf) + if n > 0 { + reply.Write(buf[:n]) + } + if err != nil { + break + } + } + if !strings.Contains(reply.String(), "handshake failed") { + t.Errorf("expected handshake failure reply, got %q", reply.String()) + } + }) + } + + // no runtime may be registered + time.Sleep(time.Second) + select { + case <-registered: + t.Fatal("runtime registered without a completed handshake") + default: + } +}