package wire import ( "reflect" "testing" "time" proto_socket "git.toki-labs.com/toki/proto-socket/go" iop "iop/proto/gen/iop" "iop/packages/go/config" "google.golang.org/protobuf/proto" "google.golang.org/protobuf/reflect/protoreflect" ) func TestEdgeParserMapEdgeHelloRoundTrip(t *testing.T) { parserMap := EdgeParserMap() typeName := proto_socket.TypeNameOf(&iop.EdgeHelloRequest{}) parse, ok := parserMap[typeName] if !ok { t.Fatalf("missing parser for %s", typeName) } req := &iop.EdgeHelloRequest{ EdgeId: "edge-dgx-group", EdgeName: "DGX Group", Version: "1.2.3", Capabilities: []string{"node-registry", "run-dispatch"}, Metadata: map[string]string{ "region": "local", }, } payload, err := proto.Marshal(req) if err != nil { t.Fatalf("marshal EdgeHelloRequest: %v", err) } got, err := parse(payload) if err != nil { t.Fatalf("parse EdgeHelloRequest: %v", err) } parsed, ok := got.(*iop.EdgeHelloRequest) if !ok { t.Fatalf("expected *iop.EdgeHelloRequest, got %T", got) } if parsed.EdgeId != req.EdgeId { t.Errorf("edge_id: got %q want %q", parsed.EdgeId, req.EdgeId) } if parsed.EdgeName != req.EdgeName { t.Errorf("edge_name: got %q want %q", parsed.EdgeName, req.EdgeName) } if parsed.Version != req.Version { t.Errorf("version: got %q want %q", parsed.Version, req.Version) } if len(parsed.Capabilities) != len(req.Capabilities) { t.Fatalf("capabilities length: got %d want %d", len(parsed.Capabilities), len(req.Capabilities)) } if parsed.Capabilities[0] != req.Capabilities[0] || parsed.Capabilities[1] != req.Capabilities[1] { t.Errorf("capabilities: got %v want %v", parsed.Capabilities, req.Capabilities) } if parsed.Metadata["region"] != "local" { t.Errorf("metadata region: got %q want %q", parsed.Metadata["region"], "local") } } func TestEdgeParserMapEdgeHelloResponseRoundTrip(t *testing.T) { parserMap := EdgeParserMap() typeName := proto_socket.TypeNameOf(&iop.EdgeHelloResponse{}) parse, ok := parserMap[typeName] if !ok { t.Fatalf("missing parser for %s", typeName) } res := &iop.EdgeHelloResponse{ Accepted: true, Protocol: Protocol, ServerTimeUnixNano: time.Now().UnixNano(), Message: "edge enrolled", } payload, err := proto.Marshal(res) if err != nil { t.Fatalf("marshal EdgeHelloResponse: %v", err) } got, err := parse(payload) if err != nil { t.Fatalf("parse EdgeHelloResponse: %v", err) } parsed, ok := got.(*iop.EdgeHelloResponse) if !ok { t.Fatalf("expected *iop.EdgeHelloResponse, got %T", got) } if !parsed.Accepted { t.Error("accepted: got false want true") } if parsed.Protocol != Protocol { t.Errorf("protocol: got %q want %q", parsed.Protocol, Protocol) } if parsed.ServerTimeUnixNano <= 0 { t.Errorf("server_time_unix_nano should be positive, got %d", parsed.ServerTimeUnixNano) } } func TestEdgeParserMapStatusRoundTrip(t *testing.T) { parserMap := EdgeParserMap() reqType := proto_socket.TypeNameOf(&iop.EdgeStatusRequest{}) parseReq, ok := parserMap[reqType] if !ok { t.Fatalf("missing parser for %s", reqType) } reqPayload, err := proto.Marshal(&iop.EdgeStatusRequest{RequestId: "edge-status-1"}) if err != nil { t.Fatalf("marshal EdgeStatusRequest: %v", err) } gotReq, err := parseReq(reqPayload) if err != nil { t.Fatalf("parse EdgeStatusRequest: %v", err) } parsedReq, ok := gotReq.(*iop.EdgeStatusRequest) if !ok { t.Fatalf("expected *iop.EdgeStatusRequest, got %T", gotReq) } if parsedReq.GetRequestId() != "edge-status-1" { t.Errorf("request_id: got %q want %q", parsedReq.GetRequestId(), "edge-status-1") } resType := proto_socket.TypeNameOf(&iop.EdgeStatusResponse{}) parseRes, ok := parserMap[resType] if !ok { t.Fatalf("missing parser for %s", resType) } res := &iop.EdgeStatusResponse{ RequestId: "edge-status-1", EdgeId: "edge-dgx-group", EdgeName: "DGX Group", ObservedTimeUnixNano: time.Now().UnixNano(), Nodes: []*iop.EdgeNodeSnapshot{ {NodeId: "node-1", Alias: "alpha", Label: "node0", Connected: true}, }, Metadata: map[string]string{"region": "local"}, } resPayload, err := proto.Marshal(res) if err != nil { t.Fatalf("marshal EdgeStatusResponse: %v", err) } gotRes, err := parseRes(resPayload) if err != nil { t.Fatalf("parse EdgeStatusResponse: %v", err) } parsedRes, ok := gotRes.(*iop.EdgeStatusResponse) if !ok { t.Fatalf("expected *iop.EdgeStatusResponse, got %T", gotRes) } if parsedRes.GetEdgeId() != "edge-dgx-group" { t.Errorf("edge_id: got %q want %q", parsedRes.GetEdgeId(), "edge-dgx-group") } if len(parsedRes.GetNodes()) != 1 { t.Fatalf("nodes length: got %d want 1", len(parsedRes.GetNodes())) } node := parsedRes.GetNodes()[0] if node.GetNodeId() != "node-1" || node.GetAlias() != "alpha" || node.GetLabel() != "node0" { t.Errorf("node snapshot: %+v", node) } if !node.GetConnected() { t.Errorf("expected node snapshot connected=true") } if parsedRes.GetMetadata()["region"] != "local" { t.Errorf("metadata region: got %q want %q", parsedRes.GetMetadata()["region"], "local") } } func TestEdgeParserMapNodeEventRoundTrip(t *testing.T) { parserMap := EdgeParserMap() typeName := proto_socket.TypeNameOf(&iop.EdgeNodeEvent{}) parse, ok := parserMap[typeName] if !ok { t.Fatalf("missing parser for %s", typeName) } event := &iop.EdgeNodeEvent{ EventId: "evt-node-1", Type: "node.connected", Source: "edge", NodeId: "node-1", Alias: "alpha", Reason: "registered", Timestamp: time.Now().UnixNano(), Metadata: map[string]string{ "rack": "r1", }, } payload, err := proto.Marshal(event) if err != nil { t.Fatalf("marshal EdgeNodeEvent: %v", err) } got, err := parse(payload) if err != nil { t.Fatalf("parse EdgeNodeEvent: %v", err) } parsed, ok := got.(*iop.EdgeNodeEvent) if !ok { t.Fatalf("expected *iop.EdgeNodeEvent, got %T", got) } if parsed.GetEventId() != event.GetEventId() { t.Errorf("event_id: got %q want %q", parsed.GetEventId(), event.GetEventId()) } if parsed.GetType() != event.GetType() { t.Errorf("type: got %q want %q", parsed.GetType(), event.GetType()) } if parsed.GetNodeId() != event.GetNodeId() { t.Errorf("node_id: got %q want %q", parsed.GetNodeId(), event.GetNodeId()) } if parsed.GetAlias() != event.GetAlias() { t.Errorf("alias: got %q want %q", parsed.GetAlias(), event.GetAlias()) } if parsed.GetTimestamp() != event.GetTimestamp() { t.Errorf("timestamp: got %d want %d", parsed.GetTimestamp(), event.GetTimestamp()) } if parsed.GetMetadata()["rack"] != "r1" { t.Errorf("metadata rack: got %q want %q", parsed.GetMetadata()["rack"], "r1") } } // TestEdgeStatusRequestDoesNotExposeNodeDirectScheduling guards that the status // request stays a correlation-only contract and never grows direct Node // addressing or scheduling fields. func TestEdgeStatusRequestDoesNotExposeNodeDirectScheduling(t *testing.T) { fields := (&iop.EdgeStatusRequest{}).ProtoReflect().Descriptor().Fields() forbidden := []protoreflect.Name{ "node_id", "node_address", "schedule_request", "target", "token", } for _, name := range forbidden { if fields.ByName(name) != nil { t.Fatalf("EdgeStatusRequest must not expose direct Node scheduling field %q", name) } } } func TestEdgeHelloContractDoesNotExposeNodeDirectScheduling(t *testing.T) { fields := (&iop.EdgeHelloRequest{}).ProtoReflect().Descriptor().Fields() forbidden := []protoreflect.Name{ "node_id", "node_address", "schedule_request", "target", } for _, name := range forbidden { if fields.ByName(name) != nil { t.Fatalf("EdgeHelloRequest must not expose direct Node scheduling field %q", name) } } } func TestWireTransportBoundary(t *testing.T) { if Protocol != "protobuf-socket" { t.Fatalf("Protocol = %q, want %q", Protocol, "protobuf-socket") } if EdgeTransport != "proto-socket-tcp" { t.Fatalf("EdgeTransport = %q, want %q", EdgeTransport, "proto-socket-tcp") } if ClientTransport != "proto-socket-ws" { t.Fatalf("ClientTransport = %q, want %q", ClientTransport, "proto-socket-ws") } } func TestProtoOwnershipGuard(t *testing.T) { forbiddenNames := []string{ "edge_config_source", "node_registry_source", "artifact_store", "build_log_store", } messages := []proto.Message{ &iop.EdgeStatusRequest{}, &iop.EdgeStatusResponse{}, &iop.EdgeHelloRequest{}, &iop.EdgeHelloResponse{}, &iop.EdgeCommandRequest{}, &iop.EdgeCommandResponse{}, &iop.EdgeCommandEvent{}, &iop.EdgeCapabilitySummary{}, &iop.EdgeDomainAgentSummary{}, } for _, msg := range messages { fields := msg.ProtoReflect().Descriptor().Fields() for i := 0; i < fields.Len(); i++ { f := fields.Get(i) name := string(f.Name()) for _, forbidden := range forbiddenNames { if name == forbidden { t.Errorf("Message %s contains forbidden field %s", msg.ProtoReflect().Descriptor().Name(), name) } } } } } func TestControlPlaneRegistryOwnershipGuard(t *testing.T) { visited := make(map[reflect.Type]bool) if hasEdgeConfig(reflect.TypeOf(EdgeRegistry{}), visited) { t.Errorf("EdgeRegistry must not store EdgeConfig or configurations containing EdgeConfig") } visitedState := make(map[reflect.Type]bool) if hasEdgeConfig(reflect.TypeOf(EdgeConnectionState{}), visitedState) { t.Errorf("EdgeConnectionState must not store EdgeConfig or configurations containing EdgeConfig") } } func hasEdgeConfig(t reflect.Type, visited map[reflect.Type]bool) bool { if visited[t] { return false } visited[t] = true if t == reflect.TypeOf(config.EdgeConfig{}) { return true } switch t.Kind() { case reflect.Struct: for i := 0; i < t.NumField(); i++ { if hasEdgeConfig(t.Field(i).Type, visited) { return true } } case reflect.Ptr, reflect.Slice, reflect.Array: return hasEdgeConfig(t.Elem(), visited) case reflect.Map: return hasEdgeConfig(t.Key(), visited) || hasEdgeConfig(t.Elem(), visited) } return false }