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/api/internal/config" apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts" "git.toki-labs.com/toki/alt/services/api/internal/workerclient" protoSocket "git.toki-labs.com/toki/proto-socket/go" ) func TestServerRespondsToHelloRequest(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() cfg := config.Config{ Host: "127.0.0.1", Port: freeTCPPort(t), SocketPath: "/socket", HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, } // Test with unavailable worker fakeWorker := &fakeWorkerClient{connectErr: workerclient.ErrUnavailable, isConnected: false} server := NewServerWithWorker(cfg, fakeWorker) if err := server.Start(ctx); err != nil { t.Fatalf("failed to start server: %v", err) } defer server.Stop() client, err := protoSocket.DialWsWithHeartbeat(ctx, cfg.Host, cfg.Port, cfg.SocketPath, 0, 0, apiContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial server: %v", err) } defer client.Close() res, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse]( &client.Communicator, &altv1.HelloRequest{ ClientName: "alt-test", ClientVersion: "test", AltProtocolVersion: "alt.v1", }, 2*time.Second, ) if err != nil { t.Fatalf("failed to send hello request: %v", err) } if res.GetServerName() != serverName { t.Errorf("server name mismatch: expected %q, got %q", serverName, res.GetServerName()) } if res.GetServerVersion() != serverVersion { t.Errorf("server version mismatch: expected %q, got %q", serverVersion, res.GetServerVersion()) } if res.GetAltProtocolVersion() != "alt.v1" { t.Errorf("protocol version mismatch: expected %q, got %q", "alt.v1", res.GetAltProtocolVersion()) } expectedCaps := map[string]bool{ "hello": true, "request-response": true, "market-read": true, "market-import": true, "backtest-read": true, "backtest-start": true, "worker-execution": true, "worker-unavailable": true, } for _, cap := range res.GetCapabilities() { if !expectedCaps[cap] { t.Errorf("unexpected capability: %q", cap) } delete(expectedCaps, cap) } if len(expectedCaps) > 0 { t.Errorf("missing expected capabilities: %v", expectedCaps) } } func TestServerRespondsToHelloRequest_WorkerAvailable(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() cfg := config.Config{ Host: "127.0.0.1", Port: freeTCPPort(t), SocketPath: "/socket", HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, } // Test with available worker fakeWorker := &fakeWorkerClient{isConnected: true} server := NewServerWithWorker(cfg, fakeWorker) if err := server.Start(ctx); err != nil { t.Fatalf("failed to start server: %v", err) } defer server.Stop() client, err := protoSocket.DialWsWithHeartbeat(ctx, cfg.Host, cfg.Port, cfg.SocketPath, 0, 0, apiContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial server: %v", err) } defer client.Close() res, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse]( &client.Communicator, &altv1.HelloRequest{ ClientName: "alt-test", ClientVersion: "test", AltProtocolVersion: "alt.v1", }, 2*time.Second, ) if err != nil { t.Fatalf("failed to send hello request: %v", err) } expectedCaps := map[string]bool{ "hello": true, "request-response": true, "market-read": true, "market-import": true, "backtest-read": true, "backtest-start": true, "worker-execution": true, "worker-available": true, } for _, cap := range res.GetCapabilities() { if !expectedCaps[cap] { t.Errorf("unexpected capability: %q", cap) } delete(expectedCaps, cap) } if len(expectedCaps) > 0 { t.Errorf("missing expected capabilities: %v", expectedCaps) } } func TestTwoServersIsolation(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) defer cancel() cfg1 := config.Config{ Host: "127.0.0.1", Port: freeTCPPort(t), SocketPath: "/socket", HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, } fakeWorker1 := &fakeWorkerClient{isConnected: true} server1 := NewServerWithWorker(cfg1, fakeWorker1) if err := server1.Start(ctx); err != nil { t.Fatalf("failed to start server 1: %v", err) } defer server1.Stop() cfg2 := config.Config{ Host: "127.0.0.1", Port: freeTCPPort(t), SocketPath: "/socket", HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, } fakeWorker2 := &fakeWorkerClient{connectErr: workerclient.ErrUnavailable, isConnected: false} server2 := NewServerWithWorker(cfg2, fakeWorker2) if err := server2.Start(ctx); err != nil { t.Fatalf("failed to start server 2: %v", err) } defer server2.Stop() // Check server 1 capabilities client1, err := protoSocket.DialWsWithHeartbeat(ctx, cfg1.Host, cfg1.Port, cfg1.SocketPath, 0, 0, apiContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial server 1: %v", err) } defer client1.Close() res1, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse]( &client1.Communicator, &altv1.HelloRequest{AltProtocolVersion: "alt.v1"}, 2*time.Second, ) if err != nil { t.Fatalf("failed to send hello request to server 1: %v", err) } hasWorkerAvailable := false for _, cap := range res1.GetCapabilities() { if cap == "worker-available" { hasWorkerAvailable = true } } if !hasWorkerAvailable { t.Error("expected server 1 to advertise worker-available") } // Check server 2 capabilities client2, err := protoSocket.DialWsWithHeartbeat(ctx, cfg2.Host, cfg2.Port, cfg2.SocketPath, 0, 0, apiContracts.ParserMap()) if err != nil { t.Fatalf("failed to dial server 2: %v", err) } defer client2.Close() res2, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse]( &client2.Communicator, &altv1.HelloRequest{AltProtocolVersion: "alt.v1"}, 2*time.Second, ) if err != nil { t.Fatalf("failed to send hello request to server 2: %v", err) } hasWorkerUnavailable := false for _, cap := range res2.GetCapabilities() { if cap == "worker-unavailable" { hasWorkerUnavailable = true } } if !hasWorkerUnavailable { t.Error("expected server 2 to advertise worker-unavailable") } } func TestSessionHandlersHaveUniqueRequestTypes(t *testing.T) { seen := make(map[string]int) for _, handler := range sessionHandlers(nil) { if handler.requestType == "" { t.Errorf("session handler has empty request type") continue } seen[handler.requestType]++ } for requestType, count := range seen { if count > 1 { t.Errorf("request type %q registered %d times; duplicate handlers would panic the communicator", requestType, count) } } } func TestSessionHandlersCoverRequiredRequests(t *testing.T) { registered := make(map[string]bool) for _, handler := range sessionHandlers(nil) { registered[handler.requestType] = true } required := []string{ protoSocket.TypeNameOf(&altv1.HelloRequest{}), } for _, requestType := range required { if !registered[requestType] { t.Errorf("required handler for %q is not registered", requestType) } } } func TestRegisterHandlersSkipsNilRegistrar(t *testing.T) { // registerHandlers must tolerate a malformed registry entry instead of // panicking during connection setup. A nil registrar is skipped without // dereferencing the client. defer func() { if r := recover(); r != nil { t.Fatalf("registerHandlers panicked on nil registrar: %v", r) } }() called := false handlers := []sessionHandler{ {requestType: "alt.v1.NilRegistrarProbe", register: nil}, {requestType: "alt.v1.LiveRegistrarProbe", register: func(*protoSocket.WsClient) { called = true }}, } registerHandlers(nil, handlers) if !called { t.Fatal("expected non-nil registrar to run after nil registrar was skipped") } } func freeTCPPort(t *testing.T) int { t.Helper() listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("failed to reserve TCP port: %v", err) } defer listener.Close() addr, ok := listener.Addr().(*net.TCPAddr) if !ok { t.Fatalf("unexpected listener address type %T", listener.Addr()) } return addr.Port } func TestCapabilitiesAndHandlersSync(t *testing.T) { handlers := sessionHandlers(nil) registered := make(map[string]bool) for _, h := range handlers { registered[h.requestType] = true } caps := capabilitiesForSession(nil) capSet := make(map[string]bool) for _, c := range caps { capSet[c] = true } // 1. Check ListInstruments/ListBars -> market-read hasMarketRequest := registered[protoSocket.TypeNameOf(&altv1.ListInstrumentsRequest{})] && registered[protoSocket.TypeNameOf(&altv1.ListBarsRequest{})] if hasMarketRequest && !capSet["market-read"] { t.Error("market-read capability missing although ListInstruments/ListBars handlers are registered") } // 1b. Check ImportDailyBars -> market-import if registered[protoSocket.TypeNameOf(&altv1.ImportDailyBarsRequest{})] && !capSet["market-import"] { t.Error("market-import capability missing although ImportDailyBars handler is registered") } // 2. Check backtest queries -> backtest-read hasBacktestReadRequest := registered[protoSocket.TypeNameOf(&altv1.ListBacktestRunsRequest{})] && registered[protoSocket.TypeNameOf(&altv1.GetBacktestRunDetailRequest{})] && registered[protoSocket.TypeNameOf(&altv1.GetBacktestResultRequest{})] && registered[protoSocket.TypeNameOf(&altv1.CompareBacktestRunsRequest{})] if hasBacktestReadRequest && !capSet["backtest-read"] { t.Error("backtest-read capability missing although backtest query handlers are registered") } // 3. Check StartBacktest -> backtest-start, worker-execution hasBacktestStartRequest := registered[protoSocket.TypeNameOf(&altv1.StartBacktestRequest{})] if hasBacktestStartRequest { if !capSet["backtest-start"] { t.Error("backtest-start capability missing although StartBacktest handler is registered") } if !capSet["worker-execution"] { t.Error("worker-execution capability missing although StartBacktest handler is registered") } } }