diff --git a/backend/app/api/v1/router.py b/backend/app/api/v1/router.py index fffd65b..29ba63e 100644 --- a/backend/app/api/v1/router.py +++ b/backend/app/api/v1/router.py @@ -27,6 +27,7 @@ from app.models.domain import ( Project, Role, SecurityGroup, + SecurityGroupMember, SecurityRule, ServiceCatalogItem, SystemSetting, @@ -64,6 +65,8 @@ from app.schemas.domain import ( ServiceCatalogCreate, ServiceCatalogRead, SecurityGroupCreate, + SecurityGroupMemberCreate, + SecurityGroupMemberRead, SecurityGroupRead, SetupCompleteRequest, SetupStatus, @@ -269,6 +272,30 @@ def workload_provider_target(db: Session, cluster: Cluster, ref: str) -> tuple[d ) +def workload_provider_targets(db: Session, cluster: Cluster, ref: str) -> list[tuple[dict, Workload]]: + resolved = workload_provider_target(db, cluster, ref) + if resolved: + return [resolved] + if not ref.startswith("sg:"): + return [] + group_ref = ref.removeprefix("sg:") + group = db.get(SecurityGroup, group_ref) or db.scalar(select(SecurityGroup).where(SecurityGroup.name == group_ref)) + if not group: + return [] + workloads = db.scalars( + select(Workload) + .join(SecurityGroupMember, SecurityGroupMember.workload_id == Workload.id) + .where(SecurityGroupMember.security_group_id == group.id, Workload.cluster_id == cluster.id) + .order_by(Workload.name) + ).all() + targets = [] + for workload in workloads: + target = workload_provider_target(db, cluster, f"workload:{workload.id}") + if target: + targets.append(target) + return targets + + def endpoint_values(db: Session, cluster: Cluster, ref: str) -> tuple[list[str | None], list[str]]: if ref == "any": return [None], [] @@ -291,6 +318,29 @@ def endpoint_values(db: Session, cluster: Cluster, ref: str) -> tuple[list[str | if values: return values, [] return [], [f"Network {network_name} has no IPAM subnets to use as provider-side matcher."] + if ref.startswith("sg:"): + group_ref = ref.removeprefix("sg:") + group = db.get(SecurityGroup, group_ref) or db.scalar(select(SecurityGroup).where(SecurityGroup.name == group_ref)) + if not group: + return [], [f"Security group {group_ref} was not found."] + workload_ids = [ + row[0] + for row in db.execute( + select(SecurityGroupMember.workload_id) + .join(Workload, Workload.id == SecurityGroupMember.workload_id) + .where(SecurityGroupMember.security_group_id == group.id, Workload.cluster_id == cluster.id) + ).all() + ] + if not workload_ids: + return [], [f"Security group {group.name} has no workloads in this cluster."] + values = [ + address.address + for address in db.scalars(select(IpAddress).where(IpAddress.workload_id.in_(workload_ids)).order_by(IpAddress.address)).all() + if address.address + ] + if values: + return values, [] + return [], [f"Security group {group.name} has no assigned member IPs."] return [], [f"Endpoint {ref} is not yet resolvable to a Proxmox firewall matcher."] @@ -491,6 +541,23 @@ def endpoint_ref_matches_flow_side( return False subnets = db.scalars(select(Subnet).where(Subnet.network_id == network.id)).all() return any(ip_value_matches(subnet.cidr, flow_ip) for subnet in subnets) + if value.startswith("sg:"): + group_ref = value.removeprefix("sg:") + group = db.get(SecurityGroup, group_ref) or db.scalar(select(SecurityGroup).where(SecurityGroup.name == group_ref)) + if not group: + return False + if side_workload: + return db.scalar( + select(func.count()) + .select_from(SecurityGroupMember) + .where(SecurityGroupMember.security_group_id == group.id, SecurityGroupMember.workload_id == side_workload.id) + ) > 0 + member_ips = db.scalars( + select(IpAddress.address) + .join(SecurityGroupMember, SecurityGroupMember.workload_id == IpAddress.workload_id) + .where(SecurityGroupMember.security_group_id == group.id) + ).all() + return flow_ip in set(member_ips) if is_ip_or_cidr(value): return ip_value_matches(value, flow_ip) return False @@ -858,47 +925,48 @@ def resolve_firewall_preview(db: Session, cluster: Cluster, preview: FirewallPre mapped = dict(rule) direction = str(rule.get("direction", "ingress")) target_ref = str(rule.get("destination") if direction == "ingress" else rule.get("source")) - target = workload_provider_target(db, cluster, target_ref) - if not target: + targets = workload_provider_targets(db, cluster, target_ref) + if not targets: conflicts.append( f"Rule {rule_index} needs a concrete {'destination' if direction == 'ingress' else 'source'} workload for Proxmox live apply." ) generated_rules.append(mapped) continue - provider_target, target_workload = target remote_ref = str(rule.get("source") if direction == "ingress" else rule.get("destination")) remote_values, endpoint_warnings = endpoint_values(db, cluster, remote_ref) warnings.extend(endpoint_warnings) if not remote_values: conflicts.append(f"Rule {rule_index} cannot resolve {remote_ref} to a Proxmox firewall source/destination matcher.") - generated_rules.append({**mapped, "provider_target": provider_target}) + for provider_target, _target_workload in targets: + generated_rules.append({**mapped, "provider_target": provider_target}) continue ports = str(rule.get("ports", "any")) protocol = str(rule.get("protocol", "any")) protocols = ["tcp", "udp"] if protocol == "tcp/udp" else [protocol] - for remote_value in remote_values: - for provider_protocol in protocols: - provider_rule = { - "type": "in" if direction == "ingress" else "out", - "action": proxmox_action(str(rule.get("action", "allow"))), - "enable": 1, - "comment": ( - f"NexaFabric policy={rule.get('policy_id')} version={rule.get('policy_version')} " - f"rule={rule_index} target={target_workload.name}" - ), - } - if provider_protocol != "any": - provider_rule["proto"] = provider_protocol - if ports != "any": - provider_rule["dport"] = ports - if remote_value: - provider_rule["source" if direction == "ingress" else "dest"] = remote_value - if rule.get("logging"): - provider_rule["log"] = "info" - mapped_rule = {**mapped, "provider_target": provider_target, "provider_rule": provider_rule} - generated_rules.append(mapped_rule) + for provider_target, target_workload in targets: + for remote_value in remote_values: + for provider_protocol in protocols: + provider_rule = { + "type": "in" if direction == "ingress" else "out", + "action": proxmox_action(str(rule.get("action", "allow"))), + "enable": 1, + "comment": ( + f"NexaFabric policy={rule.get('policy_id')} version={rule.get('policy_version')} " + f"rule={rule_index} target={target_workload.name}" + ), + } + if provider_protocol != "any": + provider_rule["proto"] = provider_protocol + if ports != "any": + provider_rule["dport"] = ports + if remote_value: + provider_rule["source" if direction == "ingress" else "dest"] = remote_value + if rule.get("logging"): + provider_rule["log"] = "info" + mapped_rule = {**mapped, "provider_target": provider_target, "provider_rule": provider_rule} + generated_rules.append(mapped_rule) return FirewallPreview( policy_id=preview.policy_id, @@ -1131,6 +1199,8 @@ def delete_cluster(cluster_id: str, user: CurrentUser, db: Session = Depends(get raise HTTPException(status_code=404, detail="Cluster not found") workload_ids = [row[0] for row in db.execute(select(Workload.id).where(Workload.cluster_id == cluster.id)).all()] if workload_ids: + for member in db.scalars(select(SecurityGroupMember).where(SecurityGroupMember.workload_id.in_(workload_ids))).all(): + db.delete(member) for address in db.scalars(select(IpAddress).where(IpAddress.workload_id.in_(workload_ids))).all(): db.delete(address) network_ids = [row[0] for row in db.execute(select(Network.id).where(Network.cluster_id == cluster.id)).all()] @@ -1862,8 +1932,36 @@ def create_project(payload: ProjectCreate, user: CurrentUser, db: Session = Depe @api_router.get("/security-groups", response_model=list[SecurityGroupRead]) -def security_groups(_: CurrentUser, db: Session = Depends(get_db)) -> list[SecurityGroup]: - return db.scalars(select(SecurityGroup).order_by(SecurityGroup.name)).all() +def security_groups(_: CurrentUser, db: Session = Depends(get_db)) -> list[dict[str, object]]: + groups = db.scalars(select(SecurityGroup).order_by(SecurityGroup.name)).all() + members_by_group: dict[str, list[dict[str, object]]] = {group.id: [] for group in groups} + if groups: + rows = db.execute( + select(SecurityGroupMember, Workload) + .join(Workload, Workload.id == SecurityGroupMember.workload_id) + .where(SecurityGroupMember.security_group_id.in_([group.id for group in groups])) + .order_by(Workload.name) + ).all() + for member, workload in rows: + members_by_group.setdefault(member.security_group_id, []).append( + { + "id": member.id, + "security_group_id": member.security_group_id, + "workload_id": member.workload_id, + "workload_name": workload.name, + "workload_external_id": workload.external_id, + } + ) + return [ + { + "id": group.id, + "project_id": group.project_id, + "name": group.name, + "description": group.description, + "members": members_by_group.get(group.id, []), + } + for group in groups + ] @api_router.post("/security-groups", response_model=SecurityGroupRead) @@ -1876,6 +1974,50 @@ def create_security_group(payload: SecurityGroupCreate, user: CurrentUser, db: S return group +def security_group_member_payload(db: Session, member: SecurityGroupMember) -> dict[str, object]: + workload = db.get(Workload, member.workload_id) + return { + "id": member.id, + "security_group_id": member.security_group_id, + "workload_id": member.workload_id, + "workload_name": workload.name if workload else None, + "workload_external_id": workload.external_id if workload else None, + } + + +@api_router.post("/security-groups/{group_id}/members", response_model=SecurityGroupMemberRead) +def add_security_group_member(group_id: str, payload: SecurityGroupMemberCreate, user: CurrentUser, db: Session = Depends(get_db)) -> dict[str, object]: + if not db.get(SecurityGroup, group_id): + raise HTTPException(status_code=404, detail="Security group not found") + if not db.get(Workload, payload.workload_id): + raise HTTPException(status_code=404, detail="Workload not found") + existing = db.scalar( + select(SecurityGroupMember).where( + SecurityGroupMember.security_group_id == group_id, + SecurityGroupMember.workload_id == payload.workload_id, + ) + ) + if existing: + return security_group_member_payload(db, existing) + member = SecurityGroupMember(security_group_id=group_id, workload_id=payload.workload_id) + db.add(member) + commit_or_400(db) + db.refresh(member) + write_audit(db, action="security_group.member_added", object_type="security_group", object_id=group_id, user_id=user.id, new_values=payload.model_dump()) + return security_group_member_payload(db, member) + + +@api_router.delete("/security-groups/{group_id}/members/{member_id}") +def delete_security_group_member(group_id: str, member_id: str, user: CurrentUser, db: Session = Depends(get_db)) -> dict[str, str]: + member = db.get(SecurityGroupMember, member_id) + if not member or member.security_group_id != group_id: + raise HTTPException(status_code=404, detail="Security group member not found") + db.delete(member) + commit_or_400(db) + write_audit(db, action="security_group.member_removed", object_type="security_group", object_id=group_id, user_id=user.id, old_values={"member_id": member_id}) + return {"status": "deleted", "id": member_id} + + @api_router.get("/security-groups/{group_id}/rules", response_model=list[SecurityRuleRead]) def security_group_rules(group_id: str, _: CurrentUser, db: Session = Depends(get_db)) -> list[SecurityRule]: return db.scalars( diff --git a/backend/app/models/domain.py b/backend/app/models/domain.py index 4d3bb1a..1c6bfa0 100644 --- a/backend/app/models/domain.py +++ b/backend/app/models/domain.py @@ -231,6 +231,17 @@ class SecurityGroup(Base, TimestampMixin): description: Mapped[str | None] = mapped_column(Text) +class SecurityGroupMember(Base, TimestampMixin): + __tablename__ = "security_group_members" + __table_args__ = (UniqueConstraint("security_group_id", "workload_id"),) + + id: Mapped[str] = mapped_column(String, primary_key=True, default=new_id) + security_group_id: Mapped[str] = mapped_column(ForeignKey("security_groups.id"), index=True) + workload_id: Mapped[str] = mapped_column(ForeignKey("workloads.id"), index=True) + security_group: Mapped[SecurityGroup] = relationship() + workload: Mapped[Workload] = relationship() + + class SecurityRule(Base, TimestampMixin): __tablename__ = "security_rules" diff --git a/backend/app/schemas/domain.py b/backend/app/schemas/domain.py index 088cd7a..d81a38b 100644 --- a/backend/app/schemas/domain.py +++ b/backend/app/schemas/domain.py @@ -127,6 +127,10 @@ class SecurityGroupCreate(BaseModel): description: str | None = None +class SecurityGroupMemberCreate(BaseModel): + workload_id: str + + class SecurityRuleCreate(BaseModel): security_group_id: str direction: str = "ingress" @@ -305,6 +309,15 @@ class SecurityGroupRead(OrmModel): project_id: str | None name: str description: str | None + members: list[dict[str, Any]] = [] + + +class SecurityGroupMemberRead(OrmModel): + id: str + security_group_id: str + workload_id: str + workload_name: str | None = None + workload_external_id: str | None = None class SecurityRuleRead(OrmModel): diff --git a/frontend/src/api/client.ts b/frontend/src/api/client.ts index f4c5f40..0051b89 100644 --- a/frontend/src/api/client.ts +++ b/frontend/src/api/client.ts @@ -99,6 +99,15 @@ export type SecurityGroup = { project_id: string | null; name: string; description: string | null; + members: SecurityGroupMember[]; +}; + +export type SecurityGroupMember = { + id: string; + security_group_id: string; + workload_id: string; + workload_name: string | null; + workload_external_id: string | null; }; export type SecurityRule = { diff --git a/frontend/src/components/SearchableSelect.tsx b/frontend/src/components/SearchableSelect.tsx new file mode 100644 index 0000000..2400aad --- /dev/null +++ b/frontend/src/components/SearchableSelect.tsx @@ -0,0 +1,74 @@ +import { useMemo, useState } from "react"; + +import { inputClass } from "./FormControls"; + +export type SearchableOption = { + label: string; + value: string; + detail?: string; +}; + +type SearchableSelectProps = { + options: SearchableOption[]; + value: string; + onChange: (value: string) => void; + placeholder?: string; +}; + +export function SearchableSelect({ options, value, onChange, placeholder = "Search..." }: SearchableSelectProps) { + const selected = options.find((option) => option.value === value); + const [query, setQuery] = useState(selected?.label ?? ""); + const [open, setOpen] = useState(false); + const filtered = useMemo(() => { + const normalized = query.trim().toLowerCase(); + if (!normalized || selected?.label === query) { + return options.slice(0, 12); + } + return options + .filter((option) => `${option.label} ${option.detail ?? ""}`.toLowerCase().includes(normalized)) + .slice(0, 12); + }, [options, query, selected?.label]); + + function choose(option: SearchableOption) { + onChange(option.value); + setQuery(option.label); + setOpen(false); + } + + return ( +