wip, rework bridge

This commit is contained in:
2026-07-13 18:33:39 +02:00
parent 265dbe2580
commit 86039811ce
6 changed files with 1137 additions and 308 deletions

View File

@@ -0,0 +1,354 @@
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")
}
}