feat: add client agent handshake transport
This commit is contained in:
@@ -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
|
||||||
|
}
|
||||||
@@ -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()}
|
||||||
|
}
|
||||||
@@ -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, "")
|
||||||
|
}
|
||||||
@@ -158,12 +158,30 @@ layer = "unit"
|
|||||||
status = "implemented"
|
status = "implemented"
|
||||||
tests = ["internal/domain/protocol_test.go:TestSelectHighestCompatibleProtocol_HP_SES_01"]
|
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]]
|
[[requirements]]
|
||||||
id = "HP-SES-05"
|
id = "HP-SES-05"
|
||||||
layer = "integration"
|
layer = "integration"
|
||||||
status = "implemented"
|
status = "implemented"
|
||||||
tests = ["internal/server/session/agent_server_test.go:TestAgentServerRegistrationAndReplacement_HP_SES_05"]
|
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]]
|
[[requirements]]
|
||||||
id = "BH-SES-04"
|
id = "BH-SES-04"
|
||||||
layer = "integration"
|
layer = "integration"
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user