This is an automated email from the ASF dual-hosted git repository.

beto pushed a commit to branch ScreenshotCachePayload-serialization
in repository https://gitbox.apache.org/repos/asf/superset.git

commit cf71fc1d91ee78d415d11707dbc7c8088b6198cd
Author: Beto Dealmeida <[email protected]>
AuthorDate: Wed Feb 5 17:45:36 2025 -0500

    fix: ScreenshotCachePayload serialization
---
 superset/charts/api.py                    |  4 +--
 superset/dashboards/api.py                |  4 +--
 superset/utils/screenshots.py             | 48 ++++++++++++++++++++++++-------
 tests/unit_tests/utils/screenshot_test.py | 11 ++++---
 4 files changed, 47 insertions(+), 20 deletions(-)

diff --git a/superset/charts/api.py b/superset/charts/api.py
index a600a3ca7f..292575feac 100644
--- a/superset/charts/api.py
+++ b/superset/charts/api.py
@@ -623,7 +623,7 @@ class ChartRestApi(BaseSupersetModelRestApi):
 
         if cache_payload.should_trigger_task(force):
             logger.info("Triggering screenshot ASYNC")
-            screenshot_obj.cache.set(cache_key, ScreenshotCachePayload())
+            screenshot_obj.cache.set(cache_key, 
ScreenshotCachePayload().to_dict())
             cache_chart_thumbnail.delay(
                 current_user=get_current_user(),
                 chart_id=chart.id,
@@ -755,7 +755,7 @@ class ChartRestApi(BaseSupersetModelRestApi):
             logger.info(
                 "Triggering thumbnail compute (chart id: %s) ASYNC", 
str(chart.id)
             )
-            screenshot_obj.cache.set(cache_key, ScreenshotCachePayload())
+            screenshot_obj.cache.set(cache_key, 
ScreenshotCachePayload().to_dict())
             cache_chart_thumbnail.delay(
                 current_user=current_user,
                 chart_id=chart.id,
diff --git a/superset/dashboards/api.py b/superset/dashboards/api.py
index c8c744ec63..c15610010b 100644
--- a/superset/dashboards/api.py
+++ b/superset/dashboards/api.py
@@ -1115,7 +1115,7 @@ class DashboardRestApi(BaseSupersetModelRestApi):
 
         if cache_payload.should_trigger_task(force):
             logger.info("Triggering screenshot ASYNC")
-            screenshot_obj.cache.set(cache_key, ScreenshotCachePayload())
+            screenshot_obj.cache.set(cache_key, 
ScreenshotCachePayload().to_dict())
             cache_dashboard_screenshot.delay(
                 username=get_current_user(),
                 guest_token=(
@@ -1296,7 +1296,7 @@ class DashboardRestApi(BaseSupersetModelRestApi):
                 "Triggering thumbnail compute (dashboard id: %s) ASYNC",
                 str(dashboard.id),
             )
-            screenshot_obj.cache.set(cache_key, ScreenshotCachePayload())
+            screenshot_obj.cache.set(cache_key, 
ScreenshotCachePayload().to_dict())
             cache_dashboard_thumbnail.delay(
                 current_user=current_user,
                 dashboard_id=dashboard.id,
diff --git a/superset/utils/screenshots.py b/superset/utils/screenshots.py
index 86f5a94ce7..9323caa7cc 100644
--- a/superset/utils/screenshots.py
+++ b/superset/utils/screenshots.py
@@ -20,7 +20,7 @@ import logging
 from datetime import datetime
 from enum import Enum
 from io import BytesIO
-from typing import TYPE_CHECKING
+from typing import TYPE_CHECKING, TypedDict
 
 from flask import current_app
 
@@ -63,13 +63,37 @@ class StatusValues(Enum):
     ERROR = "Error"
 
 
+class ScreenshotCachePayloadType(TypedDict):
+    image: bytes | None
+    timestamp: str
+    status: str
+
+
 class ScreenshotCachePayload:
-    def __init__(self, image: bytes | None = None):
+    def __init__(
+        self,
+        image: bytes | None = None,
+        status: str = StatusValues.PENDING,
+        timestamp: str = "",
+    ):
         self._image = image
-        self._timestamp = datetime.now().isoformat()
-        self.status = StatusValues.PENDING
-        if image:
-            self.status = StatusValues.UPDATED
+        self._timestamp = timestamp or datetime.now().isoformat()
+        self.status = StatusValues.UPDATED if image else status
+
+    @classmethod
+    def from_dict(cls, payload: ScreenshotCachePayloadType) -> 
ScreenshotCachePayload:
+        return cls(
+            image=payload["image"],
+            status=StatusValues(payload["status"]),
+            timestamp=payload["timestamp"],
+        )
+
+    def to_dict(self) -> ScreenshotCachePayloadType:
+        return {
+            "image": self._image,
+            "timestamp": self._timestamp,
+            "status": self.status.value,
+        }
 
     def update_timestamp(self) -> None:
         self._timestamp = datetime.now().isoformat()
@@ -177,8 +201,12 @@ class BaseScreenshot:
     def get_from_cache_key(cls, cache_key: str) -> ScreenshotCachePayload | 
None:
         logger.info("Attempting to get from cache: %s", cache_key)
         if payload := cls.cache.get(cache_key):
-            # for backwards compatability, byte objects should be converted
-            if not isinstance(payload, ScreenshotCachePayload):
+            # Initially, only bytes were stored. This was changed to store an 
instance
+            # of ScreenshotCachePayload, but since it can't be serialized in 
all
+            # backends it was further changed to just a dict.
+            if isinstance(payload, dict):
+                payload = ScreenshotCachePayload.from_dict(payload)
+            elif isinstance(payload, bytes):
                 payload = ScreenshotCachePayload(payload)
             return payload
         logger.info("Failed at getting from cache: %s", cache_key)
@@ -217,7 +245,7 @@ class BaseScreenshot:
         thumb_size = thumb_size or self.thumb_size
         logger.info("Processing url for thumbnail: %s", cache_key)
         cache_payload.computing()
-        self.cache.set(cache_key, cache_payload)
+        self.cache.set(cache_key, cache_payload.to_dict())
         image = None
         # Assuming all sorts of things can go wrong with Selenium
         try:
@@ -239,7 +267,7 @@ class BaseScreenshot:
             logger.info("Caching thumbnail: %s", cache_key)
             with 
event_logger.log_context(f"screenshot.cache.{self.thumbnail_type}"):
                 cache_payload.update(image)
-        self.cache.set(cache_key, cache_payload)
+        self.cache.set(cache_key, cache_payload.to_dict())
         logger.info("Updated thumbnail cache; Status: %s", 
cache_payload.get_status())
         return
 
diff --git a/tests/unit_tests/utils/screenshot_test.py 
b/tests/unit_tests/utils/screenshot_test.py
index 5d29d829a2..1af00a3636 100644
--- a/tests/unit_tests/utils/screenshot_test.py
+++ b/tests/unit_tests/utils/screenshot_test.py
@@ -26,7 +26,6 @@ from superset.utils.hashing import md5_sha_from_dict
 from superset.utils.screenshots import (
     BaseScreenshot,
     ScreenshotCachePayload,
-    StatusValues,
 )
 
 BASE_SCREENSHOT_PATH = "superset.utils.screenshots.BaseScreenshot"
@@ -122,7 +121,7 @@ class TestComputeAndCache:
         self._setup_compute_and_cache(mocker, screenshot_obj)
         screenshot_obj.compute_and_cache(force=False)
         cache_payload: ScreenshotCachePayload = screenshot_obj.cache.get("key")
-        assert cache_payload.status == StatusValues.UPDATED
+        assert cache_payload["status"] == "Updated"
 
     def test_screenshot_error(self, mocker: MockerFixture, screenshot_obj):
         mocks = self._setup_compute_and_cache(mocker, screenshot_obj)
@@ -130,7 +129,7 @@ class TestComputeAndCache:
         get_screenshot.side_effect = Exception
         screenshot_obj.compute_and_cache(force=False)
         cache_payload: ScreenshotCachePayload = screenshot_obj.cache.get("key")
-        assert cache_payload.status == StatusValues.ERROR
+        assert cache_payload["status"] == "Error"
 
     def test_resize_error(self, mocker: MockerFixture, screenshot_obj):
         mocks = self._setup_compute_and_cache(mocker, screenshot_obj)
@@ -138,7 +137,7 @@ class TestComputeAndCache:
         resize_image.side_effect = Exception
         screenshot_obj.compute_and_cache(force=False)
         cache_payload: ScreenshotCachePayload = screenshot_obj.cache.get("key")
-        assert cache_payload.status == StatusValues.ERROR
+        assert cache_payload["status"] == "Error"
 
     def test_skips_if_computing(self, mocker: MockerFixture, screenshot_obj):
         mocks = self._setup_compute_and_cache(mocker, screenshot_obj)
@@ -156,7 +155,7 @@ class TestComputeAndCache:
         screenshot_obj.compute_and_cache(force=True)
         get_screenshot.assert_called_once()
         cache_payload: ScreenshotCachePayload = screenshot_obj.cache.get("key")
-        assert cache_payload.status == StatusValues.UPDATED
+        assert cache_payload["status"] == "Updated"
 
     def test_skips_if_updated(self, mocker: MockerFixture, screenshot_obj):
         mocks = self._setup_compute_and_cache(mocker, screenshot_obj)
@@ -178,7 +177,7 @@ class TestComputeAndCache:
         )
         get_screenshot.assert_called_once()
         cache_payload: ScreenshotCachePayload = screenshot_obj.cache.get("key")
-        assert cache_payload._image != b"initial_value"
+        assert cache_payload["image"] != b"initial_value"
 
     def test_resize(self, mocker: MockerFixture, screenshot_obj):
         mocks = self._setup_compute_and_cache(mocker, screenshot_obj)

Reply via email to