diff --git a/deepnote_toolkit/sql/sql_cache_diagnostics.py b/deepnote_toolkit/sql/sql_cache_diagnostics.py new file mode 100644 index 0000000..62f81b7 --- /dev/null +++ b/deepnote_toolkit/sql/sql_cache_diagnostics.py @@ -0,0 +1,176 @@ +"""Log-safe descriptions of SQL cache failures. + +Failed cache requests are described without presigned URL signing parameters. +""" + +import re +from typing import Any, Optional +from urllib.parse import parse_qs, urlsplit + +import requests + +_MAX_ERROR_BODY_BYTES = 4096 +_MAX_ERROR_FIELD_CHARS = 500 +_MAX_OBJECT_PATH_CHARS = 200 +_MAX_RAW_EXCEPTION_CHARS = _MAX_ERROR_FIELD_CHARS * 4 + +# urllib3 connection errors use path-only URLs as well as full URLs +_URL_QUERY_PATTERN = re.compile( + r"((?:https?://)?[^\s?\"'<>]*/[^\s?\"'<>]*)\?[^\s\"'<>]*" +) +# SignatureDoesNotMatch bodies echo params outside any URL prefix +_AWS_CREDENTIAL_PARAM_PATTERN = re.compile( + r"(X-(?:Amz|Goog)-(?:Credential|Security-Token|Signature)=)[^&\s\"'<>]*", + re.IGNORECASE, +) +# urlsplit only splits on a literal '?' +_ENCODED_QUERY_SEPARATOR = re.compile("%3F", re.IGNORECASE) +# Parsing stops at this allowlist; other S3 error elements embed the signed query. +_S3_ERROR_FIELDS = { + "Code": "s3_error_code", + "Message": "s3_error_message", + "Expires": "s3_expires", + "ServerTime": "s3_server_time", +} + +_S3_ERROR_FIELD_PATTERN = re.compile( + rf"<({'|'.join(_S3_ERROR_FIELDS)})>(.*?)", re.DOTALL | re.IGNORECASE +) + + +class SqlCacheHttpError(Exception): + """Non-2xx from the cache object store; str() carries no URL.""" + + def __init__(self, diagnostics: dict[str, Any]) -> None: + super().__init__("SQL cache object store returned an error response") + self.diagnostics = diagnostics + + +def redact_sensitive(text: str) -> str: + """Remove credential-bearing material from text destined for logs.""" + redacted = _URL_QUERY_PATTERN.sub(r"\1?", text) + return _AWS_CREDENTIAL_PARAM_PATTERN.sub(r"\1", redacted) + + +def _redacted_snippet(text: str) -> str: + return redact_sensitive(text)[:_MAX_ERROR_FIELD_CHARS] + + +def _to_int_or_none(value: Optional[str]) -> Optional[int]: + if value is None: + return None + + try: + return int(value) + except ValueError: + return None + + +def seconds_between(start: Optional[float], end: Optional[float]) -> Optional[float]: + """Elapsed monotonic seconds, or None if either timestamp is missing.""" + if start is None or end is None: + return None + + return round(end - start, 1) + + +def safe_url_path(url: object) -> Optional[str]: + """Bounded URL path, or None when parsing fails.""" + if not isinstance(url, str): + return None + + try: + return urlsplit(url).path[:_MAX_OBJECT_PATH_CHARS] + except Exception: + return None + + +def describe_presigned_url(url: object) -> dict[str, Any]: + """Object path and declared expiry without the signing query string.""" + if not isinstance(url, str): + # urlsplit(None) returns bytes fields; bytes in logging extras drop the report + return {"object_host": None, "object_path": None, "url_expires_in": None} + + try: + parts = urlsplit(url) + expires_values = parse_qs(parts.query).get("X-Amz-Expires") + path = _ENCODED_QUERY_SEPARATOR.split(parts.path, maxsplit=1)[0] + return { + # netloc would include user:pass@ userinfo + "object_host": parts.hostname, + "object_path": redact_sensitive(path[:_MAX_OBJECT_PATH_CHARS]), + "url_expires_in": _to_int_or_none( + expires_values[0] if expires_values else None + ), + } + except Exception: + return {"object_host": None, "object_path": None, "url_expires_in": None} + + +def _read_response_body_prefix(response: requests.Response) -> Optional[str]: + try: + chunks = [] + length = 0 + # iter_content bounds wire read; .content would download the full body + for chunk in response.iter_content(_MAX_ERROR_BODY_BYTES): + chunks.append(chunk) + length += len(chunk) + if length >= _MAX_ERROR_BODY_BYTES: + break + + prefix = b"".join(chunks)[:_MAX_ERROR_BODY_BYTES] + return prefix.decode("utf-8", errors="replace") + except Exception: + return None + + +def read_body_snippet(response: requests.Response) -> Optional[str]: + """Redacted response body prefix, or None when unreadable.""" + body = _read_response_body_prefix(response) + return None if body is None else _redacted_snippet(body) + + +def describe_s3_error(response: requests.Response) -> dict[str, Any]: + """Log-safe fields from a failed object-store HTTP response.""" + diagnostics: dict[str, Any] = { + "status_code": response.status_code, + "aws_request_id": response.headers.get("x-amz-request-id"), + "aws_host_id": response.headers.get("x-amz-id-2"), + "aws_date": response.headers.get("Date"), + **{field: None for field in _S3_ERROR_FIELDS.values()}, + } + + body = _read_response_body_prefix(response) + if body is None: + return diagnostics + + found = { + element.lower(): value + for element, value in _S3_ERROR_FIELD_PATTERN.findall(body) + } + for element, field in _S3_ERROR_FIELDS.items(): + value = found.get(element.lower()) + if value is not None: + diagnostics[field] = _redacted_snippet(value) + + if diagnostics["s3_error_code"] is None: + # Non-S3 body (proxy/gateway) + diagnostics["response_body_snippet"] = _redacted_snippet(body) + + return diagnostics + + +def describe_exception(exc: BaseException) -> dict[str, Any]: + """Log-safe fields from a cache-related exception.""" + if isinstance(exc, SqlCacheHttpError): + return dict(exc.diagnostics) + + # HTTPError.message is status + URL; the response body is still available + if isinstance(exc, requests.HTTPError) and exc.response is not None: + return describe_s3_error(exc.response) + + # Truncate before redact: _URL_QUERY_PATTERN is O(n²) on long unbounded text + return { + "error_type": type(exc).__name__, + "error_message": _redacted_snippet(str(exc)[:_MAX_RAW_EXCEPTION_CHARS]), + } diff --git a/deepnote_toolkit/sql/sql_caching.py b/deepnote_toolkit/sql/sql_caching.py index db396e9..c69de4e 100644 --- a/deepnote_toolkit/sql/sql_caching.py +++ b/deepnote_toolkit/sql/sql_caching.py @@ -1,60 +1,73 @@ import hashlib import json import tempfile +import time +from io import BytesIO +from typing import IO, Any, NamedTuple, Optional import pandas as pd import requests from pyarrow import ArrowInvalid, ArrowNotImplementedError +from deepnote_toolkit.sql.sql_cache_diagnostics import ( + SqlCacheHttpError, + describe_exception, + describe_presigned_url, + read_body_snippet, + safe_url_path, + seconds_between, +) from deepnote_toolkit.sql.sql_utils import is_single_select_query from ..get_webapp_url import get_absolute_userpod_api_url, get_project_auth_headers from ..ipython_utils import output_sql_metadata from ..logging import get_logger -# Initialize logger logger = get_logger() +# The read timeout bounds the gap between two received chunks, not the transfer as +# a whole, so it cannot abort a large but healthy download. +_OBJECT_STORE_TIMEOUT: tuple[int, int] = (5, 60) -def get_sql_cache( - query, bind_params, integration_id, sql_cache_mode, return_variable_type -): - """ - Retrieves the SQL cache from webapp for a given query. - Args: - query (str): The SQL query to retrieve the cache for. - bind_params (dict): The bind parameters for the SQL query. - integration_id (str): The integration ID associated with the cache. - sql_cache_mode (str): The mode of the SQL cache. +class SqlCacheUpload(NamedTuple): + """A presigned upload URL together with the moment it was issued.""" - Returns: - tuple: A tuple containing the cached dataframe (if available) and the upload URL (if applicable). - """ + url: str + issued_at: float + + +def get_sql_cache( + query: str, + bind_params: dict, + integration_id: str, + sql_cache_mode: str, + return_variable_type: str, +) -> tuple[Optional[pd.DataFrame], Optional[SqlCacheUpload]]: + """Return a cache hit dataframe and/or a presigned upload for a cache miss.""" if not is_single_select_query(query): - # we only cache single select queries - output_sql_metadata( - { - "status": "cache_not_supported_for_query", - # We don't include the additional metadata as the query hasn't been executed/read from cache - } - ) + output_sql_metadata({"status": "cache_not_supported_for_query"}) return None, None query_hash = _generate_cache_key(query, bind_params) + # Taken before the request because the webapp signs the upload URL somewhere + # inside this round trip, which makes it a conservative upper bound + requested_at = time.monotonic() + cache_info = None try: cache_info = _request_cache_info_from_webapp( query_hash, integration_id, sql_cache_mode ) except Exception as exc: - # we failed to request the cache info from the webapp logger.error( - "Failed to request SQL cache info: %s", - exc, - extra={"sql_caching_cause": "failed_to_request_cache_info"}, + "Failed to request SQL cache info", + extra={ + "sql_caching_cause": "failed_to_request_cache_info", + **describe_exception(exc), + }, ) return None, None @@ -65,11 +78,13 @@ def get_sql_cache( try: dataframe_from_cache = _try_read_cache(download_url) except Exception as exc: - # we failed to download the dataframe from the cache logger.error( - "Failed to download dataframe from cache: %s", - exc, - extra={"sql_caching_cause": "failed_to_download_from_cache"}, + "Failed to download dataframe from cache", + extra={ + "sql_caching_cause": "failed_to_download_from_cache", + **describe_exception(exc), + **describe_presigned_url(download_url), + }, ) return None, None @@ -86,55 +101,94 @@ def get_sql_cache( return dataframe_from_cache, None if cache_info["result"] == "cacheMiss" or cache_info["result"] == "alwaysWrite": - return None, cache_info["uploadUrl"] + return None, SqlCacheUpload( + url=cache_info["uploadUrl"], issued_at=requested_at + ) return None, None -def upload_sql_cache(dataframe, upload_url): - """upload the result to the cache as a parquet file""" +def _serialize_dataframe_for_cache( + dataframe: pd.DataFrame, file_obj: IO[bytes] +) -> None: + """Write the dataframe to file_obj as parquet, falling back to pickle.""" + try: + dataframe.to_parquet(file_obj) + except (ArrowNotImplementedError, ArrowInvalid, OverflowError): + # NB-1684: pickle when parquet cannot represent the frame (e.g. int64 overflow). + file_obj.seek(0) + file_obj.truncate() + dataframe.to_pickle(file_obj) + + +def upload_sql_cache(dataframe: pd.DataFrame, upload: SqlCacheUpload) -> None: + """Best-effort cache upload; failures are logged and never raised.""" + put_started_at: Optional[float] = None try: with tempfile.TemporaryFile() as temp_file: try: - dataframe.to_parquet(temp_file) - except (ArrowNotImplementedError, ArrowInvalid, OverflowError): - # see NB-1684 - # we fallback to pickle if parquet serialization fails (which will throw either of first 2 errors) - # OverflowError: PyArrow raises this for Python int / Decimal values exceeding int64 range - temp_file.seek(0) - temp_file.truncate() - dataframe.to_pickle(temp_file) + _serialize_dataframe_for_cache(dataframe, temp_file) + except Exception as exc: + # Serialization errors embed column names; log type only. + logger.error( + "Failed to upload SQL cache", + extra={ + "sql_caching_cause": "failed_to_serialize_cache", + "error_type": type(exc).__name__, + }, + ) + return temp_file.seek(0) - # PUT the file to cache_upload_url pre-signed s3 url - response = requests.put(upload_url, data=temp_file) - response.raise_for_status() + # Presigned validity is checked when the PUT starts, not when it ends. + put_started_at = time.monotonic() + response = requests.put( + upload.url, data=temp_file, timeout=_OBJECT_STORE_TIMEOUT + ) + + response.raise_for_status() except Exception as exc: logger.error( - "Failed to upload SQL cache: %s", - exc, - extra={"sql_caching_cause": "failed_to_upload_to_cache"}, + "Failed to upload SQL cache", + extra={ + "sql_caching_cause": "failed_to_upload_to_cache", + **describe_exception(exc), + **describe_presigned_url(upload.url), + "seconds_since_url_issued": seconds_between( + upload.issued_at, put_started_at + ), + "upload_duration_seconds": seconds_between( + put_started_at, time.monotonic() + ), + }, ) -def _try_read_cache(download_url): +def _try_read_cache(download_url: str) -> pd.DataFrame: + """Download the cached object and read it as a dataframe. + + The object is fetched explicitly instead of handing the URL to pandas, which + fetches it with urllib and so raises before S3's error body is ever read. + """ + # Read the body before leaving the context: after close, content reads empty. + with requests.get( + download_url, timeout=_OBJECT_STORE_TIMEOUT, stream=True + ) as response: + try: + response.raise_for_status() + except requests.HTTPError as exc: + raise SqlCacheHttpError(describe_exception(exc)) from exc + + buffer = BytesIO(response.content) + try: - # Attempt to read as a parquet file - return pd.read_parquet(download_url) + return pd.read_parquet(buffer) except ArrowInvalid: - # ArrowInvalid means that the file at download_url is not a parquet file. - # We fallback to the pickle format if that happens, because the cache should either be in parquet or - # pickle format and we don't know which one it is, the file has no extension. - # (see .to_pickle fallback in upload_sql_cache) pass - try: - # Attempt to read as a pickle file - return pd.read_pickle(download_url) - except Exception: - # If reading as pickle also fails, re-raise this exception to be caught by the caller - raise + buffer.seek(0) + return pd.read_pickle(buffer) def _generate_cache_key(query, bind_params): @@ -143,7 +197,10 @@ def _generate_cache_key(query, bind_params): ).hexdigest() -def _request_cache_info_from_webapp(query_hash, integration_id, sql_cache_mode): +def _request_cache_info_from_webapp( + query_hash: str, integration_id: str, sql_cache_mode: str +) -> Optional[dict[str, Any]]: + """The cache info for this query, or None when caching is off or unavailable.""" # calls https://github.com/deepnote/deepnote/blob/eb96467937de12db8b588e5aa0a80244cec7eae7/apps/webapp/server/api/userpod-api.ts#L133 sql_cache_url = get_absolute_userpod_api_url( f"integrations/{integration_id}/sql-cache?sqlCacheKey={query_hash}&sqlCacheMode={sql_cache_mode}" @@ -158,8 +215,15 @@ def _request_cache_info_from_webapp(query_hash, integration_id, sql_cache_mode): ) if sql_cache_response.status_code != 200: # the caching endpoint is not available, we can't use it. We'll skip the caching logic - error_msg = f"Failed to request cache info from {sql_cache_url}, status code {sql_cache_response.status_code}, response {sql_cache_response.text}" - logger.error(error_msg, extra={"sql_caching_cause": "http_error"}) + logger.error( + "Failed to request cache info", + extra={ + "sql_caching_cause": "http_error", + "status_code": sql_cache_response.status_code, + "cache_info_path": safe_url_path(sql_cache_url), + "response_body_snippet": read_body_snippet(sql_cache_response), + }, + ) return None result_dict = sql_cache_response.json() diff --git a/deepnote_toolkit/sql/sql_execution.py b/deepnote_toolkit/sql/sql_execution.py index 642a687..6f8daf8 100644 --- a/deepnote_toolkit/sql/sql_execution.py +++ b/deepnote_toolkit/sql/sql_execution.py @@ -439,10 +439,10 @@ def _execute_sql_with_caching( integration_id = sql_alchemy_dict.get("integration_id") can_get_sql_cache = integration_id is not None and sql_caching_enabled - cache_upload_url = None + cache_upload = None if can_get_sql_cache: - dataframe_from_cache, cache_upload_url = get_sql_cache( + dataframe_from_cache, cache_upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) if dataframe_from_cache is not None: @@ -461,7 +461,7 @@ def _execute_sql_with_caching( query_with_audit_comment, bind_params, sql_alchemy_dict, - cache_upload_url, + cache_upload, return_variable_type, query_preview_source, # The original query before any transformations such as appending a LIMIT clause ) @@ -492,7 +492,7 @@ def _query_data_source( query, bind_params, sql_alchemy_dict, - cache_upload_url, + cache_upload, return_variable_type, query_preview_source, ): @@ -536,8 +536,8 @@ def _query_data_source( # if df is larger than 5GB, don't upload it. See NB-988 dataframe_is_cacheable = dataframe_size_in_bytes < 5 * 1024 * 1024 * 1024 - if cache_upload_url is not None and dataframe_is_cacheable: - upload_sql_cache(dataframe, cache_upload_url) + if cache_upload is not None and dataframe_is_cacheable: + upload_sql_cache(dataframe, cache_upload) return dataframe finally: diff --git a/tests/unit/helpers/sql_cache_fixtures.py b/tests/unit/helpers/sql_cache_fixtures.py new file mode 100644 index 0000000..91419a5 --- /dev/null +++ b/tests/unit/helpers/sql_cache_fixtures.py @@ -0,0 +1,104 @@ +from unittest import mock + +import requests + +# Distinct tokens so leak assertions can grep log output. +PRESIGNED_URL = ( + "https://bucket.s3.eu-west-1.amazonaws.com/ws/int/key" + "?X-Amz-Algorithm=AWS4-HMAC-SHA256" + "&X-Amz-Credential=CREDVALUE%2F20260729%2Feu-west-1%2Fs3%2Faws4_request" + "&X-Amz-Date=20260729T120000Z&X-Amz-Expires=900" + "&X-Amz-Security-Token=TOKENVALUE&X-Amz-Signature=SIGVALUE" +) + +# urllib3 embeds a path-only URL in the message (no scheme/host). +URLLIB3_ERROR_MESSAGE = ( + "HTTPSConnectionPool(host='bucket.s3.eu-west-1.amazonaws.com', port=443): " + "Max retries exceeded with url: /ws/int/key" + "?X-Amz-Algorithm=AWS4-HMAC-SHA256" + "&X-Amz-Credential=CREDVALUE%2F20260729%2Feu-west-1%2Fs3%2Faws4_request" + "&X-Amz-Date=20260729T120000Z&X-Amz-Expires=900" + "&X-Amz-Security-Token=TOKENVALUE&X-Amz-Signature=SIGVALUE " + "(Caused by NameResolutionError('Failed to resolve host'))" +) + +SECRETS = ("CREDVALUE", "TOKENVALUE", "SIGVALUE", "X-Amz-") + +AWS_HEADERS = { + "x-amz-request-id": "REQ123", + "x-amz-id-2": "HOSTID456", + "Date": "Wed, 29 Jul 2026 12:31:07 GMT", +} + +ACCESS_DENIED_EXPIRED_BODY = ( + b'\n' + b"AccessDeniedRequest has expired" + b"900" + b"2026-07-29T12:15:00Z" + b"2026-07-29T12:31:07Z" + b"REQ123HOSTID456" +) + +EXPIRED_TOKEN_BODY = ( + b'\n' + b"ExpiredToken" + b"The provided token has expired." + b"REQ123HOSTID456" +) + +# Real S3 bodies can embed the signed query in . +SIGNATURE_MISMATCH_BODY = ( + b'\n' + b"SignatureDoesNotMatch" + b"The request signature we calculated does not match the signature " + b"you provided. Check your key and signing method." + b"CREDVALUE" + b"AWS4-HMAC-SHA256\n20260729T120000Z\n" + b"PUT\n/ws/int/key\n" + b"X-Amz-Credential=CREDVALUE&X-Amz-Security-Token=TOKENVALUE" + b"&X-Amz-Signature=SIGVALUE\nhost:bucket.s3.eu-west-1.amazonaws.com\n" + b"" + b"REQ123HOSTID456" +) + +# Non-S3 gateways may echo the full presigned URL in HTML. +PROXY_ECHO_BODY = ( + b"502 Bad Gateway\n" + b"

502 Bad Gateway

\n" + b"

Upstream failed for request: PUT " + b"https://bucket.s3.eu-west-1.amazonaws.com/ws/int/key" + b"?X-Amz-Algorithm=AWS4-HMAC-SHA256" + b"&X-Amz-Credential=CREDVALUE%2F20260729%2Feu-west-1%2Fs3%2Faws4_request" + b"&X-Amz-Date=20260729T120000Z&X-Amz-Expires=900" + b"&X-Amz-Security-Token=TOKENVALUE&X-Amz-Signature=SIGVALUE

\n" + b"" +) + + +def s3_response(status_code, body=b"", headers=None, url=PRESIGNED_URL): + response = mock.MagicMock( + status_code=status_code, content=body, headers=headers or {}, url=url + ) + released = False + + def release(*_): + nonlocal released + released = True + response.content = b"" + + def iter_content(size): + if released: + return iter([]) + return iter([body[i : i + 20] for i in range(0, len(body), 20)]) + + response.__enter__.return_value = response + # After context exit, requests clears streamed body content. + response.__exit__.side_effect = release + response.iter_content.side_effect = iter_content + if status_code >= 400: + response.raise_for_status.side_effect = requests.HTTPError( + f"{status_code} Client Error: for url: {url}", response=response + ) + else: + response.raise_for_status.return_value = None + return response diff --git a/tests/unit/test_sql_cache_diagnostics.py b/tests/unit/test_sql_cache_diagnostics.py new file mode 100644 index 0000000..203fe1d --- /dev/null +++ b/tests/unit/test_sql_cache_diagnostics.py @@ -0,0 +1,250 @@ +import json +import time +import unittest +from unittest import mock + +from parameterized import parameterized + +from deepnote_toolkit.sql.sql_cache_diagnostics import ( + _URL_QUERY_PATTERN, + describe_exception, + describe_presigned_url, + describe_s3_error, + redact_sensitive, + safe_url_path, +) + +from .helpers.sql_cache_fixtures import ( + ACCESS_DENIED_EXPIRED_BODY, + AWS_HEADERS, + EXPIRED_TOKEN_BODY, + PRESIGNED_URL, + PROXY_ECHO_BODY, + SECRETS, + SIGNATURE_MISMATCH_BODY, + URLLIB3_ERROR_MESSAGE, + s3_response, +) + + +class TestRedactSensitive(unittest.TestCase): + def test_strips_query_string_from_urlopen_style_message(self): + redacted = redact_sensitive(URLLIB3_ERROR_MESSAGE) + + self.assertIn("bucket.s3.eu-west-1.amazonaws.com", redacted) + self.assertIn("/ws/int/key?", redacted) + for secret in SECRETS: + self.assertNotIn(secret, redacted) + + def test_query_string_strip_alone_redacts_urllib3_message(self): + """URL query stripping must redact without the AWS-param backstop.""" + stripped = _URL_QUERY_PATTERN.sub(r"\1?", URLLIB3_ERROR_MESSAGE) + + for secret in SECRETS: + self.assertNotIn(secret, stripped) + + def test_blanks_aws_params_outside_a_url(self): + redacted = redact_sensitive( + "X-Amz-Credential=CREDVALUE&X-Amz-Security-Token=TOKENVALUE" + "&X-Amz-Signature=SIGVALUE" + ) + + self.assertEqual( + redacted, + "X-Amz-Credential=&X-Amz-Security-Token=" + "&X-Amz-Signature=", + ) + + @parameterized.expand( + [ + ("question_in_prose", "Is this ok? Yes it is."), + ("bare_question", "what? nothing"), + ("plain_sentence", "The provided token has expired."), + ] + ) + def test_leaves_ordinary_text_unchanged(self, _, text): + self.assertEqual(redact_sensitive(text), text) + + +class TestDescribeException(unittest.TestCase): + def test_large_message_is_bounded_and_redacted_in_bounded_time(self): + """Truncate before redact: _URL_QUERY_PATTERN is quadratic on long strings.""" + message = URLLIB3_ERROR_MESSAGE + "&padding=" + "A" * 64_000 + + started = time.monotonic() + described = describe_exception(ValueError(message)) + elapsed = time.monotonic() - started + + self.assertEqual(described["error_type"], "ValueError") + self.assertLessEqual(len(described["error_message"]), 500) + self.assertIn("/ws/int/key?", described["error_message"]) + for secret in SECRETS: + self.assertNotIn(secret, described["error_message"]) + self.assertLess(elapsed, 2.0) + + +class TestDescribeS3Error(unittest.TestCase): + def test_extracts_code_and_message_from_xml(self): + diagnostics = describe_s3_error( + s3_response(403, ACCESS_DENIED_EXPIRED_BODY, AWS_HEADERS) + ) + + self.assertEqual(diagnostics["status_code"], 403) + self.assertEqual(diagnostics["s3_error_code"], "AccessDenied") + self.assertIn("Request has expired", diagnostics["s3_error_message"]) + self.assertEqual(diagnostics["s3_expires"], "2026-07-29T12:15:00Z") + self.assertEqual(diagnostics["s3_server_time"], "2026-07-29T12:31:07Z") + + def test_extracts_expired_token_body(self): + diagnostics = describe_s3_error(s3_response(403, EXPIRED_TOKEN_BODY)) + + self.assertEqual(diagnostics["s3_error_code"], "ExpiredToken") + self.assertEqual( + diagnostics["s3_error_message"], "The provided token has expired." + ) + + def test_captures_aws_request_headers(self): + diagnostics = describe_s3_error( + s3_response(403, EXPIRED_TOKEN_BODY, AWS_HEADERS) + ) + + self.assertEqual(diagnostics["aws_request_id"], "REQ123") + self.assertEqual(diagnostics["aws_host_id"], "HOSTID456") + self.assertEqual(diagnostics["aws_date"], "Wed, 29 Jul 2026 12:31:07 GMT") + + def test_signature_mismatch_body_surfaces_only_code_and_message(self): + """Field allowlist must ignore (signed query lives there).""" + diagnostics = describe_s3_error( + s3_response(403, SIGNATURE_MISMATCH_BODY, AWS_HEADERS) + ) + + self.assertEqual(diagnostics["s3_error_code"], "SignatureDoesNotMatch") + self.assertNotIn("response_body_snippet", diagnostics) + for value in diagnostics.values(): + for secret in SECRETS: + self.assertNotIn(secret, str(value)) + + def test_proxy_body_echoing_request_url_is_redacted(self): + """Non-S3 bodies are not on the XML allowlist; snippet redaction must apply.""" + diagnostics = describe_s3_error(s3_response(502, PROXY_ECHO_BODY)) + + self.assertIsNone(diagnostics["s3_error_code"]) + snippet = diagnostics["response_body_snippet"] + self.assertIn("502 Bad Gateway", snippet) + self.assertIn("/ws/int/key?", snippet) + for secret in SECRETS: + self.assertNotIn(secret, snippet) + + def test_non_xml_body_yields_snippet_without_code(self): + diagnostics = describe_s3_error( + s3_response(502, b"502 Bad Gateway") + ) + + self.assertIsNone(diagnostics["s3_error_code"]) + self.assertIsNone(diagnostics["aws_request_id"]) + self.assertIn("502 Bad Gateway", diagnostics["response_body_snippet"]) + self.assertLessEqual(len(diagnostics["response_body_snippet"]), 500) + + def test_oversized_body_is_bounded(self): + body = ( + b"AccessDenied" + + b"x" * 1000 + + b"" + + b"y" * 10_000 + + b"" + ) + + diagnostics = describe_s3_error(s3_response(403, body)) + + self.assertEqual(diagnostics["s3_error_code"], "AccessDenied") + for value in diagnostics.values(): + if isinstance(value, str): + self.assertLessEqual(len(value), 500) + + def test_body_prefix_is_streamed_rather_than_buffered(self): + """Must not touch Response.content (downloads the full body).""" + buffered = [] + response = mock.MagicMock(status_code=403, headers={}) + type(response).content = mock.PropertyMock( + side_effect=lambda: buffered.append("content") + ) + response.iter_content.side_effect = lambda size: iter( + [ACCESS_DENIED_EXPIRED_BODY[:size]] + ) + + diagnostics = describe_s3_error(response) + + self.assertEqual(buffered, []) + self.assertEqual(diagnostics["s3_error_code"], "AccessDenied") + + +class TestDescribePresignedUrl(unittest.TestCase): + def test_returns_path_and_expiry(self): + described = describe_presigned_url(PRESIGNED_URL) + + self.assertEqual(described["object_host"], "bucket.s3.eu-west-1.amazonaws.com") + self.assertEqual(described["object_path"], "/ws/int/key") + self.assertEqual(described["url_expires_in"], 900) + + def test_missing_expires_yields_none(self): + described = describe_presigned_url("https://example.com/x") + + self.assertEqual(described["object_path"], "/x") + self.assertIsNone(described["url_expires_in"]) + + def test_non_numeric_expires_yields_none(self): + described = describe_presigned_url("https://example.com/x?X-Amz-Expires=abc") + + self.assertIsNone(described["url_expires_in"]) + + @parameterized.expand( + [ + ("unterminated_ipv6", "https://[::1"), + ("empty", ""), + ("not_a_url", "not a url"), + ("none", None), + ("bytes", b"/ws/int/key"), + ("dict", {"url": "https://example.com/x"}), + ] + ) + def test_malformed_url_does_not_raise_or_leak(self, _, url): + described = describe_presigned_url(url) + + self.assertIsNone(described["url_expires_in"]) + for value in described.values(): + for secret in SECRETS: + self.assertNotIn(secret, str(value)) + # Log extras must json.dumps; bytes values would drop the whole report. + json.dumps(described) + json.dumps(safe_url_path(url)) + + @parameterized.expand( + [ + ( + "separator_encoded", + "%3FX-Amz-Credential=CREDVALUE&X-Amz-Security-Token=TOKENVALUE" + "&X-Amz-Signature=SIGVALUE", + ), + ( + "separator_and_equals_encoded", + "%3FX-Amz-Credential%3DCREDVALUE&X-Amz-Security-Token%3DTOKENVALUE" + "&X-Amz-Signature%3DSIGVALUE", + ), + ( + "whole_query_encoded", + "%3FX-Amz-Credential%3DCREDVALUE%26X-Amz-Security-Token%3DTOKENVALUE" + "%26X-Amz-Signature%3DSIGVALUE", + ), + ("separator_is_a_semicolon", ";X-Amz-Credential=CREDVALUE"), + ("no_separator_at_all", "X-Amz-Signature=SIGVALUE"), + ] + ) + def test_over_encoded_url_does_not_leak_signing_params_via_path(self, _, suffix): + """urlsplit only splits on a literal '?'; signing params can land in path.""" + described = describe_presigned_url( + "https://bucket.s3.eu-west-1.amazonaws.com/ws/int/key" + suffix + ) + + for value in described.values(): + for secret in ("CREDVALUE", "TOKENVALUE", "SIGVALUE"): + self.assertNotIn(secret, str(value)) diff --git a/tests/unit/test_sql_caching.py b/tests/unit/test_sql_caching.py index 7e6f287..aecd82f 100644 --- a/tests/unit/test_sql_caching.py +++ b/tests/unit/test_sql_caching.py @@ -1,24 +1,129 @@ +import json +import logging +import time import unittest +from io import BytesIO from unittest import mock from unittest.mock import patch import pandas as pd +import requests from parameterized import parameterized from pyarrow import ArrowInvalid from deepnote_toolkit.sql.sql_caching import ( + SqlCacheUpload, _generate_cache_key, + _request_cache_info_from_webapp, get_sql_cache, upload_sql_cache, ) from deepnote_toolkit.sql.sql_utils import is_single_select_query +from .helpers.sql_cache_fixtures import ( + ACCESS_DENIED_EXPIRED_BODY, + AWS_HEADERS, + EXPIRED_TOKEN_BODY, + PRESIGNED_URL, + PROXY_ECHO_BODY, + SECRETS, + SIGNATURE_MISMATCH_BODY, + URLLIB3_ERROR_MESSAGE, + s3_response, +) + +QUERY = "SELECT * FROM users" + +RESERVED_LOGRECORD_ATTRS = set( + logging.LogRecord("", 0, "", 0, "", None, None).__dict__ +) | {"message", "asctime"} + + +def _upload(url=PRESIGNED_URL, issued_at=None): + return SqlCacheUpload( + url=url, issued_at=time.monotonic() if issued_at is None else issued_at + ) + + +def _cache_hit(download_url=PRESIGNED_URL): + return { + "result": "cacheHit", + "downloadUrl": download_url, + "cacheCreatedAt": "2022-01-01 00:00:00", + } + + +def _logged_strings(mock_logger): + strings = [] + for call in mock_logger.error.call_args_list: + strings.extend(str(arg) for arg in call.args) + extra = call.kwargs.get("extra", {}) + strings.extend(str(key) for key in extra) + strings.extend(str(value) for value in extra.values()) + return strings + + +def _collect_logged_extras(): + """Exercise every error path once for shared extra-dict guards.""" + dataframe = pd.DataFrame({"a": [1, 2, 3]}) + connection_error = requests.exceptions.ConnectionError(URLLIB3_ERROR_MESSAGE) + + with patch("deepnote_toolkit.sql.sql_caching.logger") as mock_logger: + with patch("deepnote_toolkit.sql.sql_caching.requests.put") as mock_put: + mock_put.return_value = s3_response( + 403, ACCESS_DENIED_EXPIRED_BODY, AWS_HEADERS + ) + upload_sql_cache(dataframe, _upload()) + + upload_sql_cache(dataframe, SqlCacheUpload(url=None, issued_at=0.0)) + + mock_put.side_effect = connection_error + upload_sql_cache(dataframe, _upload()) + + unserializable = mock.Mock() + unserializable.to_parquet.side_effect = ValueError("column customer_email") + upload_sql_cache(unserializable, _upload()) + + with patch( + "deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp" + ) as mock_cache_info: + mock_cache_info.return_value = _cache_hit() + with patch("deepnote_toolkit.sql.sql_caching.requests.get") as mock_get: + mock_get.return_value = s3_response( + 403, EXPIRED_TOKEN_BODY, AWS_HEADERS + ) + get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + mock_get.side_effect = connection_error + get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + mock_cache_info.side_effect = connection_error + get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + with ( + patch("deepnote_toolkit.sql.sql_caching.requests.get") as mock_get, + patch( + "deepnote_toolkit.sql.sql_caching.get_absolute_userpod_api_url" + ) as mock_url, + patch( + "deepnote_toolkit.sql.sql_caching.get_project_auth_headers" + ) as mock_headers, + ): + mock_url.return_value = ( + "http://localhost:19456/userpod-api/p1/integrations/123/sql-cache" + "?sqlCacheKey=abc&sqlCacheMode=read" + ) + mock_headers.return_value = {} + mock_get.return_value = s3_response(503, b"upstream unavailable") + _request_cache_info_from_webapp("abc", "123", "read") + + return [call.kwargs["extra"] for call in mock_logger.error.call_args_list] + class TestGenerateCacheKey(unittest.TestCase): def test_empty_params_returns_valid_result(self): result = _generate_cache_key("SELECT * FROM users", {}) - # assert that the result contains only alphanumeric characters self.assertTrue(result.isalnum()) def test_different_order_of_params_produces_same_result(self): @@ -83,7 +188,7 @@ def test_cache_not_supported_for_query( mock_is_single_select_query.return_value = False - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) @@ -91,8 +196,9 @@ def test_cache_not_supported_for_query( {"status": "cache_not_supported_for_query"} ) self.assertIsNone(result_df) - self.assertIsNone(upload_url) + self.assertIsNone(upload) + @patch("deepnote_toolkit.sql.sql_caching.logger") @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") @@ -101,6 +207,7 @@ def test_failed_to_request_cache_info( mock_output_sql_metadata, mock_request_cache_info_from_webapp, mock_is_single_select_query, + mock_logger, ): query = "SELECT * FROM users" bind_params = {} @@ -113,21 +220,23 @@ def test_failed_to_request_cache_info( "Failed to request cache info" ) - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) mock_output_sql_metadata.assert_not_called() self.assertIsNone(result_df) - self.assertIsNone(upload_url) + self.assertIsNone(upload) @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") @patch("pandas.read_parquet") def test_read_from_cache_success( self, mock_read_parquet, + mock_get, mock_output_sql_metadata, mock_request_cache_info_from_webapp, mock_is_single_select_query, @@ -145,9 +254,10 @@ def test_read_from_cache_success( mock_is_single_select_query.return_value = True mock_request_cache_info_from_webapp.return_value = cache_info + mock_get.return_value = s3_response(200, b"parquet-bytes") mock_read_parquet.return_value = pd.DataFrame() - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) @@ -161,17 +271,19 @@ def test_read_from_cache_success( } ) self.assertIsInstance(result_df, pd.DataFrame) - self.assertIsNone(upload_url) + self.assertIsNone(upload) @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") @patch("pandas.read_parquet") @patch("pandas.read_pickle") def test_fallback_to_pickle_format( self, mock_read_pickle, mock_read_parquet, + mock_get, mock_output_sql_metadata, mock_request_cache_info_from_webapp, mock_is_single_select_query, @@ -189,10 +301,11 @@ def test_fallback_to_pickle_format( mock_is_single_select_query.return_value = True mock_request_cache_info_from_webapp.return_value = cache_info + mock_get.return_value = s3_response(200, b"pickle-bytes") mock_read_parquet.side_effect = ArrowInvalid mock_read_pickle.return_value = pd.DataFrame() - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, @@ -210,18 +323,22 @@ def test_fallback_to_pickle_format( } ) self.assertIsInstance(result_df, pd.DataFrame) - self.assertIsNone(upload_url) + self.assertIsNone(upload) + @patch("deepnote_toolkit.sql.sql_caching.logger") @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") @patch("pandas.read_parquet") def test_failed_to_download_from_cache( self, mock_read_parquet, + mock_get, mock_output_sql_metadata, mock_request_cache_info_from_webapp, mock_is_single_select_query, + mock_logger, ): query = "SELECT * FROM users" bind_params = {} @@ -236,14 +353,15 @@ def test_failed_to_download_from_cache( mock_is_single_select_query.return_value = True mock_request_cache_info_from_webapp.return_value = cache_info + mock_get.return_value = s3_response(200, b"parquet-bytes") mock_read_parquet.side_effect = Exception("Failed to download from cache") - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) self.assertIsNone(result_df) - self.assertIsNone(upload_url) + self.assertIsNone(upload) @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @@ -263,12 +381,14 @@ def test_cache_miss( mock_is_single_select_query.return_value = True mock_request_cache_info_from_webapp.return_value = cache_info - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) self.assertIsNone(result_df) - self.assertEqual(upload_url, cache_info["uploadUrl"]) + self.assertIsInstance(upload, SqlCacheUpload) + self.assertEqual(upload.url, cache_info["uploadUrl"]) + self.assertIsInstance(upload.issued_at, float) @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @@ -288,12 +408,14 @@ def test_always_write( mock_is_single_select_query.return_value = True mock_request_cache_info_from_webapp.return_value = cache_info - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) self.assertIsNone(result_df) - self.assertEqual(upload_url, cache_info["uploadUrl"]) + self.assertIsInstance(upload, SqlCacheUpload) + self.assertEqual(upload.url, cache_info["uploadUrl"]) + self.assertIsInstance(upload.issued_at, float) @patch("deepnote_toolkit.sql.sql_caching.is_single_select_query") @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") @@ -309,69 +431,250 @@ def test_no_cache_info( mock_is_single_select_query.return_value = True mock_request_cache_info_from_webapp.return_value = None - result_df, upload_url = get_sql_cache( + result_df, upload = get_sql_cache( query, bind_params, integration_id, sql_cache_mode, return_variable_type ) self.assertIsNone(result_df) - self.assertIsNone(upload_url) + self.assertIsNone(upload) + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") @patch("pandas.read_parquet") @patch("pandas.read_pickle") def test_read_from_cache_error_doesnt_raise( - self, mock_read_pickle, mock_read_parquet + self, + mock_read_pickle, + mock_read_parquet, + mock_get, + mock_cache_info, + mock_logger, ): + mock_cache_info.return_value = _cache_hit("https://example.com/cache") + mock_get.return_value = s3_response(200, b"not-a-dataframe") mock_read_parquet.side_effect = ArrowInvalid mock_read_pickle.side_effect = Exception("Error reading pickle") - query = "SELECT * FROM users" - bind_params = {} - integration_id = "123" - sql_cache_mode = "read" - return_variable_type = "dataframe" + result_df, upload = get_sql_cache(QUERY, {}, "123", "read", "dataframe") - result_df, upload_url = get_sql_cache( - query, bind_params, integration_id, sql_cache_mode, return_variable_type + mock_read_parquet.assert_called_once() + mock_read_pickle.assert_called_once() + self.assertIsNone(result_df) + self.assertIsNone(upload) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_download_http_error_logs_s3_diagnostics( + self, mock_get, mock_cache_info, mock_logger + ): + mock_cache_info.return_value = _cache_hit() + mock_get.return_value = s3_response(403, EXPIRED_TOKEN_BODY, AWS_HEADERS) + + result_df, upload = get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + self.assertIsNone(result_df) + self.assertIsNone(upload) + mock_logger.error.assert_called_once() + self.assertEqual( + mock_logger.error.call_args.args, + ("Failed to download dataframe from cache",), ) + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["sql_caching_cause"], "failed_to_download_from_cache") + self.assertEqual(extra["s3_error_code"], "ExpiredToken") + self.assertEqual(extra["status_code"], 403) + self.assertEqual(extra["aws_request_id"], "REQ123") + self.assertEqual(extra["object_path"], "/ws/int/key") + + @parameterized.expand( + [ + ("signature_mismatch_body", SIGNATURE_MISMATCH_BODY), + ("proxy_echoing_request_url", PROXY_ECHO_BODY), + ] + ) + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_download_http_error_logs_no_query_string( + self, _name, body, mock_get, mock_cache_info, mock_logger + ): + mock_cache_info.return_value = _cache_hit() + mock_get.return_value = s3_response(403, body, AWS_HEADERS) + + get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + mock_logger.error.assert_called_once() + for logged in _logged_strings(mock_logger): + for secret in SECRETS: + self.assertNotIn(secret, logged) + + @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_download_reads_parquet_from_fetched_bytes( + self, mock_get, mock_cache_info, mock_output_sql_metadata + ): + dataframe = pd.DataFrame({"a": [1, 2, 3]}) + buffer = BytesIO() + dataframe.to_parquet(buffer) + + mock_cache_info.return_value = _cache_hit() + mock_get.return_value = s3_response(200, buffer.getvalue()) + + result_df, _ = get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + pd.testing.assert_frame_equal(result_df, dataframe) + mock_get.assert_called_once() + self.assertEqual(mock_get.call_args.args[0], PRESIGNED_URL) + + @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_download_falls_back_to_pickle_from_same_bytes( + self, mock_get, mock_cache_info, mock_output_sql_metadata + ): + dataframe = pd.DataFrame({"a": [1, 2, 3]}) + buffer = BytesIO() + dataframe.to_pickle(buffer) + + mock_cache_info.return_value = _cache_hit() + mock_get.return_value = s3_response(200, buffer.getvalue()) + + result_df, _ = get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + pd.testing.assert_frame_equal(result_df, dataframe) + mock_get.assert_called_once() + + @patch("deepnote_toolkit.sql.sql_caching.output_sql_metadata") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_download_request_is_bounded_and_streamed( + self, mock_get, mock_cache_info, mock_output_sql_metadata + ): + dataframe = pd.DataFrame({"a": [1, 2, 3]}) + buffer = BytesIO() + dataframe.to_parquet(buffer) + + mock_cache_info.return_value = _cache_hit() + mock_get.return_value = s3_response(200, buffer.getvalue()) + + get_sql_cache(QUERY, {}, "123", "read", "dataframe") + + connect_timeout, read_timeout = mock_get.call_args.kwargs["timeout"] + self.assertIsNotNone(connect_timeout) + self.assertIsNotNone(read_timeout) + self.assertTrue(mock_get.call_args.kwargs["stream"]) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching._request_cache_info_from_webapp") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_download_network_error_is_swallowed( + self, mock_get, mock_cache_info, mock_logger + ): + mock_cache_info.return_value = _cache_hit() + mock_get.side_effect = requests.exceptions.ConnectionError( + URLLIB3_ERROR_MESSAGE + ) + + result_df, upload = get_sql_cache(QUERY, {}, "123", "read", "dataframe") self.assertIsNone(result_df) - self.assertIsNone(upload_url) + self.assertIsNone(upload) + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["error_type"], "ConnectionError") + for logged in _logged_strings(mock_logger): + for secret in SECRETS: + self.assertNotIn(secret, logged) + + +class TestRequestCacheInfoFromWebapp(unittest.TestCase): + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.get_project_auth_headers") + @patch("deepnote_toolkit.sql.sql_caching.get_absolute_userpod_api_url") + @patch("deepnote_toolkit.sql.sql_caching.requests.get") + def test_non_200_logs_constant_message_with_status_in_extra( + self, mock_get, mock_url, mock_headers, mock_logger + ): + mock_url.return_value = ( + "http://localhost:19456/userpod-api/p1/integrations/123/sql-cache" + "?sqlCacheKey=abc&sqlCacheMode=read" + ) + mock_headers.return_value = {} + mock_get.return_value = s3_response(503, b"upstream unavailable") + + self.assertIsNone(_request_cache_info_from_webapp("abc", "123", "read")) + + self.assertEqual( + mock_logger.error.call_args.args, ("Failed to request cache info",) + ) + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["status_code"], 503) + self.assertEqual(extra["response_body_snippet"], "upstream unavailable") + self.assertEqual( + extra["cache_info_path"], + "/userpod-api/p1/integrations/123/sql-cache", + ) + self.assertNotIn("sqlCacheKey", " ".join(_logged_strings(mock_logger))) class TestUploadSqlCache(unittest.TestCase): + @patch("deepnote_toolkit.sql.sql_caching.logger") @patch("deepnote_toolkit.sql.sql_caching.requests.put") - def test_upload_parquet_success(self, mock_put): - mock_put.return_value = mock.Mock(raise_for_status=mock.Mock()) + def test_upload_parquet_success(self, mock_put, mock_logger): + mock_put.return_value = mock.Mock(status_code=200) df = pd.DataFrame({"a": [1, 2, 3]}) - upload_sql_cache(df, "https://example.com/upload") + upload_sql_cache(df, _upload("https://example.com/upload")) mock_put.assert_called_once() args, _ = mock_put.call_args self.assertEqual(args[0], "https://example.com/upload") + mock_logger.error.assert_not_called() @patch("deepnote_toolkit.sql.sql_caching.requests.put") def test_overflow_error_falls_back_to_pickle(self, mock_put): """Large Python int triggers OverflowError in to_parquet, upload succeeds via pickle.""" uploaded_bytes = None - def capture_put(_url, data): + def capture_put(_url, data, **_kwargs): nonlocal uploaded_bytes uploaded_bytes = data.read() - return mock.Mock(raise_for_status=mock.Mock()) + return mock.Mock(status_code=200) mock_put.side_effect = capture_put df = pd.DataFrame({"x": pd.array([2**100, 1], dtype=object)}) - upload_sql_cache(df, "https://example.com/upload") + upload_sql_cache(df, _upload("https://example.com/upload")) roundtripped = pd.read_pickle(pd.io.common.BytesIO(uploaded_bytes)) pd.testing.assert_frame_equal(roundtripped, df) + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_arrow_failure_still_falls_back_to_pickle_and_uploads(self, mock_put): + """The narrowed serialization handler must not swallow the pickle fallback.""" + uploaded_bytes = None + + def capture_put(_url, data, **_kwargs): + nonlocal uploaded_bytes + uploaded_bytes = data.read() + return mock.Mock(status_code=200) + + mock_put.side_effect = capture_put + # nested dicts of mixed shape are not representable in parquet + df = pd.DataFrame({"x": [{"a": 1}, {"a": "two"}]}) + + upload_sql_cache(df, _upload("https://example.com/upload")) + + mock_put.assert_called_once() + roundtripped = pd.read_pickle(BytesIO(uploaded_bytes)) + pd.testing.assert_frame_equal(roundtripped, df) + @patch("deepnote_toolkit.sql.sql_caching.requests.put") def test_pickle_fallback_truncates_partial_parquet_bytes(self, mock_put): """When to_parquet writes partial bytes before failing, truncate clears them.""" - mock_put.return_value = mock.Mock(raise_for_status=mock.Mock()) + mock_put.return_value = mock.Mock(status_code=200) def write_garbage_then_overflow(f, **_kwargs): f.write(b"partial parquet data") @@ -390,7 +693,204 @@ def capture_file_state(f, **_kwargs): df.to_parquet.side_effect = write_garbage_then_overflow df.to_pickle.side_effect = capture_file_state - upload_sql_cache(df, "https://example.com/upload") + upload_sql_cache(df, _upload("https://example.com/upload")) self.assertEqual(pickle_pos, 0, "file should be at position 0") self.assertEqual(pickle_size, 0, "file should be empty after truncate") + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_http_error_logs_s3_diagnostics(self, mock_put, mock_logger): + mock_put.return_value = s3_response( + 403, ACCESS_DENIED_EXPIRED_BODY, AWS_HEADERS + ) + + upload_sql_cache(pd.DataFrame({"a": [1, 2, 3]}), _upload()) + + mock_logger.error.assert_called_once() + self.assertEqual( + mock_logger.error.call_args.args, ("Failed to upload SQL cache",) + ) + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["sql_caching_cause"], "failed_to_upload_to_cache") + self.assertEqual(extra["s3_error_code"], "AccessDenied") + self.assertEqual(extra["s3_error_message"], "Request has expired") + self.assertEqual(extra["status_code"], 403) + self.assertEqual(extra["aws_request_id"], "REQ123") + self.assertEqual(extra["aws_host_id"], "HOSTID456") + self.assertEqual(extra["object_path"], "/ws/int/key") + self.assertEqual(extra["url_expires_in"], 900) + self.assertGreaterEqual(extra["seconds_since_url_issued"], 0) + + @parameterized.expand( + [ + ("signature_mismatch_body", SIGNATURE_MISMATCH_BODY), + ("proxy_echoing_request_url", PROXY_ECHO_BODY), + ] + ) + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_http_error_logs_no_presigned_query_string( + self, _name, body, mock_put, mock_logger + ): + mock_put.return_value = s3_response(403, body, AWS_HEADERS) + + upload_sql_cache(pd.DataFrame({"a": [1, 2, 3]}), _upload()) + + mock_logger.error.assert_called_once() + for logged in _logged_strings(mock_logger): + for secret in SECRETS: + self.assertNotIn(secret, logged) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_network_error_logs_no_presigned_url(self, mock_put, mock_logger): + mock_put.side_effect = requests.exceptions.ConnectionError( + URLLIB3_ERROR_MESSAGE + ) + + upload_sql_cache(pd.DataFrame({"a": [1, 2, 3]}), _upload()) + + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["sql_caching_cause"], "failed_to_upload_to_cache") + self.assertEqual(extra["error_type"], "ConnectionError") + for logged in _logged_strings(mock_logger): + for secret in SECRETS: + self.assertNotIn(secret, logged) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_message_is_constant_across_different_failures(self, mock_put, mock_logger): + df = pd.DataFrame({"a": [1, 2, 3]}) + other_url = PRESIGNED_URL.replace("/ws/int/key", "/other/int/key2") + + mock_put.return_value = s3_response(403, ACCESS_DENIED_EXPIRED_BODY) + upload_sql_cache(df, _upload()) + mock_put.return_value = s3_response(500, b"Slow") + upload_sql_cache(df, _upload(other_url)) + mock_put.side_effect = requests.exceptions.ConnectionError( + URLLIB3_ERROR_MESSAGE + ) + upload_sql_cache(df, _upload()) + upload_sql_cache(df, _upload(other_url)) + + messages = {call.args for call in mock_logger.error.call_args_list} + self.assertEqual(messages, {("Failed to upload SQL cache",)}) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_serialization_failure_uses_distinct_cause(self, mock_put, mock_logger): + df = mock.Mock() + df.to_parquet.side_effect = ValueError("Conversion failed for column secret") + df.to_pickle.side_effect = ValueError("Conversion failed for column secret") + + upload_sql_cache(df, _upload()) + + mock_put.assert_not_called() + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["sql_caching_cause"], "failed_to_serialize_cache") + self.assertEqual(extra["error_type"], "ValueError") + self.assertNotIn("error_message", extra) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_seconds_since_url_issued_reflects_elapsed_time( + self, mock_put, mock_logger + ): + mock_put.return_value = s3_response(403, ACCESS_DENIED_EXPIRED_BODY) + + upload_sql_cache( + pd.DataFrame({"a": [1, 2, 3]}), + _upload(issued_at=time.monotonic() - 930), + ) + + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertGreaterEqual(extra["seconds_since_url_issued"], 930) + self.assertEqual(extra["url_expires_in"], 900) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.tempfile.TemporaryFile") + def test_failure_before_the_request_leaves_both_times_unset( + self, mock_temp_file, mock_logger + ): + mock_temp_file.side_effect = OSError("no space left on device") + + upload_sql_cache(pd.DataFrame({"a": [1, 2, 3]}), _upload()) + + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertIsNone(extra["seconds_since_url_issued"]) + self.assertIsNone(extra["upload_duration_seconds"]) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + @patch("deepnote_toolkit.sql.sql_caching.time.monotonic") + def test_transfer_time_is_not_counted_as_url_age( + self, mock_monotonic, mock_put, mock_logger + ): + clock = [1000.0] + mock_monotonic.side_effect = lambda: clock[0] + + def slow_put(*args, **kwargs): + clock[0] += 1000.0 + return s3_response(500, b"InternalError") + + mock_put.side_effect = slow_put + + upload_sql_cache( + pd.DataFrame({"a": [1, 2, 3]}), + SqlCacheUpload(url=PRESIGNED_URL, issued_at=100.5), + ) + + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["seconds_since_url_issued"], 899.5) + self.assertEqual(extra["upload_duration_seconds"], 1000.0) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_upload_request_bounds_connect_and_read_phases(self, mock_put, mock_logger): + mock_put.return_value = s3_response(200) + + upload_sql_cache(pd.DataFrame({"a": [1, 2, 3]}), _upload()) + + connect_timeout, read_timeout = mock_put.call_args.kwargs["timeout"] + self.assertIsNotNone(connect_timeout) + self.assertIsNotNone(read_timeout) + + @patch("deepnote_toolkit.sql.sql_caching.logger") + @patch("deepnote_toolkit.sql.sql_caching.requests.put") + def test_upload_timeout_never_raises(self, mock_put, mock_logger): + mock_put.side_effect = requests.exceptions.Timeout(URLLIB3_ERROR_MESSAGE) + + upload_sql_cache(pd.DataFrame({"a": [1, 2, 3]}), _upload()) + + extra = mock_logger.error.call_args.kwargs["extra"] + self.assertEqual(extra["error_type"], "Timeout") + + +class TestLoggedExtras(unittest.TestCase): + def setUp(self): + self.extras = _collect_logged_extras() + + def test_every_failure_path_logs_a_cause(self): + self.assertEqual( + {extra["sql_caching_cause"] for extra in self.extras}, + { + "failed_to_upload_to_cache", + "failed_to_serialize_cache", + "failed_to_download_from_cache", + "failed_to_request_cache_info", + "http_error", + }, + ) + + def test_extra_keys_avoid_reserved_logrecord_attributes(self): + """Reserved extra keys make logger.error raise into the user's cell.""" + for extra in self.extras: + self.assertEqual(set(extra) & RESERVED_LOGRECORD_ATTRS, set()) + + def test_extra_is_json_serializable(self): + """Non-JSON-serializable extras drop the whole error report.""" + for extra in self.extras: + for value in extra.values(): + self.assertIsInstance(value, (str, int, float, bool, type(None))) + json.dumps(extra)