package protosocket import ( "bytes" "context" "encoding/json" "io" "log/slog" "net/http/httptest" "strings" "testing" "time" toki "git.toki-labs.com/toki/proto-socket/go" "github.com/nomadcode/nomadcode-core/internal/notification" "google.golang.org/protobuf/types/known/structpb" "nhooyr.io/websocket" ) func requireDiagnosticsMeta(t *testing.T, env Envelope, channel, action, errorCode string) string { t.Helper() if env.Meta == nil { t.Fatal("expected diagnostics meta") } connectionID, ok := env.Meta["connection_id"].(string) if !ok || !strings.HasPrefix(connectionID, "conn-") { t.Fatalf("expected connection_id with conn- prefix, got %+v", env.Meta["connection_id"]) } if got := env.Meta["protocol_version"]; got != env.ProtocolVersion { t.Errorf("expected meta protocol_version %q, got %+v", env.ProtocolVersion, got) } if got := env.Meta["channel"]; got != channel { t.Errorf("expected meta channel %q, got %+v", channel, got) } if got := env.Meta["action"]; got != action { t.Errorf("expected meta action %q, got %+v", action, got) } if got := env.Meta["error_code"]; got != errorCode { t.Errorf("expected meta error_code %q, got %+v", errorCode, got) } if timestamp, ok := env.Meta["timestamp"].(string); !ok || timestamp == "" { t.Errorf("expected non-empty meta timestamp, got %+v", env.Meta["timestamp"]) } return connectionID } func requireJSONLog(t *testing.T, logs string, msg string, attrs map[string]string) { t.Helper() for _, line := range strings.Split(strings.TrimSpace(logs), "\n") { if line == "" { continue } var record map[string]any if err := json.Unmarshal([]byte(line), &record); err != nil { t.Fatalf("failed to parse JSON log line %q: %v", line, err) } if record["msg"] != msg { continue } for key, want := range attrs { if got := record[key]; got != want { t.Fatalf("log %q expected %s=%q, got %+v in record %+v", msg, key, want, got, record) } } if got, ok := record["timestamp"].(string); !ok || got == "" { t.Fatalf("log %q expected non-empty timestamp, got %+v in record %+v", msg, record["timestamp"], record) } return } t.Fatalf("log %q with attrs %+v not found in logs:\n%s", msg, attrs, logs) } func waitForServerClients(t *testing.T, srv *Server, want int) { t.Helper() deadline := time.Now().Add(1 * time.Second) for time.Now().Before(deadline) { srv.mu.Lock() got := len(srv.clients) srv.mu.Unlock() if got == want { return } time.Sleep(10 * time.Millisecond) } srv.mu.Lock() got := len(srv.clients) srv.mu.Unlock() t.Fatalf("expected %d server clients, got %d", want, got) } 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 TestServerAddsDiagnosticsMetaAndStructuredLogs(t *testing.T) { var logs bytes.Buffer logger := slog.New(slog.NewJSONHandler(&logs, 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-diagnostics", Type: "request", Channel: "task", Action: "unsupported.diagnostics", } 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) } connectionID := requireDiagnosticsMeta(t, responseEnv, "task", "unsupported.diagnostics", "UNSUPPORTED_ACTION") requireJSONLog(t, logs.String(), "proto-socket request handled", map[string]string{ "connection_id": connectionID, "protocol_version": ProtocolVersion, "channel": "task", "action": "unsupported.diagnostics", "error_code": "UNSUPPORTED_ACTION", }) } func TestServerBroadcastAddsDiagnosticsMetaAndStructuredLogs(t *testing.T) { var logs bytes.Buffer logger := slog.New(slog.NewJSONHandler(&logs, 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() received := make(chan struct { env Envelope err error }, 1) toki.AddListenerTyped[*structpb.Struct](&client.Communicator, func(msg *structpb.Struct) { env, err := EnvelopeFromStruct(msg) received <- struct { env Envelope err error }{env: env, err: err} }) waitForServerClients(t, srv, 1) err = srv.BroadcastEnvelope(ctx, TaskStatusChangedEnvelope(notification.TaskEvent{ Type: notification.TaskEventRunning, TaskID: "task-diagnostics", Status: "running", Title: "Diagnostics", Message: "running", })) if err != nil { t.Fatalf("broadcast failed: %v", err) } var got struct { env Envelope err error } select { case got = <-received: case <-time.After(1 * time.Second): t.Fatal("timed out waiting for broadcast envelope") } if got.err != nil { t.Fatalf("failed to parse broadcast envelope: %v", got.err) } connectionID := requireDiagnosticsMeta(t, got.env, "task", "task.status.changed", "") requireJSONLog(t, logs.String(), "proto-socket broadcast sent", map[string]string{ "connection_id": connectionID, "protocol_version": ProtocolVersion, "channel": "task", "action": "task.status.changed", "error_code": "", }) } 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() }