package protosocket import ( "context" "io" "log/slog" "net/http/httptest" "strings" "testing" "time" toki "git.toki-labs.com/toki/proto-socket/go" "google.golang.org/protobuf/types/known/structpb" "nhooyr.io/websocket" ) func TestServerAcceptsWebSocketAndClosesClients(t *testing.T) { logger := slog.New(slog.NewTextHandler(io.Discard, nil)) srv := NewServer(Config{ HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, }, logger) testServer := httptest.NewServer(srv) defer testServer.Close() wsURL := strings.Replace(testServer.URL, "http", "ws", 1) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() conn, _, err := websocket.Dial(ctx, wsURL, nil) if err != nil { t.Fatalf("failed to dial websocket: %v", err) } defer conn.Close(websocket.StatusNormalClosure, "") client := toki.NewWsClient(conn, 0, 0, ParserMap()) defer client.Close() if !client.IsAlive() { t.Error("client should be alive after connection") } disconnected := make(chan struct{}, 1) client.AddDisconnectListener(func(c *toki.WsClient) { disconnected <- struct{}{} }) if err := srv.Close(); err != nil { t.Errorf("failed to close server: %v", err) } select { case <-disconnected: // success case <-time.After(1 * time.Second): t.Error("timed out waiting for client to disconnect after server close") } } func TestServerRespondsWithErrorEnvelopeForUnsupportedAction(t *testing.T) { logger := slog.New(slog.NewTextHandler(io.Discard, nil)) srv := NewServer(Config{ HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, }, logger) defer srv.Close() testServer := httptest.NewServer(srv) defer testServer.Close() wsURL := strings.Replace(testServer.URL, "http", "ws", 1) ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) defer cancel() conn, _, err := websocket.Dial(ctx, wsURL, nil) if err != nil { t.Fatalf("failed to dial websocket: %v", err) } defer conn.Close(websocket.StatusNormalClosure, "") client := toki.NewWsClient(conn, 0, 0, ParserMap()) defer client.Close() requestEnv := Envelope{ ProtocolVersion: ProtocolVersion, ID: "req-999", Type: "request", Channel: "tasks", Action: "invalid-unsupported-action", } reqStruct, err := requestEnv.ToStruct() if err != nil { t.Fatalf("failed to create request struct: %v", err) } resStruct, err := toki.SendRequestTyped[*structpb.Struct, *structpb.Struct](&client.Communicator, reqStruct, 1*time.Second) if err != nil { t.Fatalf("failed to send request: %v", err) } responseEnv, err := EnvelopeFromStruct(resStruct) if err != nil { t.Fatalf("failed to parse response struct: %v", err) } if responseEnv.Type != "error" { t.Errorf("expected response type 'error', got %q", responseEnv.Type) } if responseEnv.CorrelationID != "req-999" { t.Errorf("expected CorrelationID 'req-999', got %q", responseEnv.CorrelationID) } if responseEnv.Error == nil { t.Fatal("expected error in response env") } if responseEnv.Error.Code != "UNSUPPORTED_ACTION" { t.Errorf("expected error code 'UNSUPPORTED_ACTION', got %q", responseEnv.Error.Code) } if responseEnv.Error.Retryable { t.Error("expected Retryable to be false for UNSUPPORTED_ACTION") } } func TestServerHeartbeatConfigSupport(t *testing.T) { logger := slog.New(slog.NewTextHandler(io.Discard, nil)) srvZero := NewServer(Config{ HeartbeatIntervalSec: 0, HeartbeatWaitSec: 0, }, logger) srvZero.Close() srvNonZero := NewServer(Config{ HeartbeatIntervalSec: 10, HeartbeatWaitSec: 2, }, logger) srvNonZero.Close() }