Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 21 additions & 9 deletions nodescraper/models/datamodel.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,29 +24,36 @@
#
###############################################################################
import io
import json
import os
import tarfile
from typing import Any, TypeVar, Union

from pydantic import BaseModel, field_validator
from pydantic import BaseModel, ConfigDict, ValidationInfo, field_validator

from nodescraper.utils import get_unique_filename

TDataModel = TypeVar("TDataModel", bound="DataModel")


class FileModel(BaseModel):
"""Binary file payload. JSON uses URL-safe base64 for ``file_contents``."""

model_config = ConfigDict(ser_json_bytes="base64", val_json_bytes="base64")

file_contents: bytes
file_name: str

@field_validator("file_contents", mode="before")
@classmethod
def file_contents_conformer(cls, value: Union[io.BytesIO, str, bytes]) -> bytes:
def file_contents_conformer(
cls, value: Union[io.BytesIO, str, bytes], info: ValidationInfo
) -> Union[str, bytes]:
if isinstance(value, io.BytesIO):
return value.getvalue()
if isinstance(value, str):
if isinstance(value, str) and info.mode != "json":
# Constructor/Python input is text.
return value.encode("utf-8")
# JSON strings stay base64 for val_json_bytes.
return value

def log_model(self, log_path: str) -> None:
Expand Down Expand Up @@ -77,15 +84,21 @@ def log_model(self, log_path: str):
get_unique_filename(log_path, f"{self.__class__.__name__.lower()}.json"),
)

exlude_fields = set()
# Write binary FileModel payloads as files; keep them in JSON as base64.
for key in self.__class__.model_fields:
data = getattr(self, key)
if isinstance(data, FileModel):
data.log_model(log_path)
exlude_fields.add(key)
elif (
isinstance(data, list)
and data
and all(isinstance(item, FileModel) for item in data)
):
for item in data:
item.log_model(log_path)

with open(log_name, "w", encoding="utf-8") as log_file:
log_file.write(self.model_dump_json(indent=2, exclude=exlude_fields))
log_file.write(self.model_dump_json(indent=2))

def merge_data(self, input_data: "DataModel") -> None:
"""Merge data into current data"""
Expand Down Expand Up @@ -122,8 +135,7 @@ def import_model(cls: type[TDataModel], model_input: Union[dict[str, Any], str])
# Build from json file
else:
with open(model_input, "r", encoding="utf-8") as input_file:
data = json.load(input_file)
return cls(**data)
return cls.model_validate_json(input_file.read())

raise ValueError("Invalid input for model data")

Expand Down
51 changes: 50 additions & 1 deletion test/unit/framework/test_datamodel.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,12 @@

from __future__ import annotations

import base64
import json
import os
from pathlib import Path

from nodescraper.models.datamodel import DataModel
from nodescraper.models.datamodel import DataModel, FileModel


class FolderDataModel(DataModel):
Expand All @@ -17,6 +18,11 @@ def build_from_folder(cls, folder_path: str) -> "FolderDataModel":
return cls(value=os.path.basename(folder_path))


class CperListDataModel(DataModel):
value: str = "ok"
cper_data: list[FileModel] = []


def test_import_model_from_dict():
"""Baseline: a dict is passed straight to the model constructor."""
assert FolderDataModel.import_model({"value": "abc"}).value == "abc"
Expand All @@ -36,3 +42,46 @@ def test_import_model_from_directory_uses_build_from_folder(tmp_path: Path):
folder.mkdir()

assert FolderDataModel.import_model(str(folder)).value == "collected"


def test_filemodel_json_roundtrips_binary_cper_bytes():
"""Non-UTF-8 CPER bytes serialize as base64 and reload intact as bytes."""
cper_bytes = b"CPER\x00\x01\xff\xff\xff\xff"
model = FileModel(file_contents=cper_bytes, file_name="corrected-1.cper")
dumped = model.model_dump_json()
payload = json.loads(dumped)
assert payload["file_contents"] == base64.urlsafe_b64encode(cper_bytes).decode("ascii")

reloaded = FileModel.model_validate_json(dumped)
assert isinstance(reloaded.file_contents, bytes)
assert reloaded.file_contents == cper_bytes


def test_filemodel_python_strings_are_utf8_not_base64():
"""A plain string that is also valid base64 must be stored as UTF-8 text."""
text = "VGVzdA=="
model = FileModel(file_contents=text, file_name="t.txt")
assert model.file_contents == text.encode("utf-8")
assert FileModel(file_contents="hello", file_name="t.txt").file_contents == b"hello"


def test_log_model_keeps_cper_list_in_json(tmp_path: Path):
"""CPER-like bytes must not crash log_model and remain in the JSON dump."""
cper_bytes = b"CPER\x00\x01\xff\xff\xff\xff"
model = CperListDataModel(
value="ok",
cper_data=[FileModel(file_contents=cper_bytes, file_name="corrected-1.cper")],
)

model.log_model(str(tmp_path))

cper_path = tmp_path / "corrected-1.cper"
assert cper_path.read_bytes() == cper_bytes

json_files = list(tmp_path.glob("*.json"))
assert len(json_files) == 1
payload = json.loads(json_files[0].read_text(encoding="utf-8"))
assert payload["value"] == "ok"
assert len(payload["cper_data"]) == 1
assert payload["cper_data"][0]["file_name"] == "corrected-1.cper"
assert base64.urlsafe_b64decode(payload["cper_data"][0]["file_contents"]) == cper_bytes
Loading