Skip to content
Closed
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
2 changes: 2 additions & 0 deletions temporalio/contrib/__init__.py
Original file line number Diff line number Diff line change
@@ -1 +1,3 @@
"""Extra modules that may have optional dependencies."""

from .sanitizer import SanitizingPayloadCodec # noqa: F401
79 changes: 79 additions & 0 deletions temporalio/contrib/sanitizer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,79 @@
from __future__ import annotations

import re
from collections.abc import Mapping, Sequence
from typing import Any

import temporalio.api.common.v1 as api_common
from temporalio.converter import PayloadCodec


_DEFAULT_PATTERNS: tuple[re.Pattern[str], ...] = (
re.compile(r"(?i)(api[-_]?key|access[-_]?token|auth[-_]?token|secret|jwt|bearer)"),
re.compile(r"(?i)(ssn|social[-_]?security|credit[-_]?card|cc[-_]?num)"),
)


def _mask_value(val: Any) -> Any:
if val is None:
return None
if isinstance(val, (int, float, bool)):
return val
return "[REDACTED]"


def _sanitize_obj(obj: Any, patterns: Sequence[re.Pattern[str]]) -> Any:
if isinstance(obj, Mapping):
return {
k: _sanitize_obj(_mask_value(v) if any(p.search(str(k)) for p in patterns) else v, patterns)
for k, v in obj.items()
}
if isinstance(obj, list):
return [_sanitize_obj(v, patterns) for v in obj]
if isinstance(obj, tuple):
return tuple(_sanitize_obj(v, patterns) for v in obj)
return obj


class SanitizingPayloadCodec(PayloadCodec):
"""PayloadCodec that redacts sensitive fields by key pattern.

This codec preserves payload encoding/type metadata and only rewrites the
payload data if the decoded JSON is a structured object. When no keys match,
it is effectively a no-op to minimize overhead.
"""

def __init__(self, *, key_patterns: Sequence[str] | None = None) -> None:
pats = key_patterns or []
self._patterns: tuple[re.Pattern[str], ...] = _DEFAULT_PATTERNS + tuple(
re.compile(p, re.IGNORECASE) for p in pats
)

async def encode(self, payloads: Sequence[api_common.Payload]) -> list[api_common.Payload]:
out: list[api_common.Payload] = []
for p in payloads:
# Only operate on json/plain payloads which are common for dict-like values
if p.metadata.get("encoding") == b"json/plain":
try:
import json

val = json.loads(p.data)
sanitized = _sanitize_obj(val, self._patterns)
if sanitized is val:
out.append(p)
continue
new = api_common.Payload()
new.metadata.update(p.metadata)
new.data = json.dumps(sanitized, separators=(",", ":")).encode("utf-8")
out.append(new)
continue
except Exception:
# On any failure, pass-through to avoid data loss
out.append(p)
continue
out.append(p)
return out

async def decode(self, payloads: Sequence[api_common.Payload]) -> list[api_common.Payload]:
# Codec is one-way (sanitizes on encode); decoding is pass-through
return list(payloads)
35 changes: 35 additions & 0 deletions tests/contrib/test_sanitizer.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,35 @@
from __future__ import annotations

import json

import pytest

from temporalio.converter import DataConverter, JSONPlainPayloadConverter
from temporalio.contrib.sanitizer import SanitizingPayloadCodec


@pytest.mark.asyncio
async def test_redacts_dict_keys():
dc = DataConverter(payload_codec=SanitizingPayloadCodec())
payloads = await dc.encode([{"api_key": "abc", "name": "bob"}])
assert payloads[0].metadata.get("encoding") == b"json/plain"
val = json.loads(payloads[0].data)
assert val["api_key"] == "[REDACTED]"
assert val["name"] == "bob"


@pytest.mark.asyncio
async def test_nested_and_list_traversal():
dc = DataConverter(payload_codec=SanitizingPayloadCodec(key_patterns=["password"]))
obj = {"user": {"password": "secret", "emails": ["a@x", "b@y"]}}
payloads = await dc.encode([obj])
val = json.loads(payloads[0].data)
assert val["user"]["password"] == "[REDACTED]"
assert val["user"]["emails"] == ["a@x", "b@y"]


@pytest.mark.asyncio
async def test_noop_when_no_sensitive_keys():
dc = DataConverter(payload_codec=SanitizingPayloadCodec())
payloads = await dc.encode([[1, 2, 3]])
assert json.loads(payloads[0].data) == [1, 2, 3]