diff --git a/openedx_learning/apps/authoring/backup_restore/serializers.py b/openedx_learning/apps/authoring/backup_restore/serializers.py index d9bcf5eca..445d372b5 100644 --- a/openedx_learning/apps/authoring/backup_restore/serializers.py +++ b/openedx_learning/apps/authoring/backup_restore/serializers.py @@ -120,3 +120,18 @@ def validate(self, attrs): attrs["children"] = children attrs.pop("container") # Remove the container field after processing return attrs + + +class CollectionSerializer(serializers.Serializer): # pylint: disable=abstract-method + """ + Serializer for collections. + """ + title = serializers.CharField(required=True) + key = serializers.CharField(required=True) + description = serializers.CharField(required=True, allow_blank=True) + created_by = serializers.IntegerField(required=True, allow_null=True) + entities = serializers.ListField( + child=serializers.CharField(), + required=True, + allow_empty=True, + ) diff --git a/openedx_learning/apps/authoring/backup_restore/toml.py b/openedx_learning/apps/authoring/backup_restore/toml.py index 5e5286e38..91bc71a06 100644 --- a/openedx_learning/apps/authoring/backup_restore/toml.py +++ b/openedx_learning/apps/authoring/backup_restore/toml.py @@ -245,3 +245,13 @@ def parse_publishable_entity_toml(content: str) -> tuple[Dict[str, Any], list]: if "version" not in pe_data: raise ValueError("Invalid publishable entity TOML: missing 'version' section") return pe_data["entity"], pe_data.get("version", []) + + +def parse_collection_toml(content: str) -> dict: + """ + Parse the collection TOML content and return a dict of its fields. + """ + collection_data: Dict[str, Any] = tomlkit.parse(content) + if "collection" not in collection_data: + raise ValueError("Invalid collection TOML: missing 'collection' section") + return collection_data["collection"] diff --git a/openedx_learning/apps/authoring/backup_restore/zipper.py b/openedx_learning/apps/authoring/backup_restore/zipper.py index 4c84ee23d..1d408db82 100644 --- a/openedx_learning/apps/authoring/backup_restore/zipper.py +++ b/openedx_learning/apps/authoring/backup_restore/zipper.py @@ -25,12 +25,14 @@ PublishableEntityVersion, ) from openedx_learning.apps.authoring.backup_restore.serializers import ( + CollectionSerializer, ComponentSerializer, ComponentVersionSerializer, ContainerSerializer, ContainerVersionSerializer, ) from openedx_learning.apps.authoring.backup_restore.toml import ( + parse_collection_toml, parse_learning_package_toml, parse_publishable_entity_toml, toml_collection, @@ -413,6 +415,7 @@ def __init__(self) -> None: self.units_map_by_key: dict[str, Any] = {} self.subsections_map_by_key: dict[str, Any] = {} self.sections_map_by_key: dict[str, Any] = {} + self.all_publishable_entities_keys: set[str] = set() # -------------------------- # Public API @@ -434,9 +437,13 @@ def load(self, zipf: zipfile.ZipFile) -> dict[str, Any]: zipf, organized_files["containers"], ContainerSerializer, ContainerVersionSerializer ) + collections_validated = self._extract_collections( + zipf, organized_files["collections"] + ) + self._write_errors() if not self.errors: - self._save(learning_package, components_validated, containers_validated) + self._save(learning_package, components_validated, containers_validated, collections_validated) return { "learning_package": learning_package.key, @@ -474,6 +481,7 @@ def _extract_entities( continue entity_data = serializer.validated_data + self.all_publishable_entities_keys.add(entity_data["key"]) entity_type = entity_data.pop("container_type", "components") results[entity_type].append(entity_data) @@ -491,6 +499,36 @@ def _extract_entities( return results + def _extract_collections( + self, + zipf: zipfile.ZipFile, + collection_files: list[str], + ) -> dict[str, Any]: + """Extraction + validation pipeline for collections.""" + results: dict[str, list[Any]] = defaultdict(list) + + for file in collection_files: + if not file.endswith(".toml"): + # Skip non-TOML files + continue + toml_content = self._read_file_from_zip(zipf, file) + collection_data = parse_collection_toml(toml_content) + serializer = CollectionSerializer(data={"created_by": None, **collection_data}) + if not serializer.is_valid(): + self.errors.append({"file": file, "errors": serializer.errors}) + continue + collection_validated = serializer.validated_data + entities_list = collection_validated["entities"] + for entity_key in entities_list: + if entity_key not in self.all_publishable_entities_keys: + self.errors.append({ + "file": file, + "errors": f"Entity key {entity_key} not found for collection {collection_validated.get('key')}" + }) + results["collections"].append(collection_validated) + + return results + # -------------------------- # Save Logic # -------------------------- @@ -499,7 +537,8 @@ def _save( self, learning_package: LearningPackage, components: dict[str, Any], - containers: dict[str, Any] + containers: dict[str, Any], + collections: dict[str, Any] ) -> None: """Persist all validated entities in two phases: published then drafts.""" @@ -508,11 +547,23 @@ def _save( self._save_units(learning_package, containers) self._save_subsections(learning_package, containers) self._save_sections(learning_package, containers) + self._save_collections(learning_package, collections) publishing_api.publish_all_drafts(learning_package.id) with publishing_api.bulk_draft_changes_for(learning_package.id): self._save_draft_versions(components, containers) + def _save_collections(self, learning_package, collections): + """Save collections and their entities.""" + for valid_collection in collections.get("collections", []): + entities = valid_collection.pop("entities", []) + collection = collections_api.create_collection(learning_package.id, **valid_collection) + collection = collections_api.add_to_collection( + learning_package_id=learning_package.id, + key=collection.key, + entities_qset=publishing_api.get_publishable_entities(learning_package.id).filter(key__in=entities) + ) + def _save_components(self, learning_package, components): """Save components and published component versions.""" for valid_component in components.get("components", []):