diff --git a/src/lambda_obs/logging.py b/src/lambda_obs/logging.py index eb4a61c..aa79b85 100644 --- a/src/lambda_obs/logging.py +++ b/src/lambda_obs/logging.py @@ -5,6 +5,7 @@ import json import sys import time +import uuid from typing import Any, Mapping, MutableMapping, TextIO @@ -17,16 +18,24 @@ def __init__( *, stream: TextIO | None = None, base: Mapping[str, Any] | None = None, + correlation_id: str | None = None, ) -> None: self.service = service self.stream = stream or sys.stdout self._base: dict[str, Any] = dict(base or {}) + if correlation_id: + self._base.setdefault("correlation_id", correlation_id) def bind(self, **kwargs: Any) -> "Logger": merged = {**self._base, **kwargs} child = Logger(self.service, stream=self.stream, base=merged) return child + def with_correlation_id(self, correlation_id: str | None = None) -> "Logger": + """Return a child logger bound to a correlation id (generated if omitted).""" + cid = correlation_id or str(uuid.uuid4()) + return self.bind(correlation_id=cid) + def _emit(self, level: str, message: str, **fields: Any) -> None: payload: MutableMapping[str, Any] = { "timestamp": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime()), diff --git a/tests/test_logging.py b/tests/test_logging.py index 30160d7..99238ba 100644 --- a/tests/test_logging.py +++ b/tests/test_logging.py @@ -24,3 +24,20 @@ def test_logger_bind_adds_fields(): payload = json.loads(buf.getvalue().strip()) assert payload["cold_start"] is True assert payload["level"] == "WARNING" + + +def test_logger_with_correlation_id(): + buf = io.StringIO() + log = Logger("demo", stream=buf).with_correlation_id("corr-123") + log.info("traced") + payload = json.loads(buf.getvalue().strip()) + assert payload["correlation_id"] == "corr-123" + + +def test_logger_generates_correlation_id(): + buf = io.StringIO() + log = Logger("demo", stream=buf).with_correlation_id() + log.info("traced") + payload = json.loads(buf.getvalue().strip()) + assert "correlation_id" in payload + assert len(payload["correlation_id"]) > 0