package node_test import ( "context" "crypto/ecdh" "crypto/ed25519" "crypto/rand" "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/transport" "iop/packages/go/credentiallease" runtime "iop/packages/go/execution" 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 } type credentialTunnelAdapter struct { countingAdapter calls int32 gotHeader string gotScheme string gotSecret string } func (a *credentialTunnelAdapter) Name() string { return "openai_compat" } func (a *credentialTunnelAdapter) Capabilities(context.Context) (runtime.Capabilities, error) { return runtime.Capabilities{AdapterName: "openai_compat", Targets: []string{"qwen"}}, nil } func (a *credentialTunnelAdapter) TunnelProvider(_ context.Context, req runtime.ProviderTunnelRequest, _ runtime.ProviderTunnelSink) error { atomic.AddInt32(&a.calls, 1) if req.Credential != nil { a.gotHeader, a.gotScheme, a.gotSecret = req.Credential.HeaderName, req.Credential.Scheme, string(req.Credential.Secret) } return nil } 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.Provider)} 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 TestNodeConsumesExactCredentialLeaseOnceAtAdapterAdmission(t *testing.T) { now := time.Unix(1700000000, 0).UTC() issuerPublic, issuerPrivate, err := ed25519.GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } recipient, err := ecdh.X25519().GenerateKey(rand.Reader) if err != nil { t.Fatal(err) } scope := credentiallease.Scope{ LeaseID: "lease-node-1", PrincipalRef: "principal-1", CredentialSlotRef: "slot-1", RouteID: "route-1", ProfileID: "openai", UpstreamTarget: "qwen", NodeID: "test-node", RecipientKeyID: "recipient-1", HeaderName: "Authorization", Scheme: "Bearer", CredentialRevision: 3, RouteRevision: 4, ProjectionGeneration: 5, IssuedAtUnixNano: now.UnixNano(), ExpiresAtUnixNano: now.Add(30 * time.Second).UnixNano(), } envelope, err := credentiallease.Issue(scope, []byte("node-secret-sentinel"), recipient.PublicKey().Bytes(), "issuer-1", issuerPrivate, rand.Reader) if err != nil { t.Fatal(err) } consumer, err := credentiallease.NewConsumer("test-node", "recipient-1", recipient.Bytes(), "issuer-1", issuerPublic, 8, func() time.Time { return now }) if err != nil { t.Fatal(err) } adapter := &credentialTunnelAdapter{} router := &fixedRouter{adapterName: "openai_compat", adapters: map[string]runtime.Provider{"openai_compat": adapter}} n, _ := makeNode(t, router) n.SetCredentialConsumer(consumer) binding := &iop.CredentialLeaseBinding{ PrincipalRef: scope.PrincipalRef, CredentialSlotRef: scope.CredentialSlotRef, RouteId: scope.RouteID, ProfileId: scope.ProfileID, UpstreamTarget: scope.UpstreamTarget, NodeId: scope.NodeID, RecipientKeyId: scope.RecipientKeyID, CredentialRevision: scope.CredentialRevision, RouteRevision: scope.RouteRevision, ProjectionGeneration: scope.ProjectionGeneration, } req := &iop.ProviderTunnelRequest{RunId: "run-lease", TunnelId: "tunnel-lease", Adapter: "openai_compat", Target: "qwen", CredentialLease: envelope.ToProto(), CredentialBinding: binding} if err := n.OnProviderTunnelRequest(context.Background(), nil, req); err != nil { t.Fatal(err) } if adapter.calls != 1 || adapter.gotHeader != "Authorization" || adapter.gotScheme != "Bearer" || adapter.gotSecret != "node-secret-sentinel" { t.Fatalf("adapter observation calls=%d header=%q scheme=%q secret=%q", adapter.calls, adapter.gotHeader, adapter.gotScheme, adapter.gotSecret) } if err := n.OnProviderTunnelRequest(context.Background(), nil, req); err == nil || adapter.calls != 1 { t.Fatalf("replay error=%v adapter_calls=%d", err, adapter.calls) } } func TestNodeOnProviderTunnelRequest_SharedAdapterCapacityRejectsSecondTunnel(t *testing.T) { adapter := newCapacityGuardTunnelAdapter() router := &fixedRouter{ adapterName: "openai_compat", adapters: map[string]runtime.Provider{"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.Provider{"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.Provider{"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"}) 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.Provider), 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.Provider)} 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.Provider)} 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): } }