From bcb31a7f4ea26a4988fb2d6ffc4f794503011ef0 Mon Sep 17 00:00:00 2001 From: cabbage Date: Sun, 6 Sep 2026 07:01:24 +0000 Subject: [PATCH] feat: add client agent handshake transport --- internal/client/agent/handshake.go | 61 ++++++++++++++++ internal/client/agent/handshake_test.go | 70 +++++++++++++++++++ internal/client/agent/websocket.go | 49 +++++++++++++ test/coverage.toml | 18 +++++ .../clientagent_integration_test.go | 42 +++++++++++ 5 files changed, 240 insertions(+) create mode 100644 internal/client/agent/handshake.go create mode 100644 internal/client/agent/handshake_test.go create mode 100644 internal/client/agent/websocket.go create mode 100644 test/integration/clientagent/clientagent_integration_test.go diff --git a/internal/client/agent/handshake.go b/internal/client/agent/handshake.go new file mode 100644 index 0000000..a929192 --- /dev/null +++ b/internal/client/agent/handshake.go @@ -0,0 +1,61 @@ +// Package agent implements the client side of the RVBox Agent protocol. +package agent + +import ( + "context" + "errors" + "fmt" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "google.golang.org/protobuf/proto" +) + +var ( + ErrProtocolHandshake = errors.New("invalid agent session handshake") + ErrSessionClosed = errors.New("agent session transport closed") +) + +// Transport has exactly the binary-message operations the protocol needs. A +// WebSocket implementation owns its deadline/ping details outside this layer. +type Transport interface { + Write(context.Context, []byte) error + Read(context.Context) ([]byte, error) + Close() error +} + +type Session struct { + ID string + Generation uint64 + Protocol *rvboxv1.ProtocolVersion +} + +func Handshake(ctx context.Context, transport Transport, hello *rvboxv1.ClientHello, limits agentproto.Limits) (Session, error) { + if transport == nil || hello == nil { + return Session{}, ErrProtocolHandshake + } + request := &rvboxv1.AgentEnvelope{Payload: &rvboxv1.AgentEnvelope_ClientHello{ClientHello: hello}} + if err := agentproto.ValidateEnvelope(request, limits, rvboxv1.Platform_PLATFORM_WINDOWS); err != nil { + return Session{}, fmt.Errorf("validate ClientHello: %w", err) + } + encoded, err := proto.Marshal(request) + if err != nil { + return Session{}, err + } + if err := transport.Write(ctx, encoded); err != nil { + return Session{}, err + } + received, err := transport.Read(ctx) + if err != nil { + return Session{}, err + } + welcome, err := agentproto.DecodeEnvelope(received, limits, rvboxv1.Platform_PLATFORM_WINDOWS) + if err != nil { + return Session{}, fmt.Errorf("decode ServerWelcome: %w", err) + } + payload := welcome.GetServerWelcome() + if payload == nil || welcome.GetSessionId() == "" || welcome.GetSessionGeneration() == 0 || payload.GetSelectedProtocol() == nil { + return Session{}, ErrProtocolHandshake + } + return Session{ID: welcome.GetSessionId(), Generation: welcome.GetSessionGeneration(), Protocol: proto.Clone(payload.GetSelectedProtocol()).(*rvboxv1.ProtocolVersion)}, nil +} diff --git a/internal/client/agent/handshake_test.go b/internal/client/agent/handshake_test.go new file mode 100644 index 0000000..9b27591 --- /dev/null +++ b/internal/client/agent/handshake_test.go @@ -0,0 +1,70 @@ +package agent + +import ( + "context" + "errors" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/timestamppb" +) + +func TestHandshakeHelloWelcome_HP_SES_07(t *testing.T) { + t.Parallel() + welcome := &rvboxv1.AgentEnvelope{SessionId: "issued-session", SessionGeneration: 7, Payload: &rvboxv1.AgentEnvelope_ServerWelcome{ServerWelcome: &rvboxv1.ServerWelcome{SelectedProtocol: &rvboxv1.ProtocolVersion{Major: 1, Minor: 0}, ServerTime: timestamppb.New(time.Date(2026, time.September, 6, 0, 0, 0, 0, time.UTC))}}} + encoded, err := proto.Marshal(welcome) + if err != nil { + t.Fatal(err) + } + transport := &fakeTransport{read: encoded} + session, err := Handshake(context.Background(), transport, validHello(), agentproto.DefaultLimits()) + if err != nil { + t.Fatal(err) + } + if session.ID != "issued-session" || session.Generation != 7 || session.Protocol.GetMajor() != 1 { + t.Fatalf("session = %#v", session) + } + var sent rvboxv1.AgentEnvelope + if err := proto.Unmarshal(transport.written, &sent); err != nil || sent.GetClientHello() == nil || sent.GetSessionId() != "" || sent.GetSessionGeneration() != 0 { + t.Fatalf("sent ClientHello valid=%t session=%q generation=%d err=%v", sent.GetClientHello() != nil, sent.GetSessionId(), sent.GetSessionGeneration(), err) + } +} + +func TestHandshakeRejectsNonWelcome_BH_SES_07(t *testing.T) { + t.Parallel() + bad := &rvboxv1.AgentEnvelope{SessionId: "issued-session", SessionGeneration: 7, Payload: &rvboxv1.AgentEnvelope_ClientCapacity{ClientCapacity: &rvboxv1.ClientCapacity{MaxRunningCommands: 1, MaxQueuedCommands: 1}}} + encoded, err := proto.Marshal(bad) + if err != nil { + t.Fatal(err) + } + _, err = Handshake(context.Background(), &fakeTransport{read: encoded}, validHello(), agentproto.DefaultLimits()) + if !errors.Is(err, ErrProtocolHandshake) { + t.Fatalf("non-welcome handshake error = %v, want ErrProtocolHandshake", err) + } +} + +type fakeTransport struct { + written []byte + read []byte +} + +func (transport *fakeTransport) Write(_ context.Context, value []byte) error { + transport.written = append([]byte(nil), value...) + return nil +} + +func (transport *fakeTransport) Read(context.Context) ([]byte, error) { + if transport.read == nil { + return nil, ErrSessionClosed + } + return append([]byte(nil), transport.read...), nil +} + +func (transport *fakeTransport) Close() error { return nil } + +func validHello() *rvboxv1.ClientHello { + return &rvboxv1.ClientHello{ClientId: "win-client", SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: "019c46f1-1d02-7000-8000-000000000061", MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now()} +} diff --git a/internal/client/agent/websocket.go b/internal/client/agent/websocket.go new file mode 100644 index 0000000..98c37ac --- /dev/null +++ b/internal/client/agent/websocket.go @@ -0,0 +1,49 @@ +package agent + +import ( + "context" + "errors" + "net/http" + + "github.com/coder/websocket" +) + +var ErrUnexpectedMessage = errors.New("agent transport received a non-binary WebSocket message") + +type WebSocketTransport struct{ connection *websocket.Conn } + +func DialWebSocket(ctx context.Context, address string, client *http.Client) (*WebSocketTransport, error) { + connection, _, err := websocket.Dial(ctx, address, &websocket.DialOptions{HTTPClient: client, CompressionMode: websocket.CompressionDisabled}) + if err != nil { + return nil, err + } + return &WebSocketTransport{connection: connection}, nil +} + +func (transport *WebSocketTransport) Write(ctx context.Context, payload []byte) error { + if transport == nil || transport.connection == nil { + return ErrSessionClosed + } + return transport.connection.Write(ctx, websocket.MessageBinary, payload) +} + +func (transport *WebSocketTransport) Read(ctx context.Context) ([]byte, error) { + if transport == nil || transport.connection == nil { + return nil, ErrSessionClosed + } + kind, payload, err := transport.connection.Read(ctx) + if err != nil { + return nil, err + } + if kind != websocket.MessageBinary { + return nil, ErrUnexpectedMessage + } + return payload, nil +} + +func (transport *WebSocketTransport) Close() error { + if transport == nil || transport.connection == nil { + return nil + } + return transport.connection.Close(websocket.StatusNormalClosure, "") +} diff --git a/test/coverage.toml b/test/coverage.toml index f799e51..5f395d5 100644 --- a/test/coverage.toml +++ b/test/coverage.toml @@ -158,12 +158,30 @@ layer = "unit" status = "implemented" tests = ["internal/domain/protocol_test.go:TestSelectHighestCompatibleProtocol_HP_SES_01"] +[[requirements]] +id = "HP-SES-07" +layer = "unit" +status = "implemented" +tests = ["internal/client/agent/handshake_test.go:TestHandshakeHelloWelcome_HP_SES_07"] + +[[requirements]] +id = "BH-SES-07" +layer = "unit" +status = "implemented" +tests = ["internal/client/agent/handshake_test.go:TestHandshakeRejectsNonWelcome_BH_SES_07"] + [[requirements]] id = "HP-SES-05" layer = "integration" status = "implemented" tests = ["internal/server/session/agent_server_test.go:TestAgentServerRegistrationAndReplacement_HP_SES_05"] +[[requirements]] +id = "HP-SES-06" +layer = "integration" +status = "implemented" +tests = ["test/integration/clientagent/clientagent_integration_test.go:TestWebSocketHelloWelcome_HP_SES_06"] + [[requirements]] id = "BH-SES-04" layer = "integration" diff --git a/test/integration/clientagent/clientagent_integration_test.go b/test/integration/clientagent/clientagent_integration_test.go new file mode 100644 index 0000000..a2ea069 --- /dev/null +++ b/test/integration/clientagent/clientagent_integration_test.go @@ -0,0 +1,42 @@ +package clientagent_test + +import ( + "context" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + rvboxv1 "github.com/rvbox/rvbox/gen/go/rvbox/v1" + "github.com/rvbox/rvbox/internal/agentproto" + "github.com/rvbox/rvbox/internal/client/agent" + "github.com/rvbox/rvbox/internal/server/session" + "github.com/rvbox/rvbox/internal/server/store" + "google.golang.org/protobuf/types/known/timestamppb" +) + +func TestWebSocketHelloWelcome_HP_SES_06(t *testing.T) { + ctx := context.Background() + persistence, err := store.Open(ctx, store.Options{DataDir: filepath.Join(t.TempDir(), "state"), BusyTimeout: time.Second}) + if err != nil { + t.Fatal(err) + } + defer persistence.Close() + server := httptest.NewServer(&session.AgentServer{Store: persistence, Registry: session.NewRegistry(), Path: "/v1/agent", Limits: agentproto.DefaultLimits()}) + defer server.Close() + address := "ws" + strings.TrimPrefix(server.URL, "http") + "/v1/agent" + transport, err := agent.DialWebSocket(ctx, address, nil) + if err != nil { + t.Fatal(err) + } + defer transport.Close() + hello := &rvboxv1.ClientHello{ClientId: "win-client", SupportedProtocol: &rvboxv1.ProtocolRange{Major: 1, MinMinor: 0, MaxMinor: 0}, DaemonVersion: "test", Platform: rvboxv1.Platform_PLATFORM_WINDOWS, Architecture: "amd64", DaemonCwd: `C:\`, SupportedShells: []rvboxv1.ShellType{rvboxv1.ShellType_SHELL_POWERSHELL}, ClientInstanceId: "019c46f1-1d02-7000-8000-000000000071", MaxRunningCommands: 1, MaxQueuedCommands: 1, SentAt: timestamppb.Now()} + result, err := agent.Handshake(ctx, transport, hello, agentproto.DefaultLimits()) + if err != nil { + t.Fatal(err) + } + if result.ID == "" || result.Generation != 1 || result.Protocol.GetMajor() != 1 { + t.Fatalf("handshake result = %#v", result) + } +}