package bridge import ( "context" "encoding/json" "errors" "fmt" "io" "os" "os/exec" "strings" "testing" "time" ) type bridgeHarness struct { bridge *Bridge requestDecoder *json.Decoder responseEncoder *json.Encoder requestReader *io.PipeReader responseWriter *io.PipeWriter } func newBridgeHarness(t *testing.T) *bridgeHarness { t.Helper() requestReader, requestWriter := io.Pipe() responseReader, responseWriter := io.Pipe() b := NewBridge("unused.py") b.started = true b.stdin = requestWriter b.protocol = responseReader b.encoder = json.NewEncoder(requestWriter) b.decoder = json.NewDecoder(responseReader) h := &bridgeHarness{ bridge: b, requestDecoder: json.NewDecoder(requestReader), responseEncoder: json.NewEncoder(responseWriter), requestReader: requestReader, responseWriter: responseWriter, } go b.readProtocol() t.Cleanup(func() { responseWriter.Close() requestReader.Close() b.finish(nil) }) return h } func testContext(t *testing.T) context.Context { t.Helper() ctx, cancel := context.WithTimeout(context.Background(), time.Second) t.Cleanup(cancel) return ctx } func (h *bridgeHarness) receiveRequest(t *testing.T) Message { t.Helper() var request Message if err := h.requestDecoder.Decode(&request); err != nil { t.Fatalf("decode request: %v", err) } return request } func (h *bridgeHarness) sendMessage(t *testing.T, msg Message) { t.Helper() if err := h.responseEncoder.Encode(msg); err != nil { t.Fatalf("encode runtime message: %v", err) } } func TestRequestCorrelatesResponse(t *testing.T) { h := newBridgeHarness(t) type requestResult struct { message Message err error } resultChannel := make(chan requestResult, 1) ctx := testContext(t) go func() { msg, err := h.bridge.request(ctx, "locals/get", nil) resultChannel <- requestResult{message: msg, err: err} }() request := h.receiveRequest(t) if request.Type != Request || request.Name != "locals/get" || request.RequestId == nil { t.Fatalf("unexpected request: %+v", request) } response, err := newMessage(Response, request.RequestId, request.Name, LocalsResponse{ Variables: map[string]string{"answer": "42"}, }) if err != nil { t.Fatalf("newMessage() error = %v", err) } h.sendMessage(t, response) result := <-resultChannel if result.err != nil { t.Fatalf("request() error = %v", result.err) } var body LocalsResponse if err := json.Unmarshal(result.message.Body, &body); err != nil { t.Fatalf("decode response body: %v", err) } if body.Variables["answer"] != "42" { t.Fatalf("response variables = %#v", body.Variables) } } func TestConcurrentRequestsAreCorrelatedByID(t *testing.T) { h := newBridgeHarness(t) ctx := testContext(t) type requestResult struct { name string body string err error } const requestCount = 8 results := make(chan requestResult, requestCount) for i := range requestCount { name := fmt.Sprintf("test/request-%d", i) go func() { msg, err := h.bridge.request(ctx, name, nil) if err != nil { results <- requestResult{name: name, err: err} return } var body struct { Name string `json:"name"` } err = json.Unmarshal(msg.Body, &body) results <- requestResult{name: name, body: body.Name, err: err} }() } requests := make([]Message, 0, requestCount) ids := make(map[uint64]struct{}, requestCount) for range requestCount { request := h.receiveRequest(t) if request.RequestId == nil { t.Fatalf("request has no ID: %+v", request) } if _, exists := ids[*request.RequestId]; exists { t.Fatalf("duplicate request ID %d", *request.RequestId) } ids[*request.RequestId] = struct{}{} requests = append(requests, request) } for i := len(requests) - 1; i >= 0; i-- { request := requests[i] response, err := newMessage(Response, request.RequestId, request.Name, struct { Name string `json:"name"` }{Name: request.Name}) if err != nil { t.Fatalf("newMessage() error = %v", err) } h.sendMessage(t, response) } for range requestCount { result := <-results if result.err != nil { t.Fatalf("request %q error = %v", result.name, result.err) } if result.body != result.name { t.Fatalf("request %q received body for %q", result.name, result.body) } } } func TestRequestReturnsProtocolError(t *testing.T) { h := newBridgeHarness(t) errChannel := make(chan error, 1) ctx := testContext(t) go func() { _, err := h.bridge.request(ctx, "execution/continue", nil) errChannel <- err }() request := h.receiveRequest(t) h.sendMessage(t, newErrorResponse( *request.RequestId, request.Name, "not_paused", "execution is not paused", )) err := <-errChannel var protocolErr *ProtocolError if !errors.As(err, &protocolErr) { t.Fatalf("request() error = %v, want ProtocolError", err) } if protocolErr.Code != "not_paused" { t.Fatalf("protocol error code = %q", protocolErr.Code) } } func TestRequestRejectsMismatchedResponseName(t *testing.T) { h := newBridgeHarness(t) errChannel := make(chan error, 1) ctx := testContext(t) go func() { _, err := h.bridge.request(ctx, "locals/get", nil) errChannel <- err }() request := h.receiveRequest(t) response, err := newMessage(Response, request.RequestId, "execution/continue", nil) if err != nil { t.Fatalf("newMessage() error = %v", err) } h.sendMessage(t, response) err = <-errChannel if err == nil || !strings.Contains(err.Error(), "response name mismatch") { t.Fatalf("request() error = %v, want response name mismatch", err) } } func TestCancelledRequestIsRemovedFromRegistry(t *testing.T) { h := newBridgeHarness(t) ctx, cancel := context.WithCancel(context.Background()) errChannel := make(chan error, 1) go func() { _, err := h.bridge.request(ctx, "locals/get", nil) errChannel <- err }() request := h.receiveRequest(t) cancel() if err := <-errChannel; !errors.Is(err, context.Canceled) { t.Fatalf("request() error = %v, want context.Canceled", err) } h.bridge.pendingMu.Lock() _, exists := h.bridge.pending[*request.RequestId] h.bridge.pendingMu.Unlock() if exists { t.Fatalf("request ID %d remained in pending registry", *request.RequestId) } } func TestReadProtocolPublishesEvent(t *testing.T) { h := newBridgeHarness(t) event, err := newMessage(Event, nil, "execution/stopped", ExecutionStoppedEvent{ Reason: "step", File: "test.py", Line: 12, }) if err != nil { t.Fatalf("newMessage() error = %v", err) } h.sendMessage(t, event) select { case received := <-h.bridge.Events(): if received.Name != "execution/stopped" { t.Fatalf("event name = %q", received.Name) } var body ExecutionStoppedEvent if err := json.Unmarshal(received.Body, &body); err != nil { t.Fatalf("decode event body: %v", err) } if body.Reason != "step" || body.File != "test.py" || body.Line != 12 { t.Fatalf("event body = %+v", body) } case <-time.After(time.Second): t.Fatal("timed out waiting for event") } } func TestReadOutputPublishesStreamAndText(t *testing.T) { b := NewBridge("unused.py") reader := io.NopCloser(strings.NewReader("hello from target\n")) go b.readOutput(reader, StdoutStream) select { case output := <-b.Output(): if output.Stream != StdoutStream { t.Fatalf("output stream = %q", output.Stream) } if output.Text != "hello from target\n" { t.Fatalf("output text = %q", output.Text) } case <-time.After(time.Second): t.Fatal("timed out waiting for output") } } func TestRequestBeforeStartReturnsErrNotRunning(t *testing.T) { b := NewBridge("unused.py") _, err := b.request(testContext(t), "locals/get", nil) if !errors.Is(err, ErrNotRunning) { t.Fatalf("request() error = %v, want ErrNotRunning", err) } } func TestCloseTerminatesRunningProcess(t *testing.T) { if os.Getenv("PYBUG_BRIDGE_HELPER_PROCESS") == "1" { time.Sleep(time.Minute) return } cmd := exec.Command(os.Args[0], "-test.run=TestCloseTerminatesRunningProcess") cmd.Env = append(os.Environ(), "PYBUG_BRIDGE_HELPER_PROCESS=1") if err := cmd.Start(); err != nil { t.Fatalf("start helper process: %v", err) } b := NewBridge("unused.py") b.cmd = cmd b.started = true go b.waitForProcess() ctx, cancel := context.WithTimeout(context.Background(), time.Second) defer cancel() if err := b.Close(ctx); err != nil { t.Fatalf("Close() error = %v", err) } select { case <-b.Done(): default: t.Fatal("Done channel was not closed") } }