355 lines
8.3 KiB
Go
355 lines
8.3 KiB
Go
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")
|
|
}
|
|
}
|