71 lines
2.9 KiB
Go
71 lines
2.9 KiB
Go
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()}
|
|
}
|