"""find_protection: unique VM + every VPG, or stop.""" from __future__ import annotations from dataclasses import dataclass, field from typing import Any, Literal from zerto_rewind_mcp.status import can_tag, status_name, substatus_name from zerto_rewind_mcp.util import pick Outcome = Literal["none", "ambiguous", "ok"] @dataclass class VpgMembership: vpg_identifier: str vpg_name: str status: int | None sub_status: int | None can_tag: bool skip_reason: str | None def as_dict(self) -> dict[str, Any]: return { "vpg_identifier": self.vpg_identifier, "vpg_name": self.vpg_name, "status": status_name(self.status), "sub_status": substatus_name(self.sub_status), "can_tag": self.can_tag, "skip_reason": self.skip_reason, } @dataclass class VmMatch: vm_identifier: str vm_name: str vpgs: list[VpgMembership] = field(default_factory=list) def as_dict(self) -> dict[str, Any]: return { "vm_identifier": self.vm_identifier, "vm_name": self.vm_name, "vpgs": [v.as_dict() for v in self.vpgs], } @dataclass class FindResult: outcome: Outcome query: str message: str vm: VmMatch | None = None matches: list[VmMatch] = field(default_factory=list) def as_dict(self) -> dict[str, Any]: payload: dict[str, Any] = { "outcome": self.outcome, "query": self.query, "message": self.message, } if self.vm is not None: payload["vm"] = self.vm.as_dict() if self.matches: payload["matches"] = [m.as_dict() for m in self.matches] return payload @property def taggable_vpgs(self) -> list[VpgMembership]: if self.vm is None: return [] return [v for v in self.vm.vpgs if v.can_tag] def membership_from_row(row: dict[str, Any]) -> VpgMembership: status = pick(row, "Status", "status") sub = pick(row, "SubStatus", "subStatus", "sub_status") ok, reason = can_tag(status, sub) return VpgMembership( vpg_identifier=str(pick(row, "VpgIdentifier", "vpgIdentifier") or ""), vpg_name=str(pick(row, "VpgName", "vpgName") or ""), status=status, sub_status=sub, can_tag=ok, skip_reason=reason, ) def group_vm_rows(rows: list[dict[str, Any]]) -> list[VmMatch]: """One GET /v1/vms row per VM-in-VPG. Group by VmIdentifier.""" by_id: dict[str, VmMatch] = {} order: list[str] = [] for row in rows: vm_id = pick(row, "VmIdentifier", "vmIdentifier") vm_name = pick(row, "VmName", "vmName") or "" if not vm_id: continue vm_id = str(vm_id) if vm_id not in by_id: by_id[vm_id] = VmMatch(vm_identifier=vm_id, vm_name=str(vm_name)) order.append(vm_id) elif vm_name and not by_id[vm_id].vm_name: by_id[vm_id].vm_name = str(vm_name) vpg = membership_from_row(row) if vpg.vpg_identifier and vpg.vpg_identifier not in { x.vpg_identifier for x in by_id[vm_id].vpgs }: by_id[vm_id].vpgs.append(vpg) return [by_id[i] for i in order] def find_from_rows(query: str, rows: list[dict[str, Any]]) -> FindResult: matches = group_vm_rows(rows) if not matches: return FindResult( outcome="none", query=query, message=( f"No protected VM matched {query!r}. " "Unprotected or unknown: refuse the change. Zerto cannot rewind this." ), ) if len(matches) > 1: names = ", ".join(f"{m.vm_name} ({m.vm_identifier})" for m in matches) return FindResult( outcome="ambiguous", query=query, message=( f"{len(matches)} VMs matched {query!r}: {names}. " "Pass a Zerto vmIdentifier. Do not mutate." ), matches=matches, ) vm = matches[0] if not vm.vpgs: return FindResult( outcome="none", query=query, message=( f"VM {vm.vm_name} ({vm.vm_identifier}) has no VPG. " "Unprotected: refuse the change." ), vm=vm, ) taggable = [v for v in vm.vpgs if v.can_tag] skipped = [v for v in vm.vpgs if not v.can_tag] skip_txt = "" if skipped: bits = "; ".join(f"{v.vpg_name}: {v.skip_reason}" for v in skipped) skip_txt = f" Skipping: {bits}." if not taggable: return FindResult( outcome="ok", query=query, message=( f"VM {vm.vm_name} is in {len(vm.vpgs)} VPG(s) but none can be tagged right now." f"{skip_txt} Refuse the change." ), vm=vm, ) return FindResult( outcome="ok", query=query, message=( f"VM {vm.vm_name} ({vm.vm_identifier}): " f"{len(taggable)} protecting VPG(s) to tag, {len(skipped)} skipped." f"{skip_txt}" ), vm=vm, )