package transport import ( "context" "net" "sync" "testing" "time" toki "git.toki-labs.com/toki/proto-socket/go" "go.uber.org/zap" "google.golang.org/protobuf/proto" edgenode "iop/apps/edge/internal/node" "iop/packages/go/config" iop "iop/proto/gen/iop" ) func TestEdgeParserMap_NodeCommandResponse(t *testing.T) { parsers := edgeParserMap() original := &iop.NodeCommandResponse{ Type: iop.NodeCommandType_NODE_COMMAND_TYPE_CAPABILITIES, Adapter: "ollama", Target: "model", SessionId: "default", Result: map[string]string{"status": "available"}, } payload, err := proto.Marshal(original) if err != nil { t.Fatalf("marshal: %v", err) } key := toki.TypeNameOf(original) parser, ok := parsers[key] if !ok { t.Fatalf("parser not found for key: %s", key) } parsed, err := parser(payload) if err != nil { t.Fatalf("parse: %v", err) } got := parsed.(*iop.NodeCommandResponse) if got.GetType() != original.GetType() || got.GetAdapter() != original.GetAdapter() || got.GetTarget() != original.GetTarget() || got.GetSessionId() != original.GetSessionId() || got.GetResult()["status"] != "available" { t.Fatalf("unexpected node command response: %+v", got) } } func TestEdgeParserMap_NodeCommandResponse_NewTypes(t *testing.T) { parsers := edgeParserMap() cases := []struct { cmdType iop.NodeCommandType result map[string]string providerSnapshots []*iop.ProviderSnapshot }{ { cmdType: iop.NodeCommandType_NODE_COMMAND_TYPE_CAPABILITIES, result: map[string]string{"targets": "codex,antigravity"}, providerSnapshots: []*iop.ProviderSnapshot{ { Adapter: "cli", Status: "available", Capacity: 4, InFlight: 2, Queued: 0, }, }, }, {iop.NodeCommandType_NODE_COMMAND_TYPE_TRANSPORT_STATUS, map[string]string{"connected": "true"}, nil}, {iop.NodeCommandType_NODE_COMMAND_TYPE_OLLAMA_API, map[string]string{"status_code": "200"}, nil}, } for _, tc := range cases { original := &iop.NodeCommandResponse{ RequestId: "req-1", Type: tc.cmdType, Adapter: "ollama", Target: "model", SessionId: "default", Result: tc.result, ProviderSnapshots: tc.providerSnapshots, } payload, err := proto.Marshal(original) if err != nil { t.Fatalf("marshal %v: %v", tc.cmdType, err) } parser, ok := parsers[toki.TypeNameOf(original)] if !ok { t.Fatalf("parser not found for NodeCommandResponse") } parsed, err := parser(payload) if err != nil { t.Fatalf("parse %v: %v", tc.cmdType, err) } got := parsed.(*iop.NodeCommandResponse) if got.GetType() != tc.cmdType { t.Fatalf("type: got %v want %v", got.GetType(), tc.cmdType) } for k, v := range tc.result { if got.GetResult()[k] != v { t.Fatalf("result[%q]: got %q want %q", k, got.GetResult()[k], v) } } if len(got.GetProviderSnapshots()) != len(tc.providerSnapshots) { t.Fatalf("provider snapshots count: got %d want %d", len(got.GetProviderSnapshots()), len(tc.providerSnapshots)) } for i, snap := range tc.providerSnapshots { g := got.GetProviderSnapshots()[i] if g.GetAdapter() != snap.GetAdapter() || g.GetStatus() != snap.GetStatus() || g.GetCapacity() != snap.GetCapacity() || g.GetInFlight() != snap.GetInFlight() || g.GetQueued() != snap.GetQueued() { t.Fatalf("provider snapshot mismatch at index %d: %+v vs %+v", i, g, snap) } } } } func TestEdgeParserMap_EdgeNodeEvent(t *testing.T) { parsers := edgeParserMap() original := &iop.EdgeNodeEvent{ Type: "node.disconnected", Source: "edge", NodeId: "node-1", Alias: "local-node", Reason: "transport_closed", } payload, err := proto.Marshal(original) if err != nil { t.Fatalf("marshal: %v", err) } key := toki.TypeNameOf(original) parser, ok := parsers[key] if !ok { t.Fatalf("parser not found for key: %s", key) } parsed, err := parser(payload) if err != nil { t.Fatalf("parse: %v", err) } got := parsed.(*iop.EdgeNodeEvent) if got.GetType() != original.GetType() || got.GetSource() != original.GetSource() || got.GetNodeId() != original.GetNodeId() || got.GetAlias() != original.GetAlias() || got.GetReason() != original.GetReason() { t.Fatalf("unexpected edge node event: %+v", got) } } func TestEdgeParserMap_ProviderTunnelFrameIsSeparateFromRunEvent(t *testing.T) { parsers := edgeParserMap() rawBody := []byte("data: {\"choices\":[{\"delta\":{\"reasoning_content\":\"raw\"}}]}\n\n") original := &iop.ProviderTunnelFrame{ RunId: "run-1", TunnelId: "tunnel-1", Sequence: 7, Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: rawBody, } payload, err := proto.Marshal(original) if err != nil { t.Fatalf("marshal: %v", err) } key := toki.TypeNameOf(original) if key == toki.TypeNameOf(&iop.RunEvent{}) { t.Fatal("provider tunnel frame must not share the RunEvent wire type") } parser, ok := parsers[key] if !ok { t.Fatalf("parser not found for key: %s", key) } parsed, err := parser(payload) if err != nil { t.Fatalf("parse: %v", err) } if _, ok := parsed.(*iop.RunEvent); ok { t.Fatal("provider tunnel frame must not parse as RunEvent") } got := parsed.(*iop.ProviderTunnelFrame) if got.GetRunId() != original.GetRunId() || got.GetTunnelId() != original.GetTunnelId() || got.GetSequence() != original.GetSequence() || got.GetKind() != original.GetKind() || string(got.GetBody()) != string(rawBody) { t.Fatalf("unexpected tunnel frame: %+v", got) } } // TestServerRoutesTunnelFramesToTunnelHandlerNotRunHandler verifies that a // ProviderTunnelFrame received from a node connection reaches only the tunnel // frame handler: the run event handler (events.Bus publisher in production) // must never observe tunnel frames, and run events must not reach the tunnel // handler. func TestServerRoutesTunnelFramesToTunnelHandlerNotRunHandler(t *testing.T) { edgeConn, nodeConn := net.Pipe() defer edgeConn.Close() defer nodeConn.Close() edgeClient := toki.NewTcpClient(edgeConn, 0, 0, edgeParserMap()) nodeClient := toki.NewTcpClient(nodeConn, 0, 0, toki.ParserMap{}) reg := edgenode.NewRegistry() reg.RegisterIfAbsent(&edgenode.NodeEntry{ NodeID: "node-1", Client: edgeClient, }) s := &Server{ registry: reg, logger: zap.NewNop(), } var mu sync.Mutex var tunnelFrames []*iop.ProviderTunnelFrame var runEvents []*iop.RunEvent s.SetTunnelFrameHandler(func(nodeID string, gen uint64, f *iop.ProviderTunnelFrame) { mu.Lock() tunnelFrames = append(tunnelFrames, f) mu.Unlock() }) s.SetRunEventHandler(func(e *iop.RunEvent) { mu.Lock() runEvents = append(runEvents, e) mu.Unlock() }) s.onNodeConnected(edgeClient) rawBody := []byte("data: {\"choices\":[{\"delta\":{\"content\":\"raw\"}}]}\n\n") if err := nodeClient.Send(&iop.ProviderTunnelFrame{ RunId: "run-1", TunnelId: "tunnel-1", Sequence: 1, Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: rawBody, }); err != nil { t.Fatalf("send tunnel frame: %v", err) } if err := nodeClient.Send(&iop.RunEvent{RunId: "run-1", Type: "delta", Delta: "normalized"}); err != nil { t.Fatalf("send run event: %v", err) } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { mu.Lock() done := len(tunnelFrames) == 1 && len(runEvents) == 1 mu.Unlock() if done { break } time.Sleep(10 * time.Millisecond) } mu.Lock() defer mu.Unlock() if len(tunnelFrames) != 1 { t.Fatalf("tunnel handler frames: got %d want 1", len(tunnelFrames)) } if string(tunnelFrames[0].GetBody()) != string(rawBody) { t.Fatalf("tunnel frame body mismatch: %q", tunnelFrames[0].GetBody()) } if len(runEvents) != 1 { t.Fatalf("run handler events: got %d want 1", len(runEvents)) } if runEvents[0].GetDelta() != "normalized" { t.Fatalf("run handler must only see run events, got %+v", runEvents[0]) } } func TestReceptionIdentityFence_RunEventAndTunnel(t *testing.T) { edgeConn1, nodeConn1 := net.Pipe() defer edgeConn1.Close() defer nodeConn1.Close() edgeConn2, nodeConn2 := net.Pipe() defer edgeConn2.Close() defer nodeConn2.Close() edgeClient1 := toki.NewTcpClient(edgeConn1, 0, 0, edgeParserMap()) nodeClient1 := toki.NewTcpClient(nodeConn1, 0, 0, toki.ParserMap{}) edgeClient2 := toki.NewTcpClient(edgeConn2, 0, 0, edgeParserMap()) nodeClient2 := toki.NewTcpClient(nodeConn2, 0, 0, toki.ParserMap{}) registry := edgenode.NewRegistry() s := &Server{ registry: registry, logger: zap.NewNop(), } entry1 := &edgenode.NodeEntry{NodeID: "node-1", Client: edgeClient1} if !registry.RegisterIfAbsent(entry1) { t.Fatal("failed to register client 1") } var mu sync.Mutex type lifecycleCall struct { nodeID string generation uint64 runID string } type tunnelCall struct { nodeID string generation uint64 runID string } var lifecycles []lifecycleCall var tunnels []tunnelCall var observedRunEvents []*iop.RunEvent s.SetRunLifecycleHandler(func(nodeID string, gen uint64, e *iop.RunEvent) { mu.Lock() lifecycles = append(lifecycles, lifecycleCall{nodeID: nodeID, generation: gen, runID: e.GetRunId()}) mu.Unlock() }) s.SetTunnelFrameHandler(func(nodeID string, gen uint64, f *iop.ProviderTunnelFrame) { mu.Lock() tunnels = append(tunnels, tunnelCall{nodeID: nodeID, generation: gen, runID: f.GetRunId()}) mu.Unlock() }) s.SetRunEventHandler(func(e *iop.RunEvent) { mu.Lock() observedRunEvents = append(observedRunEvents, e) mu.Unlock() }) s.onNodeConnected(edgeClient1) s.onNodeConnected(edgeClient2) // Send from Client 1 (current owner, gen 1) with spoofed payload NodeId "spoofed-node" if err := nodeClient1.Send(&iop.RunEvent{RunId: "run-c1", Type: "complete", NodeId: "spoofed-node"}); err != nil { t.Fatalf("send run event client 1: %v", err) } if err := nodeClient1.Send(&iop.ProviderTunnelFrame{RunId: "run-c1", TunnelId: "t1", Sequence: 1, Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, NodeId: "spoofed-node"}); err != nil { t.Fatalf("send tunnel client 1: %v", err) } deadline := time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { mu.Lock() done := len(lifecycles) == 1 && len(tunnels) == 1 && len(observedRunEvents) == 1 mu.Unlock() if done { break } time.Sleep(10 * time.Millisecond) } mu.Lock() if len(lifecycles) != 1 || lifecycles[0].nodeID != "node-1" || lifecycles[0].generation != entry1.ConnectionGeneration { t.Fatalf("lifecycle client 1: got %+v, want node-1 gen %d", lifecycles, entry1.ConnectionGeneration) } if len(tunnels) != 1 || tunnels[0].nodeID != "node-1" || tunnels[0].generation != entry1.ConnectionGeneration { t.Fatalf("tunnel client 1: got %+v, want node-1 gen %d", tunnels, entry1.ConnectionGeneration) } mu.Unlock() // Reconnect: unregister client 1, register client 2 for node-1 registry.UnregisterIfClient("node-1", edgeClient1) entry2 := &edgenode.NodeEntry{NodeID: "node-1", Client: edgeClient2} if !registry.RegisterIfAbsent(entry2) { t.Fatal("failed to register client 2") } // Now client 1 is stale. Send from client 1 again. if err := nodeClient1.Send(&iop.RunEvent{RunId: "run-stale-c1", Type: "complete", NodeId: "node-1"}); err != nil { t.Fatalf("send stale run event client 1: %v", err) } if err := nodeClient1.Send(&iop.ProviderTunnelFrame{RunId: "run-stale-c1", TunnelId: "t2", Sequence: 1, Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, NodeId: "node-1"}); err != nil { t.Fatalf("send stale tunnel client 1: %v", err) } // Send from client 2 (new current owner, gen 2) if err := nodeClient2.Send(&iop.RunEvent{RunId: "run-c2", Type: "complete", NodeId: "node-1"}); err != nil { t.Fatalf("send run event client 2: %v", err) } if err := nodeClient2.Send(&iop.ProviderTunnelFrame{RunId: "run-c2", TunnelId: "t3", Sequence: 1, Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, NodeId: "node-1"}); err != nil { t.Fatalf("send tunnel client 2: %v", err) } deadline = time.Now().Add(2 * time.Second) for time.Now().Before(deadline) { mu.Lock() done := len(lifecycles) == 2 && len(tunnels) == 2 && len(observedRunEvents) == 3 mu.Unlock() if done { break } time.Sleep(10 * time.Millisecond) } mu.Lock() defer mu.Unlock() // Stale client 1 must be dropped from lifecycles & tunnels if len(lifecycles) != 2 { t.Fatalf("expected exactly 2 lifecycles (client 1 active + client 2 active), got %d: %+v", len(lifecycles), lifecycles) } if lifecycles[1].nodeID != "node-1" || lifecycles[1].generation != entry2.ConnectionGeneration || lifecycles[1].runID != "run-c2" { t.Fatalf("lifecycle client 2: got %+v, want run-c2 gen %d", lifecycles[1], entry2.ConnectionGeneration) } if len(tunnels) != 2 { t.Fatalf("expected exactly 2 tunnels, got %d: %+v", len(tunnels), tunnels) } if tunnels[1].nodeID != "node-1" || tunnels[1].generation != entry2.ConnectionGeneration || tunnels[1].runID != "run-c2" { t.Fatalf("tunnel client 2: got %+v, want run-c2 gen %d", tunnels[1], entry2.ConnectionGeneration) } // Observability fanout sees all 3 run events (message-only fanout) if len(observedRunEvents) != 3 { t.Fatalf("expected 3 observed run events, got %d", len(observedRunEvents)) } } func TestServerEnrichesRunEventNodeAlias(t *testing.T) { registry := edgenode.NewRegistry() registry.Register(&edgenode.NodeEntry{NodeID: "node-1", Alias: "alias-1"}) s := &Server{registry: registry} event := &iop.RunEvent{RunId: "run-1", Type: "start", NodeId: "node-1"} s.enrichRunEvent(event) if event.GetNodeAlias() != "alias-1" { t.Fatalf("expected alias-1, got %q", event.GetNodeAlias()) } preset := &iop.RunEvent{RunId: "run-2", Type: "start", NodeId: "node-1", NodeAlias: "preset"} s.enrichRunEvent(preset) if preset.GetNodeAlias() != "preset" { t.Fatalf("expected preset to be preserved, got %q", preset.GetNodeAlias()) } unknown := &iop.RunEvent{RunId: "run-3", Type: "start", NodeId: "node-x"} s.enrichRunEvent(unknown) if unknown.GetNodeAlias() != "" { t.Fatalf("expected empty alias for unknown node, got %q", unknown.GetNodeAlias()) } } // TestEdgeParserMap_NodeConfigRefreshResponse verifies the parser map includes // the NodeConfigRefreshResponse type so Edge can receive Node ack messages. func TestEdgeParserMap_NodeConfigRefreshResponse(t *testing.T) { parsers := edgeParserMap() original := &iop.NodeConfigRefreshResponse{ RequestId: "req-refresh-1", Status: iop.NodeConfigRefreshStatus_NODE_CONFIG_REFRESH_STATUS_APPLIED, } payload, err := proto.Marshal(original) if err != nil { t.Fatalf("marshal: %v", err) } key := toki.TypeNameOf(original) parser, ok := parsers[key] if !ok { t.Fatalf("parser not found for key: %s", key) } parsed, err := parser(payload) if err != nil { t.Fatalf("parse: %v", err) } got := parsed.(*iop.NodeConfigRefreshResponse) if got.GetRequestId() != original.GetRequestId() || got.GetStatus() != original.GetStatus() { t.Fatalf("unexpected response: %+v", got) } } // TestServerPushConfigRefreshSkippedWhenNodeDisconnected verifies that a node // configured in NodeStore but absent from the live registry is recorded as // skipped. This represents the real disconnected-node path: the disconnect // listener removes the entry from the registry, so configured-but-disconnected // nodes never appear in registry.All(). func TestServerPushConfigRefreshSkippedWhenNodeDisconnected(t *testing.T) { nodeStore, err := edgenode.LoadFromConfig([]config.NodeDefinition{ {ID: "node-1", Alias: "alias-1", Token: "tok-1"}, }) if err != nil { t.Fatalf("load node store: %v", err) } // Registry has no entry for node-1 (simulates post-disconnect state). registry := edgenode.NewRegistry() s := &Server{ registry: registry, nodeStore: nodeStore, logger: zap.NewNop(), HeartbeatInterval: 30, HeartbeatWait: 45, } req := &iop.NodeConfigRefreshRequest{RequestId: "req-1"} results := s.PushConfigRefresh(context.Background(), req) if len(results) != 1 { t.Fatalf("expected 1 result, got %d", len(results)) } if results[0].NodeID != "node-1" { t.Fatalf("unexpected node id: %q", results[0].NodeID) } if results[0].Status != "skipped" { t.Fatalf("expected status=skipped for absent registry entry, got %q", results[0].Status) } } // TestServerPushConfigRefreshEmptyNodeStoreProducesNoResults verifies that // PushConfigRefresh with an empty NodeStore returns an empty result slice. func TestServerPushConfigRefreshEmptyNodeStoreProducesNoResults(t *testing.T) { nodeStore, err := edgenode.LoadFromConfig(nil) if err != nil { t.Fatalf("load node store: %v", err) } s := &Server{ registry: edgenode.NewRegistry(), nodeStore: nodeStore, logger: zap.NewNop(), HeartbeatInterval: 30, HeartbeatWait: 45, } results := s.PushConfigRefresh(context.Background(), &iop.NodeConfigRefreshRequest{RequestId: "req-empty"}) if len(results) != 0 { t.Fatalf("expected 0 results for empty node store, got %d", len(results)) } } func TestBuildConfigPayload_OllamaVllmOneof(t *testing.T) { rec := &edgenode.NodeRecord{ Adapters: config.AdaptersConf{ OllamaInstances: []config.OllamaInstanceConf{ {Name: "ollama", Enabled: true, BaseURL: "http://localhost:11434"}, }, VllmInstances: []config.VllmInstanceConf{ {Name: "vllm", Enabled: true, Endpoint: "http://localhost:8000"}, }, }, } payload, err := edgenode.BuildConfigPayload(rec) if err != nil { t.Fatalf("build: %v", err) } var ollama, vllm *iop.AdapterConfig for _, a := range payload.GetAdapters() { switch a.GetType() { case "ollama": ollama = a case "vllm": vllm = a } } if ollama == nil || ollama.GetOllama().GetBaseUrl() != "http://localhost:11434" { t.Fatalf("ollama: %+v", ollama) } if vllm == nil || vllm.GetVllm().GetEndpoint() != "http://localhost:8000" { t.Fatalf("vllm: %+v", vllm) } } func TestBuildConfigPayload_AllAdaptersSettingsNil(t *testing.T) { rec := &edgenode.NodeRecord{ Adapters: config.AdaptersConf{ Ollama: config.OllamaConf{Enabled: true, BaseURL: "http://localhost:11434"}, Vllm: config.VllmConf{Enabled: true, Endpoint: "http://localhost:8000"}, Mock: config.MockConf{Enabled: true}, }, } payload, err := edgenode.BuildConfigPayload(rec) if err != nil { t.Fatalf("build: %v", err) } for _, a := range payload.GetAdapters() { if a.GetSettings() != nil { t.Fatalf("adapter %q must not populate legacy Settings in new Edge-generated payload", a.GetType()) } } var mockFound bool for _, a := range payload.GetAdapters() { if a.GetType() == "mock" { mockFound = true if a.GetMock() == nil { t.Fatal("mock adapter must use typed MockAdapterConfig oneof") } } } if !mockFound { t.Fatal("expected mock adapter in payload") } } func TestEdgeParserMap_ExecutionFailureRoundTrip(t *testing.T) { parsers := edgeParserMap() failure := &iop.ExecutionFailure{ Code: "response_stalled", Message: "provider response stalled", Retryable: true, Metadata: map[string]string{ "failure_code": "response_stalled", "provider_health": "available", "liveness_classification": "request_stalled", "idle_duration_ms": "5000", "run_id": "run-1", "attempt_id": "run-1", "attempt_fence": "confirmed", "adapter": "ollama", "target": "llama3", "health_observation_seq": "1", }, } t.Run("RunEvent with ExecutionFailure", func(t *testing.T) { event := &iop.RunEvent{ RunId: "run-1", Type: "error", Error: "provider response stalled", Failure: failure, NodeId: "node-1", Metadata: failure.Metadata, } data, err := proto.Marshal(event) if err != nil { t.Fatalf("marshal: %v", err) } parser, ok := parsers[toki.TypeNameOf(event)] if !ok { t.Fatalf("parser not found for RunEvent") } parsed, err := parser(data) if err != nil { t.Fatalf("parse: %v", err) } got := parsed.(*iop.RunEvent) if got.GetFailure() == nil { t.Fatal("expected non-nil Failure on parsed RunEvent") } if got.GetFailure().GetCode() != "response_stalled" || !got.GetFailure().GetRetryable() { t.Fatalf("unexpected Failure: %+v", got.GetFailure()) } if got.GetFailure().GetMetadata()["provider_health"] != "available" { t.Fatalf("unexpected metadata: %+v", got.GetFailure().GetMetadata()) } }) t.Run("ProviderTunnelFrame with ExecutionFailure", func(t *testing.T) { frame := &iop.ProviderTunnelFrame{ RunId: "run-1", TunnelId: "tunnel-1", Sequence: 5, Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_ERROR, Error: "provider response stalled", Failure: failure, NodeId: "node-1", Metadata: failure.Metadata, } data, err := proto.Marshal(frame) if err != nil { t.Fatalf("marshal: %v", err) } parser, ok := parsers[toki.TypeNameOf(frame)] if !ok { t.Fatalf("parser not found for ProviderTunnelFrame") } parsed, err := parser(data) if err != nil { t.Fatalf("parse: %v", err) } got := parsed.(*iop.ProviderTunnelFrame) if got.GetFailure() == nil { t.Fatal("expected non-nil Failure on parsed ProviderTunnelFrame") } if got.GetFailure().GetCode() != "response_stalled" || !got.GetFailure().GetRetryable() { t.Fatalf("unexpected Failure: %+v", got.GetFailure()) } if got.GetFailure().GetMetadata()["liveness_classification"] != "request_stalled" { t.Fatalf("unexpected metadata: %+v", got.GetFailure().GetMetadata()) } }) }