From 79f74d6dc02f384fa1e2889a10fd8bdd5350e16d Mon Sep 17 00:00:00 2001 From: "google-labs-jules[bot]" <161369871+google-labs-jules[bot]@users.noreply.github.com> Date: Sun, 19 Jul 2026 06:37:14 +0000 Subject: [PATCH] test: add test for write_deidentified_stream Adds a mock-based test for write_deidentified_stream in test_spark_streaming.py. Co-authored-by: zrt219 <199104500+zrt219@users.noreply.github.com> --- .../unit/integrations/test_spark_streaming.py | 37 +++++++++++++++++++ 1 file changed, 37 insertions(+) diff --git a/openmed/tests/unit/integrations/test_spark_streaming.py b/openmed/tests/unit/integrations/test_spark_streaming.py index 939aad0..11f749a 100644 --- a/openmed/tests/unit/integrations/test_spark_streaming.py +++ b/openmed/tests/unit/integrations/test_spark_streaming.py @@ -6,11 +6,13 @@ from typing import Any, Sequence import pytest +from unittest.mock import patch, MagicMock from openmed.integrations.spark_streaming import ( DEFAULT_BATCH_ID_COLUMN, SparkDeidentifyColumn, SparkDeidentifySink, + write_deidentified_stream, _coerce_columns, _redact_partition, _SparkPartitionConfig, @@ -227,3 +229,38 @@ def deidentify_many( assert after_replay == before_replay spark.sql(f"DROP TABLE IF EXISTS {table}") + + +@patch("openmed.integrations.spark_streaming.SparkDeidentifySink") +def test_write_deidentified_stream(mock_sink_cls: Any) -> None: + mock_sink = mock_sink_cls.return_value + mock_sink.start.return_value = "query_result" + + result = write_deidentified_stream( + streaming_df="mock_df", + target_table="my_table", + columns=["note"], + checkpoint_location="/path/to/checkpoint", + query_name="my_query", + output_mode="update", + trigger={"processingTime": "1 minute"}, + options={"checkpointLocation": "/path/to/checkpoint"}, + extra_sink_arg="extra_value" + ) + + mock_sink_cls.assert_called_once_with( + columns=["note"], + target_table="my_table", + checkpoint_location="/path/to/checkpoint", + extra_sink_arg="extra_value" + ) + + mock_sink.start.assert_called_once_with( + "mock_df", + query_name="my_query", + output_mode="update", + trigger={"processingTime": "1 minute"}, + options={"checkpointLocation": "/path/to/checkpoint"} + ) + + assert result == "query_result"