508 lines
15 KiB
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):
|
|
}
|
|
}
|