package socket import ( "context" "net" "testing" "time" altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" "git.toki-labs.com/toki/alt/services/worker/internal/config" workerContracts "git.toki-labs.com/toki/alt/services/worker/internal/contracts" protoSocket "git.toki-labs.com/toki/proto-socket/go" ) func TestWorkerSocketServerHello(t *testing.T) { // Get temporary port l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("failed to listen on temporary port: %v", err) } port := l.Addr().(*net.TCPAddr).Port l.Close() cfg := config.Config{ Host: "127.0.0.1", Port: port, SocketPath: "/socket", } ctx, cancel := context.WithCancel(context.Background()) defer cancel() server := NewServer(cfg, BacktestDeps{}) if err := server.Start(ctx); err != nil { t.Fatalf("failed to start worker socket server: %v", err) } defer server.Stop() // Wait slightly for server setup time.Sleep(50 * time.Millisecond) // Dial client, err := protoSocket.DialWs(ctx, "127.0.0.1", port, "/socket", workerContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial worker socket server: %v", err) } defer client.Close() // Send HelloRequest req := &altv1.HelloRequest{ AltProtocolVersion: "alt.v1", } res, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, req, 2*time.Second) if err != nil { t.Fatalf("failed to send HelloRequest: %v", err) } if res.ServerName != "alt-worker" { t.Errorf("expected ServerName to be %q, got %q", "alt-worker", res.ServerName) } if res.AltProtocolVersion != "alt.v1" { t.Errorf("expected AltProtocolVersion to be %q, got %q", "alt.v1", res.AltProtocolVersion) } // Assert no backtest-start or worker-execution since BacktestDeps{} has no Starter. for _, cap := range res.Capabilities { if cap == "backtest-start" || cap == "worker-execution" { t.Errorf("expected capability %q to be omitted when Starter is nil", cap) } } } func TestSessionHandlers(t *testing.T) { handlers := sessionHandlers(BacktestDeps{}) hasHello := false seenTypes := make(map[string]bool) for _, h := range handlers { if h.requestType == "" { t.Error("handler requestType cannot be empty") } if h.register == nil { t.Errorf("handler for %q has nil register function", h.requestType) } if seenTypes[h.requestType] { t.Errorf("duplicate handler registration for requestType %q", h.requestType) } seenTypes[h.requestType] = true if h.requestType == "alt.v1.HelloRequest" { hasHello = true } } if !hasHello { t.Error("expected alt.v1.HelloRequest to be registered in session handlers") } } func TestRegisterHandlers(t *testing.T) { t.Run("skip nil register", func(t *testing.T) { handlers := []sessionHandler{ { requestType: "test.DummyRequest", register: nil, }, } registerHandlers(nil, handlers) }) } func TestWorkerHelloCapabilitiesReflectDeps(t *testing.T) { tests := []struct { name string deps Deps expected []string }{ { name: "no deps", deps: Deps{}, expected: []string{"hello"}, }, { name: "market deps", deps: Deps{ Instruments: &fakeInstrumentStore{}, Bars: &fakeBarStore{}, }, expected: []string{"hello", "market-read"}, }, { name: "backtest read deps", deps: Deps{ Analysis: &fakeAnalysisStore{}, Results: &fakeResultStore{}, }, expected: []string{"hello", "backtest-read"}, }, { name: "starter deps", deps: Deps{ Starter: &fakeStarter{}, }, expected: []string{"hello", "backtest-start", "worker-execution"}, }, { name: "all deps", deps: Deps{ Instruments: &fakeInstrumentStore{}, Bars: &fakeBarStore{}, Analysis: &fakeAnalysisStore{}, Results: &fakeResultStore{}, Starter: &fakeStarter{}, }, expected: []string{"hello", "market-read", "backtest-read", "backtest-start", "worker-execution"}, }, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { caps := capabilitiesForDeps(tt.deps) if len(caps) != len(tt.expected) { t.Fatalf("expected caps %v, got %v", tt.expected, caps) } for i, v := range tt.expected { if caps[i] != v { t.Errorf("expected cap at %d to be %q, got %q", i, v, caps[i]) } } }) } } func TestWorkerCapabilitiesDoNotClaimUnavailableStart(t *testing.T) { caps := capabilitiesForDeps(Deps{}) for _, cap := range caps { if cap == "backtest-start" || cap == "worker-execution" { t.Errorf("capabilities for empty Deps{} should not claim unavailable start, but got %q", cap) } } }