diff --git a/src/glassflow/client.py b/src/glassflow/client.py index 5334e72..9cfa15a 100644 --- a/src/glassflow/client.py +++ b/src/glassflow/client.py @@ -16,6 +16,7 @@ from . import __version__ from .config import GlassflowConfig, resolve_config +from .heartbeat import HeartbeatSender, OpenRootSpanTracker from .instrumentation import enable_instrumentations from .masking import MaskingSpanExporter from .semconv import TRACER_NAME @@ -50,9 +51,15 @@ class GlassflowClient: configuration is available as ``client.config``. """ - def __init__(self, provider: TracerProvider, config: GlassflowConfig) -> None: + def __init__( + self, + provider: TracerProvider, + config: GlassflowConfig, + heartbeat: HeartbeatSender | None = None, + ) -> None: self._provider = provider self.config = config + self._heartbeat = heartbeat self._is_shutdown = False def get_tracer(self, name: str = TRACER_NAME) -> trace.Tracer: @@ -68,8 +75,14 @@ def flush(self, timeout_millis: int = 30_000) -> bool: return self._provider.force_flush(timeout_millis) def shutdown(self) -> None: - """Drain pending spans and stop. Releases the global init() slot.""" + """Drain pending spans and stop. Releases the global init() slot. + + Also stops the heartbeat thread and sends its final ``stopped`` ping, + so the backend can tell a clean shutdown from a vanished agent. + """ global _current_client + if self._heartbeat is not None: + self._heartbeat.stop() self._provider.shutdown() self._is_shutdown = True with _lock: @@ -89,6 +102,10 @@ def init( mask: Callable[[Any], Any] | None = None, instruments: Sequence[str] | None = None, span_exporter: SpanExporter | None = None, + heartbeat: bool | None = None, + heartbeat_interval: float | None = None, + agent_name: str | None = None, + heartbeat_transport: Callable[[dict[str, Any]], None] | None = None, set_global: bool = True, ) -> GlassflowClient: """Initialize the SDK: build a tracer provider that exports OTLP traces. @@ -112,6 +129,16 @@ def init( Instrumentors are process-global, so with ``set_global=False`` they are only enabled when ``instruments`` is passed explicitly. span_exporter: Override the default OTLP exporter (useful for testing). + heartbeat: Enable the agent-lifetime heartbeat thread + (``GLASSFLOW_HEARTBEAT``; default off this release). Pings + ``/v1/heartbeat`` from init until process exit so the + platform can tell a live-but-idle agent from a vanished one. + heartbeat_interval: Seconds between pings (default 15, clamped to + ``[5, 300]`` — the backend derives staleness from this). + agent_name: Identity heartbeats group under; defaults to + ``service_name``. + heartbeat_transport: Override the heartbeat HTTP transport + (useful for testing, like ``span_exporter``). set_global: Register the provider as the global OpenTelemetry provider. """ global _current_client @@ -134,6 +161,10 @@ def init( mask=mask, instruments=instruments, span_exporter=span_exporter, + heartbeat=heartbeat, + heartbeat_interval=heartbeat_interval, + agent_name=agent_name, + heartbeat_transport=heartbeat_transport, set_global=set_global, ) @@ -150,6 +181,10 @@ def _do_init( mask: Callable[[Any], Any] | None, instruments: Sequence[str] | None, span_exporter: SpanExporter | None, + heartbeat: bool | None, + heartbeat_interval: float | None, + agent_name: str | None, + heartbeat_transport: Callable[[dict[str, Any]], None] | None, set_global: bool, ) -> GlassflowClient: global _current_client @@ -161,6 +196,9 @@ def _do_init( disabled=disabled, sample_rate=sample_rate, capture_content=capture_content, + heartbeat=heartbeat, + heartbeat_interval=heartbeat_interval, + agent_name=agent_name, ) # telemetry.sdk.* is reserved for the OTel SDK itself (Resource.create fills # it); we identify as a distribution via telemetry.distro.*. @@ -197,7 +235,24 @@ def _do_init( if not config.disabled and (set_global or instruments is not None): enable_instrumentations(provider, instruments) - client = GlassflowClient(provider, config) + # Heartbeat: process-lifetime liveness, independent of trace + # traffic. The tracker rides the provider as a span processor so payloads + # can carry the currently-open root trace ids; disabled kills it too. + sender: HeartbeatSender | None = None + if config.heartbeat and not config.disabled: + tracker = OpenRootSpanTracker() + provider.add_span_processor(tracker) + sender = HeartbeatSender( + url=config.heartbeat_endpoint, + headers=config.headers, + interval=config.heartbeat_interval, + agent_name=config.agent_name, + tracker=tracker, + transport=heartbeat_transport, + ) + sender.start() + + client = GlassflowClient(provider, config, heartbeat=sender) if set_global: _current_client = client return client diff --git a/src/glassflow/config.py b/src/glassflow/config.py index b5d5a23..b62fc28 100644 --- a/src/glassflow/config.py +++ b/src/glassflow/config.py @@ -21,6 +21,15 @@ ENV_DISABLED = "GLASSFLOW_DISABLED" ENV_SAMPLE_RATE = "GLASSFLOW_SAMPLE_RATE" ENV_CAPTURE_CONTENT = "GLASSFLOW_CAPTURE_CONTENT" +ENV_HEARTBEAT = "GLASSFLOW_HEARTBEAT" +ENV_HEARTBEAT_INTERVAL = "GLASSFLOW_HEARTBEAT_INTERVAL" +ENV_AGENT_NAME = "GLASSFLOW_AGENT_NAME" + +# The backend expresses staleness as multiples of the interval, so the clamp +# bounds are part of the heartbeat wire contract. +HEARTBEAT_INTERVAL_MIN = 5.0 +HEARTBEAT_INTERVAL_MAX = 300.0 +DEFAULT_HEARTBEAT_INTERVAL = 15.0 _TRUENESS = frozenset({"1", "true", "yes", "on"}) @@ -57,12 +66,20 @@ class GlassflowConfig: disabled: bool = False sample_rate: float = 1.0 capture_content: bool = True + heartbeat: bool = False + heartbeat_interval: float = DEFAULT_HEARTBEAT_INTERVAL + agent_name: str = DEFAULT_SERVICE_NAME @property def traces_endpoint(self) -> str: """Full OTLP/HTTP traces URL (``/v1/traces``).""" return self.endpoint.rstrip("/") + "/v1/traces" + @property + def heartbeat_endpoint(self) -> str: + """Heartbeat URL (``/v1/heartbeat``) — same host as traces.""" + return self.endpoint.rstrip("/") + "/v1/heartbeat" + def _clamp_sample_rate(value: float) -> float: """Clamp to [0.0, 1.0] — an out-of-range value must degrade, not crash init().""" @@ -73,6 +90,21 @@ def _clamp_sample_rate(value: float) -> float: return clamped +def _clamp_heartbeat_interval(value: float) -> float: + """Clamp to the contract bounds — out-of-range degrades, never crashes init().""" + if HEARTBEAT_INTERVAL_MIN <= value <= HEARTBEAT_INTERVAL_MAX: + return value + clamped = min(max(value, HEARTBEAT_INTERVAL_MIN), HEARTBEAT_INTERVAL_MAX) + logger.warning( + "heartbeat_interval %s is outside [%s, %s]; clamped to %s", + value, + HEARTBEAT_INTERVAL_MIN, + HEARTBEAT_INTERVAL_MAX, + clamped, + ) + return clamped + + def resolve_config( *, endpoint: str | None = None, @@ -82,6 +114,9 @@ def resolve_config( disabled: bool | None = None, sample_rate: float | None = None, capture_content: bool | None = None, + heartbeat: bool | None = None, + heartbeat_interval: float | None = None, + agent_name: str | None = None, ) -> GlassflowConfig: """Resolve SDK configuration from arguments, environment, then defaults. @@ -104,6 +139,15 @@ def resolve_config( (``GLASSFLOW_SAMPLE_RATE``). capture_content: When ``False``, content attributes are stripped at export (``GLASSFLOW_CAPTURE_CONTENT``). + heartbeat: Enable the agent-lifetime heartbeat thread + (``GLASSFLOW_HEARTBEAT``). Off by default this release. + heartbeat_interval: Seconds between pings + (``GLASSFLOW_HEARTBEAT_INTERVAL``), clamped to ``[5, 300]`` — + the backend derives staleness from this, so the bounds are part + of the wire contract. + agent_name: Identity heartbeats group under (``GLASSFLOW_AGENT_NAME``); + defaults to ``service_name`` so the agents view and the traces + view agree on what an "agent" is. Returns: The resolved, immutable ``GlassflowConfig``. @@ -119,6 +163,14 @@ def resolve_config( _env_bool(ENV_CAPTURE_CONTENT, default=True) if capture_content is None else capture_content ) + resolved_heartbeat = _env_bool(ENV_HEARTBEAT, default=False) if heartbeat is None else heartbeat + resolved_heartbeat_interval = _clamp_heartbeat_interval( + _env_float(ENV_HEARTBEAT_INTERVAL, default=DEFAULT_HEARTBEAT_INTERVAL) + if heartbeat_interval is None + else heartbeat_interval + ) + resolved_agent_name = agent_name or os.getenv(ENV_AGENT_NAME) or resolved_service_name + resolved_headers = dict(headers or {}) has_auth = any(key.lower() == "authorization" for key in resolved_headers) if resolved_api_key and not has_auth: @@ -132,4 +184,7 @@ def resolve_config( disabled=resolved_disabled, sample_rate=resolved_sample_rate, capture_content=resolved_capture_content, + heartbeat=resolved_heartbeat, + heartbeat_interval=resolved_heartbeat_interval, + agent_name=resolved_agent_name, ) diff --git a/src/glassflow/heartbeat.py b/src/glassflow/heartbeat.py new file mode 100644 index 0000000..9f3ca5a --- /dev/null +++ b/src/glassflow/heartbeat.py @@ -0,0 +1,258 @@ +"""Agent-lifetime heartbeat sender (payload v1). + +The heartbeat answers one question traces cannot: is this agent process +alive right now? Spans export only when they finish, so an idle or crashed +agent is indistinguishable from a healthy quiet one. Heartbeats are a +process-lifetime signal, fully independent of trace traffic: a daemon +thread pings ``POST /v1/heartbeat`` from ``init()`` until process exit. + +Contract highlights (the spec is normative; this module implements it): + +- First ping immediately at start (the agent appears without waiting an + interval), then every ``interval`` seconds. +- Graceful shutdown (``client.shutdown()`` / ``atexit``) sends a final + ``stopped: true`` ping. No signal handlers are installed — a library + must not own process signals; an unhandled SIGTERM/SIGKILL means no + stopped ping, and the backend's stale→gone path covers exactly that. +- Never raises into user code. Pings have a short timeout, are never + retried or queued (liveness is only true fresh — a late heartbeat is + misinformation), and delivery problems warn once per process. +- ``fork()``: the child re-arms with a NEW ``instance_id`` — one identity + never speaks for two processes. +""" + +from __future__ import annotations + +import atexit +import json +import logging +import os +import threading +import urllib.request +import uuid +import weakref +from collections.abc import Callable +from datetime import datetime, timezone +from typing import Any + +from opentelemetry.context import Context +from opentelemetry.sdk.trace import ReadableSpan, Span, SpanProcessor + +from . import __version__ + +logger = logging.getLogger(__name__) + +PAYLOAD_VERSION = 1 +OPEN_TRACES_CAP = 32 +_PING_TIMEOUT_S = 3.0 +# The final stopped ping runs inside atexit: it must never hold the user's +# process exit hostage, so it gets a tighter budget than regular pings (a +# missed stopped ping just reports as gone instead of stopped — acceptable). +_FINAL_PING_TIMEOUT_S = 1.0 + +# os.register_at_fork callbacks can NEVER be unregistered, so per-sender +# registration would leak a callback + sender reference for every init() +# in a long-lived app. One module-level handler + a weak registry instead: +# stopped/collected senders simply vanish from the set. +_active_senders: weakref.WeakSet[HeartbeatSender] = weakref.WeakSet() +_fork_hook_installed = False +_fork_hook_lock = threading.Lock() + + +def _reset_active_senders_in_child() -> None: # pragma: no cover — fork-only + for sender in list(_active_senders): + sender._reset_in_child() # noqa: SLF001 — module-internal + + +def _install_fork_hook() -> None: + global _fork_hook_installed + with _fork_hook_lock: + if _fork_hook_installed or not hasattr(os, "register_at_fork"): + return + os.register_at_fork(after_in_child=_reset_active_senders_in_child) + _fork_hook_installed = True + + +class OpenRootSpanTracker(SpanProcessor): + """Tracks trace ids of currently-open root spans. + + A root span is one started with no parent context; children of the same + trace never touch the set. This is what lets the backend derive + ``running`` vs ``ready`` from the heartbeat payload alone. + """ + + def __init__(self) -> None: + self._lock = threading.Lock() + # trace_id (32-hex) -> count of open root spans in that trace (a + # trace id normally has one root, but be safe about duplicates). + self._open: dict[str, int] = {} + + def on_start(self, span: Span, parent_context: Context | None = None) -> None: + if span.parent is not None: + return + context = span.get_span_context() + if context is None: + return + self._on_root_start(format(context.trace_id, "032x")) + + def on_end(self, span: ReadableSpan) -> None: + if span.parent is not None: + return + context = span.get_span_context() + if context is None: + return + trace_id = format(context.trace_id, "032x") + with self._lock: + count = self._open.get(trace_id, 0) - 1 + if count <= 0: + self._open.pop(trace_id, None) + else: + self._open[trace_id] = count + + def _on_root_start(self, trace_id: str) -> None: + with self._lock: + self._open[trace_id] = self._open.get(trace_id, 0) + 1 + + def open_trace_ids(self) -> list[str]: + """Trace ids of currently-open root spans (insertion order).""" + with self._lock: + return list(self._open) + + def shutdown(self) -> None: # pragma: no cover — SpanProcessor API + pass + + def force_flush(self, timeout_millis: int = 30_000) -> bool: # pragma: no cover + return True + + +def _http_transport(url: str, headers: dict[str, str]) -> Callable[[dict[str, Any], float], None]: + """Default transport: a plain POST with a per-call timeout, no retries. + + TLS certificate verification is urllib's default and is deliberately not + configurable here — a liveness signal must not become a reason to accept + unverified endpoints. + """ + + def send(payload: dict[str, Any], timeout: float) -> None: + request = urllib.request.Request( + url, + data=json.dumps(payload).encode("utf-8"), + headers={"Content-Type": "application/json", **headers}, + method="POST", + ) + with urllib.request.urlopen(request, timeout=timeout): + pass # 2xx is success; the body is ignored by contract + + return send + + +class HeartbeatSender: + """Daemon thread pinging the heartbeat endpoint for the process lifetime.""" + + def __init__( + self, + *, + url: str, + headers: dict[str, str], + interval: float, + agent_name: str, + tracker: OpenRootSpanTracker, + transport: Callable[[dict[str, Any]], None] | None = None, + ping_timeout: float = _PING_TIMEOUT_S, + final_ping_timeout: float = _FINAL_PING_TIMEOUT_S, + ) -> None: + self._interval = interval + self._agent_name = agent_name + self._tracker = tracker + self._ping_timeout = ping_timeout + self._final_ping_timeout = final_ping_timeout + if transport is not None: + # Injected transports (tests) take the payload only. + self._send: Callable[[dict[str, Any], float], None] = lambda p, _t: transport(p) + else: + self._send = _http_transport(url, headers) + self._instance_id = str(uuid.uuid4()) + self._stop_event = threading.Event() + self._stopped = False + self._lock = threading.Lock() + self._delivery_warned = False + self._thread: threading.Thread | None = None + + @property + def instance_id(self) -> str: + """Identity of one process lifetime; fresh per process (and per fork).""" + return self._instance_id + + def start(self) -> None: + """Start the daemon thread; first ping goes out immediately.""" + self._thread = threading.Thread(target=self._run, name="glassflow-heartbeat", daemon=True) + self._thread.start() + atexit.register(self.stop) + # A forked child must never reuse the parent's identity; the shared + # module-level fork hook re-arms every live sender with a fresh + # instance_id (same pattern the OTel exporter uses). + _active_senders.add(self) + _install_fork_hook() + + def stop(self) -> None: + """Stop the thread and send the final ``stopped`` ping. Idempotent. + + Exit-latency budget: an in-flight ping can hold the join for at most + ``ping_timeout``; the final ping gets ``final_ping_timeout``. A dead + endpoint therefore delays process exit by a bounded few seconds, not + by the interval. + """ + with self._lock: + if self._stopped: + return + self._stopped = True + _active_senders.discard(self) + atexit.unregister(self.stop) + self._stop_event.set() + if self._thread is not None: + self._thread.join(timeout=self._ping_timeout + 1.0) + self._send_ping(stopped=True) + + def _run(self) -> None: + self._send_ping() + while not self._stop_event.wait(self._interval): + self._send_ping() + + def _reset_in_child(self) -> None: # pragma: no cover — fork-only path + if self._stopped: + return + self._instance_id = str(uuid.uuid4()) + self._stop_event = threading.Event() + self.start() + + def __del__(self) -> None: # pragma: no cover — GC-timing dependent + atexit.unregister(self.stop) + + def _build_payload(self, *, stopped: bool = False) -> dict[str, Any]: + open_ids = self._tracker.open_trace_ids() + payload: dict[str, Any] = { + "v": PAYLOAD_VERSION, + "instance_id": self._instance_id, + "agent_name": self._agent_name, + # RFC3339 UTC with sub-second precision, Z suffix + "sent_at": datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "Z", + "sdk_language": "python", + "sdk_version": __version__, + "open_traces": open_ids[:OPEN_TRACES_CAP], + "open_trace_count": len(open_ids), + } + if stopped: + # Present-and-true only on the final ping; false is never sent. + payload["stopped"] = True + return payload + + def _send_ping(self, *, stopped: bool = False) -> None: + timeout = self._final_ping_timeout if stopped else self._ping_timeout + try: + self._send(self._build_payload(stopped=stopped), timeout) + except Exception as exc: # noqa: BLE001 — never raises into user code + if not self._delivery_warned: + self._delivery_warned = True + logger.warning("heartbeat delivery failed (%s); further failures log at DEBUG", exc) + else: + logger.debug("heartbeat delivery failed: %s", exc) diff --git a/tests/test_generation.py b/tests/test_generation.py index 449bff0..962ec40 100644 --- a/tests/test_generation.py +++ b/tests/test_generation.py @@ -244,7 +244,7 @@ def test_manual_generation(exported_spans: InMemorySpanExporter) -> None: assert attrs["gen_ai.usage.input_tokens"] == 1 -# --- record_first_token: the TTFT anchor for streaming (GLA2-175) --- +# --- record_first_token: the TTFT anchor for streaming --- def _first_token_events(span: ReadableSpan) -> list[Event]: diff --git a/tests/test_heartbeat.py b/tests/test_heartbeat.py new file mode 100644 index 0000000..c110c24 --- /dev/null +++ b/tests/test_heartbeat.py @@ -0,0 +1,448 @@ +"""Agent-lifetime heartbeat sender. + +Implements the heartbeat spec: payload v1, emission semantics (init-to-exit +lifetime, immediate first ping, stopped ping on graceful shutdown, silent +failure), and the config surface (off by default, interval clamped [5, 300], +agent_name defaults to service_name). +""" + +from __future__ import annotations + +import json +import logging +import re +import threading +import uuid +from http.server import BaseHTTPRequestHandler, HTTPServer +from typing import Any + +import pytest +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from glassflow import init +from glassflow.config import resolve_config +from glassflow.heartbeat import HeartbeatSender, OpenRootSpanTracker + +# --------------------------------------------------------------------------- +# Config surface +# --------------------------------------------------------------------------- + + +def test_heartbeat_disabled_by_default() -> None: + assert resolve_config().heartbeat is False + + +def test_heartbeat_enabled_via_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("GLASSFLOW_HEARTBEAT", "true") + assert resolve_config().heartbeat is True + + +def test_heartbeat_argument_wins_over_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("GLASSFLOW_HEARTBEAT", "true") + assert resolve_config(heartbeat=False).heartbeat is False + + +def test_heartbeat_interval_default() -> None: + assert resolve_config().heartbeat_interval == 15.0 + + +def test_heartbeat_interval_env(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("GLASSFLOW_HEARTBEAT_INTERVAL", "30") + assert resolve_config().heartbeat_interval == 30.0 + + +def test_heartbeat_interval_clamped_low_with_warning(caplog: pytest.LogCaptureFixture) -> None: + with caplog.at_level(logging.WARNING): + assert resolve_config(heartbeat_interval=1.0).heartbeat_interval == 5.0 + assert any("heartbeat_interval" in r.message for r in caplog.records) + + +def test_heartbeat_interval_clamped_high() -> None: + assert resolve_config(heartbeat_interval=1000.0).heartbeat_interval == 300.0 + + +def test_agent_name_defaults_to_service_name() -> None: + config = resolve_config(service_name="checkout-agent") + assert config.agent_name == "checkout-agent" + + +def test_agent_name_override(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("GLASSFLOW_AGENT_NAME", "env-agent") + assert resolve_config(service_name="svc").agent_name == "env-agent" + assert resolve_config(service_name="svc", agent_name="arg-agent").agent_name == "arg-agent" + + +def test_heartbeat_endpoint_property() -> None: + config = resolve_config(endpoint="https://ingest.example.com/") + assert config.heartbeat_endpoint == "https://ingest.example.com/v1/heartbeat" + + +# --------------------------------------------------------------------------- +# Open-root-span tracker +# --------------------------------------------------------------------------- + + +def _tracked_client() -> tuple[Any, OpenRootSpanTracker]: + tracker = OpenRootSpanTracker() + client = init( + set_global=False, + service_name="t", + span_exporter=InMemorySpanExporter(), + ) + client._provider.add_span_processor(tracker) # noqa: SLF001 — test wiring + return client, tracker + + +def test_tracker_counts_open_root_spans() -> None: + client, tracker = _tracked_client() + assert tracker.open_trace_ids() == [] + with client.get_tracer().start_as_current_span("root") as root: + expected = format(root.get_span_context().trace_id, "032x") + assert tracker.open_trace_ids() == [expected] + # a child span of the same trace is NOT a new open root + with client.get_tracer().start_as_current_span("child"): + assert tracker.open_trace_ids() == [expected] + assert tracker.open_trace_ids() == [] + + +def test_tracker_handles_concurrent_roots() -> None: + client, tracker = _tracked_client() + tracer = client.get_tracer() + a = tracer.start_span("a") + b = tracer.start_span("b") + assert len(tracker.open_trace_ids()) == 2 + a.end() + assert len(tracker.open_trace_ids()) == 1 + b.end() + assert tracker.open_trace_ids() == [] + + +# --------------------------------------------------------------------------- +# Payload (spec v1) +# --------------------------------------------------------------------------- + + +def _sender( + sent: list[dict[str, Any]], + tracker: OpenRootSpanTracker | None = None, + interval: float = 3600.0, +) -> HeartbeatSender: + return HeartbeatSender( + url="https://ingest.example.com/v1/heartbeat", + headers={}, + interval=interval, + agent_name="checkout-agent", + tracker=tracker or OpenRootSpanTracker(), + transport=sent.append, + ) + + +def test_payload_matches_spec_v1() -> None: + sent: list[dict[str, Any]] = [] + sender = _sender(sent) + sender._send_ping() # noqa: SLF001 — payload unit test + payload = sent[0] + assert payload["v"] == 1 + uuid.UUID(payload["instance_id"]) # valid UUID + assert payload["agent_name"] == "checkout-agent" + # RFC3339 UTC with sub-second precision + assert re.fullmatch(r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}\.\d+Z", payload["sent_at"]) + assert payload["sdk_language"] == "python" + assert payload["sdk_version"] + assert payload["open_traces"] == [] + assert payload["open_trace_count"] == 0 + assert "stopped" not in payload + + +def test_payload_open_traces_capped_at_32() -> None: + sent: list[dict[str, Any]] = [] + tracker = OpenRootSpanTracker() + for i in range(40): + tracker._on_root_start(format(i + 1, "032x")) # noqa: SLF001 + sender = _sender(sent, tracker=tracker) + sender._send_ping() # noqa: SLF001 + payload = sent[0] + assert len(payload["open_traces"]) == 32 + assert payload["open_trace_count"] == 40 + + +def test_instance_id_constant_across_pings() -> None: + sent: list[dict[str, Any]] = [] + sender = _sender(sent) + sender._send_ping() # noqa: SLF001 + sender._send_ping() # noqa: SLF001 + assert sent[0]["instance_id"] == sent[1]["instance_id"] + + +# --------------------------------------------------------------------------- +# Lifecycle: immediate first ping, stopped ping, idempotent stop +# --------------------------------------------------------------------------- + + +def test_start_sends_immediate_ping() -> None: + sent: list[dict[str, Any]] = [] + first_ping = threading.Event() + + def transport(payload: dict[str, Any]) -> None: + sent.append(payload) + first_ping.set() + + sender = HeartbeatSender( + url="u", + headers={}, + interval=3600.0, + agent_name="a", + tracker=OpenRootSpanTracker(), + transport=transport, + ) + sender.start() + try: + assert first_ping.wait(timeout=5.0), "no ping arrived after start()" + assert "stopped" not in sent[0] + finally: + sender.stop() + + +def test_stop_sends_stopped_ping_exactly_once() -> None: + sent: list[dict[str, Any]] = [] + sender = _sender(sent) + sender.start() + sender.stop() + sender.stop() # idempotent + stopped_pings = [p for p in sent if p.get("stopped") is True] + assert len(stopped_pings) == 1 + assert sent[-1] is stopped_pings[0] + + +def test_client_shutdown_stops_heartbeat() -> None: + sent: list[dict[str, Any]] = [] + client = init( + set_global=False, + service_name="hb-svc", + heartbeat=True, + heartbeat_transport=sent.append, + span_exporter=InMemorySpanExporter(), + ) + client.shutdown() + assert any(p.get("stopped") is True for p in sent) + assert sent[0]["agent_name"] == "hb-svc" # agent_name defaulted to service_name + + +def test_open_traces_flow_into_payloads() -> None: + sent: list[dict[str, Any]] = [] + client = init( + set_global=False, + service_name="hb-svc", + heartbeat=True, + heartbeat_transport=sent.append, + span_exporter=InMemorySpanExporter(), + ) + try: + span = client.get_tracer().start_span("root") + trace_id = format(span.get_span_context().trace_id, "032x") + client._heartbeat._send_ping() # noqa: SLF001 — deterministic mid-run ping + span.end() + client._heartbeat._send_ping() # noqa: SLF001 + with_open = [p for p in sent if trace_id in p.get("open_traces", [])] + assert with_open, "open root trace never appeared in a payload" + assert sent[-1]["open_traces"] == [] + finally: + client.shutdown() + + +def test_heartbeat_off_by_default_no_thread() -> None: + client = init(set_global=False, service_name="svc", span_exporter=InMemorySpanExporter()) + assert client._heartbeat is None # noqa: SLF001 + client.shutdown() + + +def test_disabled_kill_switch_disables_heartbeat() -> None: + sent: list[dict[str, Any]] = [] + client = init( + set_global=False, + service_name="svc", + disabled=True, + heartbeat=True, + heartbeat_transport=sent.append, + ) + client.shutdown() + assert sent == [] + + +# --------------------------------------------------------------------------- +# Failure behavior: never raises, warns once +# --------------------------------------------------------------------------- + + +def test_transport_failure_never_raises_and_warns_once( + caplog: pytest.LogCaptureFixture, +) -> None: + def broken(_: dict[str, Any]) -> None: + raise ConnectionError("endpoint down") + + sender = HeartbeatSender( + url="u", + headers={}, + interval=3600.0, + agent_name="a", + tracker=OpenRootSpanTracker(), + transport=broken, + ) + with caplog.at_level(logging.DEBUG, logger="glassflow.heartbeat"): + sender._send_ping() # noqa: SLF001 + sender._send_ping() # noqa: SLF001 + warnings = [r for r in caplog.records if r.levelno == logging.WARNING] + assert len(warnings) == 1 + + +# --------------------------------------------------------------------------- +# Real HTTP transport (no mocks): local server, auth header, 204 +# --------------------------------------------------------------------------- + + +def test_http_transport_end_to_end() -> None: + received: list[tuple[dict[str, Any], str | None]] = [] + + class Handler(BaseHTTPRequestHandler): + def do_POST(self) -> None: # noqa: N802 — http.server API + body = self.rfile.read(int(self.headers["Content-Length"])) + received.append((json.loads(body), self.headers.get("Authorization"))) + self.send_response(204) + self.end_headers() + + def log_message(self, *args: Any) -> None: # silence test output + pass + + server = HTTPServer(("127.0.0.1", 0), Handler) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + url = f"http://127.0.0.1:{server.server_port}/v1/heartbeat" + sender = HeartbeatSender( + url=url, + headers={"Authorization": "Bearer test-key"}, + interval=3600.0, + agent_name="e2e-agent", + tracker=OpenRootSpanTracker(), + ) + sender._send_ping() # noqa: SLF001 + assert len(received) == 1 + payload, auth = received[0] + assert payload["v"] == 1 + assert payload["agent_name"] == "e2e-agent" + assert auth == "Bearer test-key" + finally: + server.shutdown() + thread.join(timeout=5) + + +# --------------------------------------------------------------------------- +# atexit: a real subprocess that exits cleanly sends the stopped ping +# --------------------------------------------------------------------------- + + +def test_atexit_sends_stopped_ping_from_real_process() -> None: + import os + import subprocess + import sys + + heartbeats: list[dict[str, Any]] = [] + + class Handler(BaseHTTPRequestHandler): + def do_POST(self) -> None: # noqa: N802 — http.server API + body = self.rfile.read(int(self.headers.get("Content-Length", "0"))) + if self.path == "/v1/heartbeat": + heartbeats.append(json.loads(body)) + self.send_response(204) + self.end_headers() + + def log_message(self, *args: Any) -> None: + pass + + server = HTTPServer(("127.0.0.1", 0), Handler) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + env = { + **os.environ, + "GLASSFLOW_ENDPOINT": f"http://127.0.0.1:{server.server_port}", + "GLASSFLOW_HEARTBEAT": "1", + "GLASSFLOW_SERVICE_NAME": "atexit-agent", + } + # init() then exit normally WITHOUT calling shutdown(): the atexit + # hook alone must produce the stopped ping. + result = subprocess.run( + [sys.executable, "-c", "import glassflow; glassflow.init(instruments=[])"], + env=env, + timeout=30, + capture_output=True, + ) + assert result.returncode == 0, result.stderr.decode() + assert heartbeats, "no heartbeat arrived from the subprocess" + assert "stopped" not in heartbeats[0] + assert heartbeats[-1].get("stopped") is True + assert all(p["agent_name"] == "atexit-agent" for p in heartbeats) + finally: + server.shutdown() + thread.join(timeout=5) + + +# --------------------------------------------------------------------------- +# Shutdown latency & resource hygiene (security/robustness review) +# --------------------------------------------------------------------------- + + +def test_stop_with_dead_slow_endpoint_returns_quickly() -> None: + """A dead endpoint must not hold the user's process exit hostage.""" + import socket + import time + + # A socket that accepts but never responds — the worst-case endpoint. + listener = socket.socket() + listener.bind(("127.0.0.1", 0)) + listener.listen(1) + try: + url = f"http://127.0.0.1:{listener.getsockname()[1]}/v1/heartbeat" + sender = HeartbeatSender( + url=url, + headers={}, + interval=3600.0, + agent_name="a", + tracker=OpenRootSpanTracker(), + ping_timeout=0.3, + final_ping_timeout=0.3, + ) + sender.start() + time.sleep(0.05) # let the first (hanging) ping get in flight + started = time.monotonic() + sender.stop() + elapsed = time.monotonic() - started + # bound: in-flight ping timeout + final ping timeout + slack + assert elapsed < 2.0, f"stop() took {elapsed:.2f}s against a dead endpoint" + finally: + listener.close() + + +def test_final_ping_timeout_defaults_shorter_than_ping_timeout() -> None: + sender = HeartbeatSender( + url="u", + headers={}, + interval=3600.0, + agent_name="a", + tracker=OpenRootSpanTracker(), + ) + assert sender._final_ping_timeout < sender._ping_timeout # noqa: SLF001 + assert sender._final_ping_timeout == 1.0 # noqa: SLF001 + + +def test_stopped_senders_leave_the_fork_registry() -> None: + """init/shutdown cycles must not accumulate fork-handler references.""" + from glassflow.heartbeat import _active_senders + + before = len(_active_senders) + sent: list[dict[str, Any]] = [] + sender = _sender(sent) + sender.start() + assert len(_active_senders) == before + 1 + sender.stop() + assert len(_active_senders) == before