iop/apps/node/internal/node/provider_tunnel_test.go

508 lines
15 KiB
Go

package node_test
import (
"context"
"errors"
"net"
"strings"
"sync/atomic"
"testing"
"time"
toki "git.toki-labs.com/toki/proto-socket/go"
"go.uber.org/zap"
"google.golang.org/protobuf/proto"
"iop/apps/node/internal/node"
"iop/apps/node/internal/runtime"
"iop/apps/node/internal/transport"
iop "iop/proto/gen/iop"
)
// --- tunnel test doubles ---
// mockTunnelAdapter extends countingAdapter and adds TunnelProvider.
type mockTunnelAdapter struct {
countingAdapter
t *testing.T
expectedReq runtime.ProviderTunnelRequest
respondErr error
}
func (a *mockTunnelAdapter) Name() string { return "openai_compat" }
func (a *mockTunnelAdapter) Capabilities(_ context.Context) (runtime.Capabilities, error) {
return runtime.Capabilities{AdapterName: "openai_compat", Targets: []string{"qwen"}}, nil
}
func (a *mockTunnelAdapter) TunnelProvider(ctx context.Context, req runtime.ProviderTunnelRequest, sink runtime.ProviderTunnelSink) error {
if req.RunID != a.expectedReq.RunID || req.TunnelID != a.expectedReq.TunnelID {
a.t.Errorf("unexpected tunnel req: %+v", req)
}
if a.respondErr != nil {
return a.respondErr
}
err := sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
RunID: req.RunID,
TunnelID: req.TunnelID,
Sequence: 0,
Kind: runtime.ProviderTunnelFrameKindResponseStart,
StatusCode: 200,
Headers: map[string]string{"Content-Type": "application/json"},
Timestamp: time.Now(),
})
if err != nil {
return err
}
err = sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
RunID: req.RunID,
TunnelID: req.TunnelID,
Sequence: 1,
Kind: runtime.ProviderTunnelFrameKindBody,
Body: []byte(`{"choices":[{"delta":{"content":"ok"}}]}`),
Timestamp: time.Now(),
})
if err != nil {
return err
}
err = sink.EmitTunnelFrame(ctx, runtime.ProviderTunnelFrame{
RunID: req.RunID,
TunnelID: req.TunnelID,
Sequence: 2,
Kind: runtime.ProviderTunnelFrameKindEnd,
End: true,
Timestamp: time.Now(),
})
return err
}
// cancelAwareTunnelAdapter blocks until its context is cancelled.
type cancelAwareTunnelAdapter struct {
countingAdapter
started chan struct{}
observedCancel chan struct{}
}
func (a *cancelAwareTunnelAdapter) Name() string { return "openai_compat" }
func (a *cancelAwareTunnelAdapter) Capabilities(_ context.Context) (runtime.Capabilities, error) {
return runtime.Capabilities{AdapterName: "openai_compat", Targets: []string{"qwen"}}, nil
}
func (a *cancelAwareTunnelAdapter) TunnelProvider(ctx context.Context, _ runtime.ProviderTunnelRequest, _ runtime.ProviderTunnelSink) error {
close(a.started)
<-ctx.Done()
close(a.observedCancel)
return ctx.Err()
}
type capacityGuardTunnelAdapter struct {
started chan string
release chan struct{}
tunnelCalls int32
executeCalls int32
}
func newCapacityGuardTunnelAdapter() *capacityGuardTunnelAdapter {
return &capacityGuardTunnelAdapter{
started: make(chan string, 4),
release: make(chan struct{}),
}
}
func (a *capacityGuardTunnelAdapter) Name() string { return "openai_compat" }
func (a *capacityGuardTunnelAdapter) Capabilities(_ context.Context) (runtime.Capabilities, error) {
return runtime.Capabilities{
AdapterName: "openai_compat",
Targets: []string{"qwen"},
MaxConcurrency: 1,
}, nil
}
func (a *capacityGuardTunnelAdapter) Execute(ctx context.Context, spec runtime.ExecutionSpec, _ runtime.EventSink) error {
atomic.AddInt32(&a.executeCalls, 1)
a.started <- spec.RunID
select {
case <-a.release:
return nil
case <-ctx.Done():
return runtime.ErrRunCancelled
}
}
func (a *capacityGuardTunnelAdapter) TunnelProvider(ctx context.Context, req runtime.ProviderTunnelRequest, _ runtime.ProviderTunnelSink) error {
atomic.AddInt32(&a.tunnelCalls, 1)
a.started <- req.RunID
select {
case <-a.release:
return nil
case <-ctx.Done():
return ctx.Err()
}
}
// buildSessionTestPipeForNode creates a net.Pipe used as the transport.Session
// transport layer for tunnel tests. The edge side is returned as a TcpClient
// that the test can observe emitted frames from.
func buildSessionTestPipeForNode(t *testing.T) (edgeSide *toki.TcpClient, sess *transport.Session) {
t.Helper()
edgeConn, nodeConn := net.Pipe()
edgeParserMap := toki.ParserMap{
toki.TypeNameOf(&iop.ProviderTunnelFrame{}): func(b []byte) (proto.Message, error) {
m := &iop.ProviderTunnelFrame{}
return m, proto.Unmarshal(b, m)
},
}
nodeParserMap := toki.ParserMap{}
edgeSide = toki.NewTcpClient(edgeConn, 0, 0, edgeParserMap)
nodeSide := toki.NewTcpClient(nodeConn, 0, 0, nodeParserMap)
t.Cleanup(func() { edgeSide.Close(); nodeSide.Close() })
sess = transport.ExportNewSession(nodeSide, zap.NewNop(), "node-id-1", "alias-1")
return edgeSide, sess
}
// --- tunnel tests ---
func TestNodeOnProviderTunnelRequest_Success(t *testing.T) {
mta := &mockTunnelAdapter{
t: t,
expectedReq: runtime.ProviderTunnelRequest{
RunID: "run-tunnel-1",
TunnelID: "tunnel-1",
},
}
router := &fixedRouter{adapterName: "openai_compat", adapters: make(map[string]runtime.Adapter)}
router.adapters["openai_compat"] = mta
n, _ := makeNode(t, router)
req := &iop.ProviderTunnelRequest{
RunId: "run-tunnel-1",
TunnelId: "tunnel-1",
Adapter: "openai_compat",
Target: "qwen",
Method: "POST",
Path: "/v1/chat/completions",
}
err := n.OnProviderTunnelRequest(context.Background(), nil, req)
if err != nil {
t.Fatalf("OnProviderTunnelRequest failed: %v", err)
}
}
func TestNodeOnProviderTunnelRequest_SharedAdapterCapacityRejectsSecondTunnel(t *testing.T) {
adapter := newCapacityGuardTunnelAdapter()
router := &fixedRouter{
adapterName: "openai_compat",
adapters: map[string]runtime.Adapter{"openai_compat": adapter},
}
n, _ := makeNode(t, router)
firstErr := make(chan error, 1)
go func() {
firstErr <- n.OnProviderTunnelRequest(context.Background(), nil, &iop.ProviderTunnelRequest{
RunId: "run-tunnel-capacity-1",
TunnelId: "tunnel-capacity-1",
Adapter: "openai_compat",
Target: "qwen",
})
}()
select {
case got := <-adapter.started:
if got != "run-tunnel-capacity-1" {
t.Fatalf("first started run = %q", got)
}
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for first tunnel to start")
}
edgeSide, sess := buildSessionTestPipeForNode(t)
frameCh := make(chan *iop.ProviderTunnelFrame, 1)
toki.AddListenerTyped[*iop.ProviderTunnelFrame](&edgeSide.Communicator, func(tf *iop.ProviderTunnelFrame) {
frameCh <- tf
})
err := n.OnProviderTunnelRequest(context.Background(), sess, &iop.ProviderTunnelRequest{
RunId: "run-tunnel-capacity-2",
TunnelId: "tunnel-capacity-2",
Adapter: "openai_compat",
Target: "qwen",
})
if !errors.Is(err, node.ErrConcurrencyLimitExceeded) {
t.Fatalf("second tunnel error = %v, want ErrConcurrencyLimitExceeded", err)
}
if got := atomic.LoadInt32(&adapter.tunnelCalls); got != 1 {
t.Fatalf("upstream tunnel calls = %d, want 1", got)
}
select {
case frame := <-frameCh:
if frame.GetKind() != iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR {
t.Fatalf("frame kind = %v, want ERROR", frame.GetKind())
}
if !strings.Contains(frame.GetError(), "concurrency unavailable") {
t.Fatalf("frame error = %q, want concurrency result", frame.GetError())
}
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for concurrency error frame")
}
close(adapter.release)
if err := <-firstErr; err != nil {
t.Fatalf("first tunnel failed: %v", err)
}
}
func TestNodeAdapterCapacityIsSharedByNormalizedAndTunnelExecution(t *testing.T) {
adapter := newCapacityGuardTunnelAdapter()
router := &fixedRouter{
adapterName: "openai_compat",
adapters: map[string]runtime.Adapter{"openai_compat": adapter},
}
n, _ := makeNode(t, router)
tunnelErr := make(chan error, 1)
go func() {
tunnelErr <- n.OnProviderTunnelRequest(context.Background(), nil, &iop.ProviderTunnelRequest{
RunId: "run-shared-tunnel",
TunnelId: "tunnel-shared",
Adapter: "openai_compat",
Target: "qwen",
})
}()
select {
case got := <-adapter.started:
if got != "run-shared-tunnel" {
t.Fatalf("first started run = %q", got)
}
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for tunnel to start")
}
err := n.OnRunRequest(context.Background(), &transport.Session{}, &iop.RunRequest{
RunId: "run-shared-normalized",
Adapter: "openai_compat",
Target: "qwen",
})
if !errors.Is(err, node.ErrConcurrencyLimitExceeded) {
t.Fatalf("normalized run error = %v, want ErrConcurrencyLimitExceeded", err)
}
if got := atomic.LoadInt32(&adapter.executeCalls); got != 0 {
t.Fatalf("normalized upstream execute calls = %d, want 0", got)
}
close(adapter.release)
if err := <-tunnelErr; err != nil {
t.Fatalf("tunnel failed: %v", err)
}
}
func TestNodeOnProviderTunnelRequest_CancelRequestCancelsProviderContext(t *testing.T) {
adapter := &cancelAwareTunnelAdapter{
started: make(chan struct{}),
observedCancel: make(chan struct{}),
}
router := &fixedRouter{adapterName: "openai_compat", adapters: map[string]runtime.Adapter{"openai_compat": adapter}}
n, _ := makeNode(t, router)
req := &iop.ProviderTunnelRequest{
RunId: "run-tunnel-cancel",
TunnelId: "tunnel-cancel",
Adapter: "openai_compat",
Target: "qwen",
Method: "POST",
Path: "/v1/chat/completions",
TimeoutSec: 30,
}
errCh := make(chan error, 1)
go func() {
errCh <- n.OnProviderTunnelRequest(context.Background(), nil, req)
}()
select {
case <-adapter.started:
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for tunnel adapter to start")
}
err := n.OnCancel(context.Background(), nil, &iop.CancelRequest{
RunId: "run-tunnel-cancel",
Action: iop.CancelAction_CANCEL_ACTION_CANCEL_RUN,
})
if err != nil {
t.Fatalf("OnCancel failed: %v", err)
}
select {
case <-adapter.observedCancel:
case <-time.After(2 * time.Second):
t.Fatal("tunnel adapter did not observe cancel request")
}
select {
case err := <-errCh:
if !errors.Is(err, context.Canceled) {
t.Fatalf("OnProviderTunnelRequest error = %v, want context.Canceled", err)
}
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for tunnel request to finish after cancel")
}
}
func TestNodeOnProviderTunnelRequest_LookupFailure(t *testing.T) {
router := &fixedRouter{
adapterName: "nonexistent",
adapters: make(map[string]runtime.Adapter),
lookupErrors: map[string]error{
"nonexistent": errors.New("adapter lookup error"),
},
}
n, _ := makeNode(t, router)
edgeSide, sess := buildSessionTestPipeForNode(t)
req := &iop.ProviderTunnelRequest{
RunId: "run-tunnel-1",
TunnelId: "tunnel-1",
Adapter: "nonexistent",
Target: "qwen",
}
frameCh := make(chan *iop.ProviderTunnelFrame, 10)
toki.AddListenerTyped[*iop.ProviderTunnelFrame](&edgeSide.Communicator, func(tf *iop.ProviderTunnelFrame) {
frameCh <- tf
})
err := n.OnProviderTunnelRequest(context.Background(), sess, req)
if err == nil {
t.Fatal("expected error, got nil")
}
select {
case frame := <-frameCh:
if frame.GetRunId() != "run-tunnel-1" || frame.GetTunnelId() != "tunnel-1" {
t.Errorf("unexpected IDs in frame: run_id=%q tunnel_id=%q", frame.GetRunId(), frame.GetTunnelId())
}
if frame.GetSequence() != 0 {
t.Errorf("expected sequence 0, got %d", frame.GetSequence())
}
if frame.GetKind() != iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR {
t.Errorf("expected ERROR frame, got %v", frame.GetKind())
}
if !strings.Contains(frame.GetError(), "adapter lookup error") {
t.Errorf("expected error message containing 'adapter lookup error', got %q", frame.GetError())
}
if frame.GetNodeId() != "node-id-1" || frame.GetNodeAlias() != "alias-1" {
t.Errorf("unexpected node ID or alias: node_id=%q alias=%q", frame.GetNodeId(), frame.GetNodeAlias())
}
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for ERROR frame")
}
// Verify no more frames
select {
case frame := <-frameCh:
t.Fatalf("unexpected duplicate frame received: %+v", frame)
default:
}
}
func TestNodeOnProviderTunnelRequest_UnsupportedAdapter(t *testing.T) {
// countingAdapter does not implement ProviderTunnelAdapter
mta := &countingAdapter{}
router := &fixedRouter{adapterName: "test", adapters: make(map[string]runtime.Adapter)}
router.adapters["test"] = mta
n, _ := makeNode(t, router)
edgeSide, sess := buildSessionTestPipeForNode(t)
req := &iop.ProviderTunnelRequest{
RunId: "run-tunnel-1",
TunnelId: "tunnel-1",
Adapter: "test",
Target: "qwen",
}
frameCh := make(chan *iop.ProviderTunnelFrame, 10)
toki.AddListenerTyped[*iop.ProviderTunnelFrame](&edgeSide.Communicator, func(tf *iop.ProviderTunnelFrame) {
frameCh <- tf
})
err := n.OnProviderTunnelRequest(context.Background(), sess, req)
if err == nil {
t.Fatal("expected error, got nil")
}
select {
case frame := <-frameCh:
if frame.GetRunId() != "run-tunnel-1" || frame.GetTunnelId() != "tunnel-1" {
t.Errorf("unexpected IDs in frame: run_id=%q tunnel_id=%q", frame.GetRunId(), frame.GetTunnelId())
}
if frame.GetSequence() != 0 {
t.Errorf("expected sequence 0, got %d", frame.GetSequence())
}
if frame.GetKind() != iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR {
t.Errorf("expected ERROR frame, got %v", frame.GetKind())
}
if !strings.Contains(frame.GetError(), "does not support tunneling") {
t.Errorf("expected error message containing 'does not support tunneling', got %q", frame.GetError())
}
case <-time.After(2 * time.Second):
t.Fatal("timeout waiting for ERROR frame")
}
// Verify no more frames
select {
case frame := <-frameCh:
t.Fatalf("unexpected duplicate frame received: %+v", frame)
default:
}
}
func TestNodeOnProviderTunnelRequest_AdapterErrorNoDuplicate(t *testing.T) {
// Adapter directly returns error
mta := &mockTunnelAdapter{
t: t,
expectedReq: runtime.ProviderTunnelRequest{
RunID: "run-tunnel-1",
TunnelID: "tunnel-1",
},
respondErr: errors.New("adapter runtime error"),
}
router := &fixedRouter{adapterName: "openai_compat", adapters: make(map[string]runtime.Adapter)}
router.adapters["openai_compat"] = mta
n, _ := makeNode(t, router)
edgeSide, sess := buildSessionTestPipeForNode(t)
req := &iop.ProviderTunnelRequest{
RunId: "run-tunnel-1",
TunnelId: "tunnel-1",
Adapter: "openai_compat",
Target: "qwen",
}
frameCh := make(chan *iop.ProviderTunnelFrame, 10)
toki.AddListenerTyped[*iop.ProviderTunnelFrame](&edgeSide.Communicator, func(tf *iop.ProviderTunnelFrame) {
frameCh <- tf
})
err := n.OnProviderTunnelRequest(context.Background(), sess, req)
if err == nil {
t.Fatal("expected error, got nil")
}
// The mockTunnelAdapter inside does not emit frames if respondErr is set, it just returns errors.
// Since we removed sendTunnelError for adapter failures from node.go, no frames should be received at all.
// Wait a brief moment to ensure no frames were sent.
select {
case frame := <-frameCh:
t.Fatalf("unexpected frame received: %+v", frame)
case <-time.After(100 * time.Millisecond):
}
}