diff --git a/buckaroo/dataflow/dataflow.py b/buckaroo/dataflow/dataflow.py index 2ba5d133a..af096966b 100644 --- a/buckaroo/dataflow/dataflow.py +++ b/buckaroo/dataflow/dataflow.py @@ -404,8 +404,10 @@ def _summary_sd(self, change): # Recorded once the stats exist: after a run that raises, the key still # names the frame summary_sd belongs to, which _populate_sd_cache checks. self._summary_sd_cache_key = key - self.summary_sd = result_summary_sd + # errs first: assigning summary_sd runs _populate_sd_cache, which files + # them with the sd it caches (see _summary_errs_cache). self.errs = errs + self.summary_sd = result_summary_sd @observe('summary_sd', 'processed_result') @exception_protect('merged_sd-protector') @@ -465,6 +467,9 @@ def __init__(self, orig_df, debug=False, raise ValueError(f"stats_tier must be one of {STATS_TIERS}, got {stats_tier!r}") # Set before super().__init__ — assigning raw_df runs the whole cascade. self.stats_tier = stats_tier + # The errs of the runs behind summary_stats_cache entries that had any, + # by the same key (see set_stats_tier). + self._summary_errs_cache: TDict[str, ErrDict] = {} self.init_sd: InitSD if init_sd is None: self.init_sd = {} @@ -708,6 +713,8 @@ def _populate_sd_cache(self, _change): current = (id(self.processed_df), id(self.analysis_klasses), self.stats_tier) if keys['filt'] not in new_cache and self._summary_sd_cache_key == current: new_cache[keys['filt']] = dict(self.summary_sd or {}) + if getattr(self, 'errs', None): + self._summary_errs_cache[keys['filt']] = self.errs cache_grew = True # raw + clean: fresh compute, but only on cache miss. @@ -811,23 +818,35 @@ def set_stats_tier(self, tier: str, summary: Optional[Tuple[SDType, ErrDict]] = clean scopes from the cache, computing the ones it lacks. """ filt_key = self._scope_cache_key(split_chain_by_scope(self.operations)['filt'], tier=tier) - if summary is None: - cached = self.summary_stats_cache.get(filt_key) - if cached is None: - self.stats_tier = tier - return - summary = (cached, {}) - else: - # Stored under the state's key at the new tier, so merged_sd reads - # the stats given rather than an entry already there. - self.summary_stats_cache = {**self.summary_stats_cache, filt_key: summary[0]} - sd, errs = summary - # Recorded first, so the stats_tier observer finds this frame's stats - # present and leaves them to the assignments below. - self._summary_sd_cache_key = (id(self.processed_df), id(self.analysis_klasses), tier) - self.stats_tier = tier - self.summary_sd = sd - self.errs = errs + prior_tier, prior_key = self.stats_tier, self._summary_sd_cache_key + try: + if summary is None: + cached = self.summary_stats_cache.get(filt_key) + if cached is None: + self.stats_tier = tier + return + # The errs of the run that made the entry, kept apart because + # the cache holds sds only. + summary = (cached, self._summary_errs_cache.get(filt_key, {})) + else: + # Stored under the state's key at the new tier, so merged_sd reads + # the stats given rather than an entry already there. + self.summary_stats_cache = {**self.summary_stats_cache, filt_key: summary[0]} + if summary[1]: + self._summary_errs_cache[filt_key] = summary[1] + sd, errs = summary + # Recorded first, so the stats_tier observer finds this frame's stats + # present and leaves them to the assignments below. + self._summary_sd_cache_key = (id(self.processed_df), id(self.analysis_klasses), tier) + self.stats_tier = tier + self.errs = errs + self.summary_sd = sd + except Exception: + # The key goes back first, so the observer the tier change fires + # finds the frame's stats present and runs nothing. + self._summary_sd_cache_key = prior_key + self.stats_tier = prior_tier + raise # ### end summary stats block diff --git a/buckaroo/server/handlers.py b/buckaroo/server/handlers.py index 17edadaa5..be4de020a 100644 --- a/buckaroo/server/handlers.py +++ b/buckaroo/server/handlers.py @@ -16,7 +16,8 @@ from buckaroo.server.focus import find_or_create_session_window from buckaroo.dataflow.dataflow import STATS_TIERS from buckaroo.server.session import ( - DEFAULT_STATS_DELIVERY, DEFAULT_STATS_TIER, STATS_DELIVERIES, build_state_message, dataflow_stats_tier) + DEFAULT_STATS_DELIVERY, DEFAULT_STATS_TIER, STATS_DELIVERIES, begin_stats_generation, dataflow_stats_tier) +from buckaroo.server.stats_wire import broadcast_state, refresh_session_snapshot from buckaroo.server import telemetry from buckaroo.pluggable_analysis_framework import perf_log @@ -194,17 +195,10 @@ def _push_state_to_clients(self, session, metadata: dict): if not session.ws_clients: return - for client in list(session.ws_clients): - try: - # Reset per-client live search first (dataset changed), - # then build the msg so the injected ``buckaroo_state`` - # mirrors the post-reset value. - client.search_string = "" - msg = build_state_message(session, metadata=metadata, - search_string=client.search_string) - client.write_message(json.dumps(msg)) - except Exception: - session.ws_clients.discard(client) + # Each client's search is reset first (dataset changed), then its + # message is built so the injected ``buckaroo_state`` mirrors the + # post-reset value. + broadcast_state(session, metadata=metadata, reset_search=True) def _handle_browser_window(self, session_id: str) -> str: """Handle browser window management.""" @@ -391,6 +385,11 @@ async def post(self): **component_config, } + # /load builds no deferred stats, whatever policy a prior /load_expr on + # this session left behind, and the generation moves on with the data. + session.stats_tier, session.stats_delivery = "full", "inline" + begin_stats_generation(session) + # Notify connected clients and open browser self._push_state_to_clients(session, metadata) browser_action = "skipped" if no_browser else self._handle_browser_window(session_id) @@ -670,17 +669,13 @@ async def post(self): **dvc.get("component_config", {}), **component_config} - if session.ws_clients: - for client in list(session.ws_clients): - try: - # Reset per-client live search (#851): a term from - # the prior expression would silently filter the new one. - client.search_string = "" - msg = build_state_message(session, metadata=metadata, - search_string=client.search_string) - client.write_message(json.dumps(msg)) - except Exception: - session.ws_clients.discard(client) + # A new expression is a new stats generation; a deferred session starts + # it pending. + begin_stats_generation(session) + + # Reset per-client live search (#851): a term from the prior + # expression would silently filter the new one. + broadcast_state(session, metadata=metadata, reset_search=True) if no_browser or not self.application.settings.get("open_browser", False): browser_action = "skipped" @@ -824,18 +819,12 @@ async def post(self): session.df_meta = display_state["df_meta"] session.mode = "viewer" telemetry.arm_session(session, tele_sink) + # A viewer session has no dataflow and so no deferred stats. + session.stats_tier, session.stats_delivery = DEFAULT_STATS_TIER, DEFAULT_STATS_DELIVERY + begin_stats_generation(session) - # Push to WebSocket clients - if session.ws_clients: - for client in list(session.ws_clients): - try: - # Reset per-client live search (#851). - client.search_string = "" - msg = build_state_message(session, - search_string=client.search_string) - client.write_message(json.dumps(msg)) - except Exception: - session.ws_clients.discard(client) + # Push to WebSocket clients. Reset per-client live search (#851). + broadcast_state(session, reset_search=True) # Browser window if no_browser or not self.application.settings.get("open_browser", False): @@ -953,31 +942,13 @@ async def post(self, session_id): if bs.get("quick_command_args"): xorq_dataflow.quick_command_args = bs["quick_command_args"] - refreshed = get_buckaroo_display_state(xorq_dataflow) session.xorq_dataflow = xorq_dataflow session.stats_tier = stats_tier session.stats_delivery = stats_delivery - session.df_display_args = refreshed["df_display_args"] - session.df_data_dict = refreshed["df_data_dict"] - session.df_meta = refreshed["df_meta"] - session.buckaroo_options = refreshed["buckaroo_options"] - session.command_config = refreshed["command_config"] + refresh_session_snapshot(session, xorq_dataflow) + begin_stats_generation(session) - if session.component_config and session.df_display_args: - for key in session.df_display_args: - dvc = session.df_display_args[key].get("df_viewer_config") - if dvc is not None: - dvc["component_config"] = { - **dvc.get("component_config", {}), - **session.component_config} - - for client in list(session.ws_clients): - try: - msg = build_state_message(session, - search_string=getattr(client, "search_string", "")) - client.write_message(json.dumps(msg)) - except Exception: - session.ws_clients.discard(client) + broadcast_state(session) klass_count = len(extra_klasses) log.info("reload_expr session=%s project_root=%s klasses=%d", diff --git a/buckaroo/server/session.py b/buckaroo/server/session.py index 260d9e6e8..e2eecc2e8 100644 --- a/buckaroo/server/session.py +++ b/buckaroo/server/session.py @@ -26,6 +26,25 @@ def dataflow_stats_tier(stats_tier: str, stats_delivery: str) -> str: return "schema" if stats_delivery == "deferred" else stats_tier +# What ``df_meta.stats.status`` says about a session's stats for its current +# ``stats_gen``: ``complete`` (the stats of the tier the session is headed for +# are in the snapshot), ``pending`` (a deferred session that has not produced +# them yet), ``not_computed`` (the session is headed for the schema tier, so +# none will arrive) and ``error`` (the run failed; sticky until the next +# generation). +STATS_STATUSES = ("complete", "pending", "not_computed", "error") + + +def initial_stats_status(stats_tier: str, stats_delivery: str) -> tuple[str, Optional[str]]: + """The status, and its reason, of a session whose stats generation has just + started, from the policy pair alone.""" + if stats_tier != "full": + return "not_computed", "host" + if stats_delivery == "deferred": + return "pending", None + return "complete", None + + @dataclass class SessionState: session_id: str @@ -64,6 +83,16 @@ class SessionState: # tier the dataflow is built at follows from the pair (dataflow_stats_tier). stats_tier: str = DEFAULT_STATS_TIER stats_delivery: str = DEFAULT_STATS_DELIVERY + # The stats generation: a counter the server owns, bumped whenever the + # dataflow state the stats describe changes (/load, /load_expr, /load_compare, + # /reload_expr, and a buckaroo_state_change that touches a dataflow field). + # It rides on df_meta.stats and on every stats message, so a client can drop + # a reply for a state it has left. Independent of any client-owned request + # sequence (state_seq). ``stats_status`` and ``stats_reason`` live here, not + # in ``df_meta``, because the dataflow rebuilds df_meta wholesale. + stats_gen: int = 0 + stats_status: str = "complete" + stats_reason: Optional[str] = None # Companion telemetry sink (#943): a fire-and-forget POST callable, built # once from the /load_expr payload's telemetry_url on the IOLoop (where # make_http_sink captures AsyncHTTPClient/IOLoop.current()). Stored here so @@ -111,6 +140,33 @@ def release_xorq_state(self) -> None: this field is the runtime escape hatch.""" +def begin_stats_generation(session: "SessionState") -> None: + """Start a new stats generation: the session's dataflow state has changed, + so stats for the previous one no longer describe it. The status restarts + from the session's policy pair.""" + session.stats_gen += 1 + session.stats_status, session.stats_reason = initial_stats_status( + session.stats_tier, session.stats_delivery) + + +def stats_meta(session: "SessionState") -> Optional[dict]: + """The ``df_meta.stats`` value for a session: ``{status, tier, gen}`` plus + ``reason`` when there is one. ``tier`` is the tier reached so far, which is + the one the session is headed for only once it is complete. + + ``None`` for a session on the default policy (inline delivery, full tier, + complete), which sends the message it always has; a client reads a missing + ``stats`` as complete.""" + if session.stats_status == "complete" and session.stats_delivery != "deferred": + return None + stats: dict = {"status": session.stats_status, + "tier": session.stats_tier if session.stats_status == "complete" else "schema", + "gen": session.stats_gen} + if session.stats_reason: + stats["reason"] = session.stats_reason + return stats + + def build_state_message(session: "SessionState", metadata: dict | None = None, search_string: str = "", reply_seq: int | None = None) -> dict: """Build the full ``initial_state`` WebSocket payload from a session. @@ -137,10 +193,16 @@ def build_state_message(session: "SessionState", metadata: dict | None = None, Returns: A dict ready to be JSON-serialised and sent to WebSocket clients. """ + # The dataflow rebuilds df_meta wholesale, so the stats status is injected + # here, into a copy, rather than stored in it. + df_meta = session.df_meta + stats = stats_meta(session) + if stats is not None: + df_meta = {**df_meta, "stats": stats} msg: dict = {"type": "initial_state", "protocol_version": PROTOCOL_VERSION, "metadata": metadata if metadata is not None else session.metadata, "prompt": session.prompt, "df_display_args": session.df_display_args, "df_data_dict": session.df_data_dict, - "df_meta": session.df_meta, "mode": session.mode} + "df_meta": df_meta, "mode": session.mode} if reply_seq is not None: msg["reply_seq"] = reply_seq if session.mode == "buckaroo": diff --git a/buckaroo/server/stats_wire.py b/buckaroo/server/stats_wire.py new file mode 100644 index 000000000..2089cf56b --- /dev/null +++ b/buckaroo/server/stats_wire.py @@ -0,0 +1,260 @@ +"""Server side of the stats wire protocol (rows-first s3). + +A deferred session (``stats_delivery="deferred"``) publishes its dataflow at the +schema tier and leaves the rest of the stats to later. Two kinds of client +reach them: + +* A client that advertised ``?caps=stats_update`` on its WebSocket URL gets a + stats-free ``initial_state`` (``df_meta.stats.status == "pending"``) and is + owed the stats for that generation. The server pushes them as a + ``stats_update`` once it has finished sending the client's first row reply + (``push_stats``, ADR-002 D2), and the client merges it. The client may also + ask with ``stats_request {stats_gen, scope}``; a request for a generation the + session has left gets ``stats_aborted``. +* Any other client gets complete messages: ``build_state_message_for`` runs the + missing stats synchronously before it builds a message for one, which is + today's cost, paid on the loop. + +Every send site goes through ``build_state_message_for`` (or ``broadcast_state``, +which calls it per client), because the session holds one shared snapshot and +the client is known only to the handler that owns the connection. + +A push and a request are each the whole run: one synchronous call that computes +the full stats and applies the final assignment (``set_stats_tier``, +``refresh_session_snapshot``). Resumable units generalize it later. +""" +import json +import logging +import time +import traceback +from contextlib import nullcontext +from typing import Any, Callable, Optional + +from buckaroo.pluggable_analysis_framework import perf_log +from buckaroo.server.data_loading import get_buckaroo_display_state +from buckaroo.server.session import SessionState, build_state_message + +log = logging.getLogger("buckaroo.server.stats_wire") + +# The capability a client advertises, as one of the comma-separated values of +# ``?caps=`` on the WebSocket URL: it merges ``stats_update`` messages. The +# connection is the only place it can be recorded, because ``open()`` sends the +# first message before the client has said anything. +STATS_UPDATE_CAP = "stats_update" + +# The scopes a ``stats_request`` can name. ``raw`` is the unfiltered table of +# the session's current generation; filtered stats exist only through a +# dataflow-field change (``quick_command_args.search``), which bumps the +# generation. +STATS_SCOPES = ("raw",) + + +def parse_caps(raw: str) -> frozenset: + """The capabilities in a ``?caps=a,b`` value.""" + return frozenset(cap for cap in (part.strip() for part in raw.split(",")) if cap) + + +def client_has_cap(client: Any, cap: str) -> bool: + """Whether a WebSocket handler recorded ``cap`` at open. Anything else with + no ``caps`` (a test double, a future non-WebSocket client) has none.""" + return cap in getattr(client, "caps", ()) + + +def session_dataflow(session: Optional[SessionState]) -> Any: + """The dataflow behind a buckaroo-mode session, or ``None`` (viewer and lazy + sessions have none).""" + if session is None or session.mode != "buckaroo": + return None + return session.xorq_dataflow if session.backend == "xorq" else session.dataflow + + +def refresh_session_snapshot(session: SessionState, dataflow: Any) -> None: + """Copy the dataflow's display state onto the session snapshot that new + clients and every push read, and re-apply ``component_config`` so theme + settings survive. The dataflow reaches no client until this runs.""" + refreshed = get_buckaroo_display_state(dataflow) + session.df_display_args = refreshed["df_display_args"] + session.df_data_dict = refreshed["df_data_dict"] + session.df_meta = refreshed["df_meta"] + session.buckaroo_options = refreshed["buckaroo_options"] + session.command_config = refreshed["command_config"] + if session.component_config and session.df_display_args: + for key in session.df_display_args: + dvc = session.df_display_args[key].get("df_viewer_config") + if dvc is not None: + dvc["component_config"] = {**dvc.get("component_config", {}), **session.component_config} + + +def complete_stats(session: SessionState) -> bool: + """Run the stats a pending session is missing, in one synchronous call, and + publish them: the final assignment (``set_stats_tier``), then the session + snapshot refreshed and the status set to ``complete`` in the same step, so + ``all_stats`` and ``df_meta.stats`` cannot disagree. Returns whether the + session is complete. + + The run is timed as ``stats.complete`` on the session's telemetry sink, or + on the sink already bound when the session has none. It is not a + ``firstpull.*`` span: a state change or a legacy client's connect completes + stats long after the load. + + A failure is the session's state for this generation (``error``, reason + ``stats_failed``) and is not retried by the next request; the next + generation starts clean.""" + if session.stats_status == "complete": + return True + dataflow = session_dataflow(session) + if session.stats_status != "pending" or dataflow is None: + return False + with ( + perf_log.telemetry_context(session.session_id, session.tele_sink) if session.tele_sink else nullcontext(), + perf_log.perf_span("stats.complete", session=session.session_id, stats_gen=session.stats_gen), + ): + try: + dataflow.set_stats_tier("full") + refresh_session_snapshot(session, dataflow) + except Exception: + log.error("stats run failed session=%s stats_gen=%s: %s", session.session_id, session.stats_gen, + traceback.format_exc()) + session.stats_status, session.stats_reason = "error", "stats_failed" + return False + session.stats_status, session.stats_reason = "complete", None + return True + + +def build_state_message_for(session: SessionState, client: Any, metadata: Optional[dict] = None, + reply_seq: Optional[int] = None) -> dict: + """The ``initial_state`` message for one client. + + A client without the ``stats_update`` capability on a pending deferred + session gets its missing stats run first, so its message is complete; a + capable client gets the session snapshot as it is (stats-free while the + session is pending) and is recorded as owed that generation's stats + (``client.stats_owed``), which ``push_stats`` delivers after its first row + reply. The search term is the recipient's + own (#851), and ``reply_seq`` is passed through to ``build_state_message`` + (#998).""" + if (session.stats_delivery == "deferred" and session.stats_status == "pending" + and not client_has_cap(client, STATS_UPDATE_CAP)): + complete_stats(session) + elif session.stats_status == "pending" and client_has_cap(client, STATS_UPDATE_CAP): + client.stats_owed = session.stats_gen + return build_state_message(session, metadata=metadata, search_string=getattr(client, "search_string", ""), + reply_seq=reply_seq) + + +def broadcast_state(session: SessionState, metadata: Optional[dict] = None, reset_search: bool = False, + reply_to: Any = None, reply_seq: Optional[int] = None, + highlight: Optional[Callable[[Any, str], Any]] = None) -> None: + """Send every connected client its own ``initial_state``. A client whose + write fails is dropped from the session. + + ``reset_search`` clears each client's live search term first (#851), for a + push that replaces the dataset. Clients that merge ``stats_update`` go + first: a legacy client's message completes the session's stats, and a + message built after that would carry them, so the capable client would + never see the pending state its own frame is meant to describe. + + ``reply_seq`` goes on the copy for ``reply_to`` only, the client whose + ``buckaroo_state_change`` this answers (#998); the others made no change, + so every copy is current for them. ``highlight(df_display_args, term)``, + when given, puts each client's own live-search highlight on its copy.""" + for client in sorted(session.ws_clients, key=lambda c: not client_has_cap(c, STATS_UPDATE_CAP)): + try: + if reset_search: + client.search_string = "" + msg = build_state_message_for(session, client, metadata=metadata, + reply_seq=reply_seq if client is reply_to else None) + if highlight is not None: + msg["df_display_args"] = highlight(msg["df_display_args"], getattr(client, "search_string", "")) + client.write_message(json.dumps(msg)) + except Exception: + session.ws_clients.discard(client) + + +def _aborted(stats_gen: Any, scope: Any, reason: str, session: Optional[SessionState] = None) -> dict: + """``stats_aborted``: the request was not run. ``stats_gen`` echoes the + request so the client can pair them; ``current_gen`` is the session's, so a + stale client can ask again without waiting for the next ``initial_state``.""" + msg: dict = {"type": "stats_aborted", "stats_gen": stats_gen, "scope": scope, "reason": reason} + if session is not None: + msg["current_gen"] = session.stats_gen + return msg + + +def _answer_stats_request(session: Optional[SessionState], stats_gen: Any, scope: Any, started: float) -> dict: + dataflow = session_dataflow(session) + if session is None or dataflow is None: + return _aborted(stats_gen, scope, "no_data") + if stats_gen != session.stats_gen: + return _aborted(stats_gen, scope, "stale", session) + if scope not in STATS_SCOPES: + return _aborted(stats_gen, scope, "unsupported_scope", session) + if session.stats_status == "not_computed": + return _aborted(stats_gen, scope, "not_requestable", session) + if session.stats_status == "error" or (session.stats_status == "pending" and not complete_stats(session)): + return _aborted(stats_gen, scope, "error", session) + # Complete: from here the answer is the dataflow's own all_stats, with no + # query, whether this request ran the stats or an earlier one did. + return {"type": "stats_update", "stats_gen": stats_gen, "scope": scope, "tier": dataflow.stats_tier, + "final": True, "payload": session.df_data_dict["all_stats"], + "elapsed_ms": round((time.perf_counter() - started) * 1000, 1)} + + +def handle_stats_request(session: Optional[SessionState], msg: dict) -> dict: + """Answer a ``stats_request {stats_gen, scope, columns?}``: a + ``stats_update`` carrying the complete ``all_stats`` as an inline wide + ``DFEnvelope`` (self-contained, so binary pairing stays single-slot), or a + ``stats_aborted``. ``columns`` is a hint a whole-run request ignores. + + The ``stats.request`` span records the request and its ``outcome`` + (``update`` or the abort reason), which is where updates sent, requests + dropped as stale and errors are counted. The caller binds the session's + telemetry sink around this call.""" + stats_gen, scope, columns = msg.get("stats_gen"), msg.get("scope", "raw"), msg.get("columns") + started = time.perf_counter() + with perf_log.perf_span("stats.request", session=session.session_id if session else None, stats_gen=stats_gen, + scope=scope, columns=len(columns) if isinstance(columns, list) else None) as span: + try: + reply = _answer_stats_request(session, stats_gen, scope, started) + except Exception: + log.error("stats_request error session=%s: %s", session.session_id if session else None, + traceback.format_exc()) + reply = _aborted(stats_gen, scope, "error", session) + _record_outcome(span, reply) + return reply + + +def _record_outcome(span: Any, reply: dict) -> None: + if reply["type"] == "stats_update": + span.set_attr(outcome="update", tier=reply["tier"]) + else: + span.set_attr(outcome=reply["reason"]) + + +def push_stats(session: Optional[SessionState], client: Any, stats_gen: int, rows_done: float) -> Optional[dict]: + """The ``stats_update`` (or ``stats_aborted``) owed to ``client`` for + ``stats_gen``, or ``None`` when nothing is owed any more: the client was + already pushed this generation, or the session has moved to another one (the + ``initial_state`` that started it registered its own push). ``rows_done`` is + the ``perf_counter`` time the row reply finished sending. + + The caller runs this on the write future of the client's first row reply + and sends the result, so rows reach the socket before any stats work starts. + The stats run once per session: a later connection's push is answered from + the session. The ``stats.push`` span records the gap since the rows and the + ``outcome``; the caller binds the session's telemetry sink around this call.""" + if getattr(client, "stats_owed", None) != stats_gen: + return None + client.stats_owed = None + if session is None or session.stats_gen != stats_gen: + return None + started = time.perf_counter() + with perf_log.perf_span("stats.push", session=session.session_id, stats_gen=stats_gen, + gap_ms=round((started - rows_done) * 1000, 1)) as span: + try: + reply = _answer_stats_request(session, stats_gen, "raw", started) + except Exception: + log.error("stats push error session=%s: %s", session.session_id, traceback.format_exc()) + reply = _aborted(stats_gen, "raw", "error", session) + _record_outcome(span, reply) + return reply diff --git a/buckaroo/server/websocket_handler.py b/buckaroo/server/websocket_handler.py index 7ef51aeb5..bc33fc4fc 100644 --- a/buckaroo/server/websocket_handler.py +++ b/buckaroo/server/websocket_handler.py @@ -2,6 +2,7 @@ import json import logging import os +import time import traceback from contextlib import nullcontext from urllib.parse import urlparse @@ -9,8 +10,9 @@ import tornado.websocket from buckaroo.pluggable_analysis_framework import perf_log -from buckaroo.server.data_loading import (handle_infinite_request, handle_infinite_request_buckaroo, handle_infinite_request_lazy, get_buckaroo_display_state) -from buckaroo.server.session import build_state_message +from buckaroo.server.data_loading import (handle_infinite_request, handle_infinite_request_buckaroo, handle_infinite_request_lazy) +from buckaroo.server.session import begin_stats_generation, dataflow_stats_tier +from buckaroo.server.stats_wire import (broadcast_state, build_state_message_for, handle_stats_request, parse_caps, push_stats, refresh_session_snapshot) def _handle_infinite_request_xorq(xorq_dataflow, payload_args, search_string=""): @@ -38,6 +40,15 @@ def open(self, session_id): # highlight overlay below — never broadcast, never stored on the # session. self.search_string = "" + # Capabilities the client advertised on the URL (``?caps=a,b``). Recorded + # per connection because this method sends the first message before the + # client can say anything, and the other send sites push one shared + # snapshot (stats_wire.build_state_message_for reads it per client). + self.caps = parse_caps(self.get_query_argument("caps", "")) + # The stats_gen this connection is owed a ``stats_update`` for: set when + # it is sent a pending frame (stats_wire.build_state_message_for), + # cleared by the push that follows its next row reply. + self.stats_owed = None sessions = self.application.settings["sessions"] sessions.add_ws_client(session_id, self) @@ -45,7 +56,7 @@ def open(self, session_id): # search_string="" — fresh connection, no per-client typing yet. session = sessions.get(session_id) if session and (session.df is not None or session.ldf is not None or session.xorq_dataflow is not None): - self.write_message(json.dumps(build_state_message(session, search_string=self.search_string))) + self.write_message(json.dumps(build_state_message_for(session, self))) def on_message(self, message): try: @@ -63,6 +74,22 @@ def on_message(self, message): # can drop a reply that an overlapping later change superseded. # Optional; a client that sends none gets none back. self._handle_buckaroo_state_change(msg.get("new_state") or {}, state_seq=msg.get("state_seq")) + elif msg_type == "stats_request": + self._handle_stats_request(msg) + + def _handle_stats_request(self, msg): + """Answer a client's ``stats_request`` with a ``stats_update`` or a + ``stats_aborted``. Synchronous, like ``infinite_request``: the request + runs the whole stats computation in this call (see ``stats_wire``). + + This branch is its own async context, so the session's telemetry sink + is bound here for the ``stats.request`` span and the stats spans under + it.""" + sessions = self.application.settings["sessions"] + session = sessions.get(self.session_id) + with perf_log.telemetry_context(self.session_id, session.tele_sink if session else None): + reply = handle_stats_request(session, msg) + self.write_message(json.dumps(reply)) def _handle_buckaroo_state_change(self, new_state, state_seq=None): sessions = self.application.settings["sessions"] @@ -109,36 +136,36 @@ def _handle_buckaroo_state_change(self, new_state, state_seq=None): log.debug("buckaroo_state_change no-op session=%s — skipping rebroadcast", self.session_id) return - # Propagate changes to the dataflow (mirrors BuckarooWidgetBase._buckaroo_state) - if old_state.get("post_processing") != new_state.get("post_processing"): - dataflow.post_processing_method = new_state.get("post_processing", "") - if old_state.get("cleaning_method") != new_state.get("cleaning_method"): - dataflow.cleaning_method = new_state.get("cleaning_method", "") - if old_state.get("quick_command_args") != new_state.get("quick_command_args"): - dataflow.quick_command_args = new_state.get("quick_command_args", {}) - - # Re-extract state from the dataflow — same helper works for both - # ServerDataflow and XorqServerDataflow (verified by probe). - buckaroo_state = get_buckaroo_display_state(dataflow) - session.df_display_args = buckaroo_state["df_display_args"] - session.df_data_dict = buckaroo_state["df_data_dict"] - session.df_meta = buckaroo_state["df_meta"] + # A deferred session whose stats were completed is at the full tier; + # the change reruns the cascade at the schema tier again (13 ms + # against the full stats), and the stats follow as requests. + # If applying the change raises, the session still describes the + # snapshot it had, so the tier goes back to the one that matches. + prior_tier = dataflow.stats_tier + try: + if session.stats_delivery == "deferred": + dataflow.stats_tier = dataflow_stats_tier(session.stats_tier, session.stats_delivery) + + # Propagate changes to the dataflow (mirrors BuckarooWidgetBase._buckaroo_state) + if old_state.get("post_processing") != new_state.get("post_processing"): + dataflow.post_processing_method = new_state.get("post_processing", "") + if old_state.get("cleaning_method") != new_state.get("cleaning_method"): + dataflow.cleaning_method = new_state.get("cleaning_method", "") + if old_state.get("quick_command_args") != new_state.get("quick_command_args"): + dataflow.quick_command_args = new_state.get("quick_command_args", {}) + + # Re-extract state from the dataflow — same helper works for both + # ServerDataflow and XorqServerDataflow (verified by probe). + refresh_session_snapshot(session, dataflow) + except Exception: + dataflow.stats_tier = prior_tier + raise + # The state the stats describe has changed, so the generation moves on. + begin_stats_generation(session) # Strip search_string before snapshotting onto the session — it # belongs to this client only (#851), so a future client that # connects shouldn't inherit it via build_state_message. session.buckaroo_state = {k: v for k, v in new_state.items() if k != "search_string"} - session.buckaroo_options = buckaroo_state["buckaroo_options"] - session.command_config = buckaroo_state["command_config"] - - # Re-apply component_config so theme settings survive state changes - if session.component_config and session.df_display_args: - for key in session.df_display_args: - dvc = session.df_display_args[key].get("df_viewer_config") - if dvc is not None: - dvc["component_config"] = { - **dvc.get("component_config", {}), - **session.component_config, - } # Broadcast updated state to all connected clients. Each # client gets its own search_string re-injected so a @@ -148,15 +175,7 @@ def _handle_buckaroo_state_change(self, new_state, state_seq=None): # doesn't drop it. # Only the originating client gets reply_seq (#998): the # others made no change, so every copy is current for them. - for client in list(session.ws_clients): - try: - search = getattr(client, "search_string", "") - msg = build_state_message(session, search_string=search, - reply_seq=state_seq if client is self else None) - msg["df_display_args"] = self._with_highlight(session.df_display_args, search) - client.write_message(json.dumps(msg)) - except Exception: - session.ws_clients.discard(client) + broadcast_state(session, reply_to=self, reply_seq=state_seq, highlight=self._with_highlight) except Exception: tb = traceback.format_exc() log.error("buckaroo_state_change error session=%s: %s", self.session_id, tb) @@ -207,10 +226,14 @@ def _send_client_state(self, session, buckaroo_state, reply_seq=None): is re-injected so the reply doesn't clear the search box (Codex P1 on #854). ``reply_seq`` is the change's ``state_seq`` (#998), so the client drops this reply if it has sent a later change. + + The message is built first: for a client without the + ``stats_update`` capability it completes the session's stats, which + replaces the display config, and the highlight goes on that one. """ - msg = build_state_message(session, search_string=self.search_string, reply_seq=reply_seq) + msg = build_state_message_for(session, self, reply_seq=reply_seq) msg["buckaroo_state"] = {**buckaroo_state, "search_string": self.search_string} - msg["df_display_args"] = self._with_highlight(session.df_display_args, self.search_string) + msg["df_display_args"] = self._with_highlight(msg["df_display_args"], self.search_string) try: self.write_message(json.dumps(msg)) except Exception: @@ -258,6 +281,8 @@ def _dispatch(pa): tele_sink = session.tele_sink if not_seen else None first_payload = not_seen and (perf_log.enabled() or tele_sink is not None) + last_write = None + def _dispatch_and_send(pa, span_label): # Dispatch one window and send its two-frame reply: a JSON text # frame, then the binary Parquet frame when non-empty. On the initial @@ -267,11 +292,12 @@ def _dispatch_and_send(pa, span_label): # v2. span = (perf_log.perf_span(span_label, session=self.session_id) if first_payload else nullcontext()) + nonlocal last_write with span: resp, parquet = _dispatch(pa) - self.write_message(json.dumps(resp)) + last_write = self.write_message(json.dumps(resp)) if parquet: - self.write_message(parquet, binary=True) + last_write = self.write_message(parquet, binary=True) try: with perf_log.telemetry_context(self.session_id, tele_sink): @@ -289,6 +315,30 @@ def _dispatch_and_send(pa, span_label): log.error("infinite_request error session=%s: %s", self.session_id, tb) self.write_message(json.dumps({"type": "infinite_resp", "key": payload_args, "length": 0, "error_info": tb if _BUCKAROO_DEBUG else "Request failed"})) + if self.stats_owed is not None and last_write is not None: + # D2: the stats follow the rows. The continuation runs once the last + # frame of this reply is written, so none of the stats work can + # delay it. + owed, rows_done = self.stats_owed, time.perf_counter() + last_write.add_done_callback(lambda _: self._push_stats(owed, rows_done)) + + def _push_stats(self, stats_gen, rows_done): + """Send the ``stats_update`` this connection is owed for ``stats_gen``. + Does nothing when the connection has closed since the row reply, or + when ``push_stats`` finds nothing owed (a newer generation, or already + pushed). A newer generation's own pending frame registered its own + push.""" + if self.ws_connection is None: + return + sessions = self.application.settings["sessions"] + session = sessions.get(self.session_id) + with perf_log.telemetry_context(self.session_id, session.tele_sink if session else None): + reply = push_stats(session, self, stats_gen, rows_done) + if reply is not None: + try: + self.write_message(json.dumps(reply)) + except tornado.websocket.WebSocketClosedError: + log.debug("stats push write failed for session=%s", self.session_id) def on_close(self): sessions = self.application.settings["sessions"] diff --git a/packages/buckaroo-js-core/pw-tests/server-load-expr.spec.ts b/packages/buckaroo-js-core/pw-tests/server-load-expr.spec.ts index 329b34b17..755c17e9d 100644 --- a/packages/buckaroo-js-core/pw-tests/server-load-expr.spec.ts +++ b/packages/buckaroo-js-core/pw-tests/server-load-expr.spec.ts @@ -32,9 +32,9 @@ function buildExprDir(buildsRoot: string): string { return out.trim().split('\n').pop()!; } -async function loadExpr(request: any, sessionId: string, buildDir: string) { +async function loadExpr(request: any, sessionId: string, buildDir: string, extra: Record = {}) { const resp = await request.post(`${BASE}/load_expr`, { - data: { session: sessionId, build_dir: buildDir, no_browser: true }, + data: { session: sessionId, build_dir: buildDir, no_browser: true, ...extra }, }); if (!resp.ok()) { throw new Error(`/load_expr failed (${resp.status()}): ${await resp.text()}`); @@ -46,6 +46,17 @@ async function getPinnedRowCount(page: any): Promise { return await page.locator('.df-viewer .ag-floating-top-container .ag-row').count(); } +// The text of every pinned cell in one column, top to bottom. In the summary +// view these are the stats: dtype, length, min, max, distinct count and so on. +async function getPinnedCellTexts(page: any, colId: string): Promise { + const texts = await page.locator(`.df-viewer .ag-floating-top-container [col-id="${colId}"]`).allInnerTexts(); + return texts.map((t: string) => t.trim()); +} + +async function showSummaryView(page: any) { + await page.locator('.status-bar').locator('select').first().selectOption('summary'); +} + test.describe('POST /load_expr', () => { let buildsRoot: string; let buildDir: string; @@ -113,4 +124,31 @@ test.describe('POST /load_expr', () => { await expect.poll(() => getCellText(page, COL.idx, 1), { timeout: 15_000 }).toBe('3'); expect(await getCellText(page, COL.name, 1)).toBe('alpha'); }); + + test('eager stats reach the DOM: the summary view shows the computed values', async ({ page, request }) => { + const session = `lx-eager-${Date.now()}`; + await loadExpr(request, session, buildDir, { stats_delivery: 'inline' }); + + await page.goto(`${BASE}/s/${session}`); + await waitForGrid(page); + await showSummaryView(page); + + // 'name' has 7 distinct values; 'idx' runs 0 to 9. + await expect.poll(() => getPinnedCellTexts(page, COL.name), { timeout: 15_000 }).toContain('7'); + expect(await getPinnedCellTexts(page, COL.idx)).toEqual(expect.arrayContaining(['0', '9'])); + }); + + test('deferred stats: the rows render first and the stats then reach the DOM', async ({ page, request }) => { + const session = `lx-deferred-${Date.now()}`; + await loadExpr(request, session, buildDir, { stats_delivery: 'deferred' }); + + await page.goto(`${BASE}/s/${session}`); + await waitForGrid(page); + expect(await getCellText(page, COL.idx, 0)).toBe('0'); + expect(await getCellText(page, COL.name, 0)).toBe('alpha'); + + await showSummaryView(page); + await expect.poll(() => getPinnedCellTexts(page, COL.name), { timeout: 15_000 }).toContain('7'); + expect(await getPinnedCellTexts(page, COL.idx)).toEqual(expect.arrayContaining(['0', '9'])); + }); }); diff --git a/tests/unit/server/test_load_expr.py b/tests/unit/server/test_load_expr.py index fa5f2e852..3c4a99cf7 100644 --- a/tests/unit/server/test_load_expr.py +++ b/tests/unit/server/test_load_expr.py @@ -1,5 +1,6 @@ """End-to-end tests for POST /load_expr — server load path for XorqBuckarooInfiniteWidget over a xorq/ibis expression.""" +import datetime import gc import io import json @@ -8,12 +9,15 @@ import sys import tempfile import weakref +from contextlib import contextmanager from pathlib import Path +from types import SimpleNamespace from unittest.mock import MagicMock, patch import pandas as pd import pyarrow.parquet as pq import pytest +import tornado.gen import tornado.httpclient import tornado.testing import tornado.websocket @@ -33,10 +37,15 @@ from buckaroo.dataflow.sd_cache import split_chain_by_scope # noqa: E402 from buckaroo.jlisp.lisp_utils import s as lisp_sym # noqa: E402 +from buckaroo.pluggable_analysis_framework import perf_log # noqa: E402 from buckaroo.pluggable_analysis_framework.col_analysis import ColAnalysis # noqa: E402 from buckaroo.pluggable_analysis_framework.xorq_stat_pipeline import XorqStatPipeline # noqa: E402 +from buckaroo.serialization_utils import resolve_summary_stats_payload # noqa: E402 from buckaroo.server import telemetry, xorq_loading # noqa: E402 from buckaroo.server.app import make_app as _make_app # noqa: E402 +from buckaroo.server.session import SessionState, begin_stats_generation # noqa: E402 +from buckaroo.server.stats_wire import complete_stats # noqa: E402 +from buckaroo.server.websocket_handler import DataStreamHandler # noqa: E402 pytestmark = pytest.mark.skipif( sys.platform == "win32", @@ -1589,6 +1598,659 @@ async def test_reload_expr_accepts_a_new_pair_and_stores_it(self): shutil.rmtree(project_root, ignore_errors=True) +def _build_stats_wire_dir(builds_root): + """Build the ``_stats_tier_expr`` table to ``builds_root``. Its float column + has a 1e9 maximum, so full stats change that column's ``minWidth`` and a + message built from the schema tier differs from a complete one there.""" + expr = xo.memtable({ + "price": [12.5, 18.9, 7.4, 22.1, 1e9], + "qty": [1, 2, 1, 3, 2], + "category": ["a", "b", "a", "c", "b"]}, name="t") + return str(xo.build_expr(expr, builds_dir=builds_root)) + + +def _state_change(**changes): + new_state = {"post_processing": "", "cleaning_method": "", "quick_command_args": {}, + "df_display": "main", "show_commands": False, "sampled": False, "search_string": ""} + return json.dumps({"type": "buckaroo_state_change", "new_state": {**new_state, **changes}}) + + +def _stats_request(stats_gen, **fields): + return json.dumps({"type": "stats_request", "stats_gen": stats_gen, "scope": "raw", **fields}) + + +async def _read_json(ws, timeout=3.0): + """The next frame on ``ws``, decoded. Raises AssertionError, not a hang, + when none arrives.""" + try: + frame = await tornado.gen.with_timeout( + datetime.timedelta(seconds=timeout), ws.read_message()) + except tornado.gen.TimeoutError: + raise AssertionError(f"no frame within {timeout}s") from None + assert frame is not None, "the connection closed" + return json.loads(frame) + + +def _rows_by_stat(payload): + """The decoded ``all_stats`` rows of a payload, keyed by stat name.""" + return {row["index"]: row for row in resolve_summary_stats_payload(payload)} + + +def _comparable(frame): + """An ``initial_state`` frame as a JSON string, for comparing a legacy + client's frame with what an inline session sends: ``all_stats`` decoded + (its parquet bytes follow column order) and ``df_meta.stats`` dropped.""" + frame = json.loads(json.dumps(frame)) + frame["df_data_dict"]["all_stats"] = _rows_by_stat(frame["df_data_dict"]["all_stats"]) + frame["df_meta"].pop("stats", None) + return _as_json(frame) + + +@contextmanager +def _count_stat_queries(): + """Record every query ``XorqStatPipeline`` sends while the block runs.""" + queries = [] + original = XorqStatPipeline._execute + + def spy(pipeline, query): + queries.append(query) + return original(pipeline, query) + + with patch.object(XorqStatPipeline, "_execute", spy): + yield queries + + +def _deferred_session(): + """A pending deferred session over a schema-tier dataflow, as /load_expr + leaves one.""" + session = SessionState(session_id="s", path="", mode="buckaroo", backend="xorq", + xorq_dataflow=_build_dataflow(stats_tier="schema"), stats_tier="full", stats_delivery="deferred") + begin_stats_generation(session) + return session + + +class TestCompleteStats: + """``complete_stats`` takes a deferred session's dataflow to the full tier.""" + + def test_a_cache_hit_keeps_the_errs_of_the_run_that_filled_the_cache(self): + session = _deferred_session() + dataflow = session.xorq_dataflow + errs = {"a": {"stat": "boom"}} + real = dataflow._get_summary_sd + dataflow._get_summary_sd = lambda df, scope="filt": (real(df, scope)[0], errs) if dataflow.stats_tier == "full" else real(df, + scope) + assert complete_stats(session) + + # A search and its clearing, each reset to the schema tier as a state + # change does, bring the dataflow back to a state whose stats are cached. + for args in ({"search": ["a"]}, {}): + dataflow.stats_tier = "schema" + dataflow.quick_command_args = args + begin_stats_generation(session) + assert complete_stats(session) + + assert dataflow.errs == errs + + def test_spans_reach_the_bound_sink_when_the_session_has_none_and_the_run_is_not_a_first_pull(self): + session = _deferred_session() + records = [] + + with perf_log.telemetry_context("outer", records.append): + assert complete_stats(session) + + names = {record["name"] for record in records} + assert "stats.complete" in names + assert "firstpull.stats_total" not in names + + +class TestStatsWire(tornado.testing.AsyncHTTPTestCase): + """``stats_request``, ``stats_update`` and ``stats_aborted`` on a deferred + ``/load_expr`` session, the ``stats_gen`` counter and ``df_meta.stats``, and + the per-connection capability (``?caps=stats_update``) that keeps a client + without it on complete messages (rows-first s3).""" + + def get_app(self): + return make_app() + + def setUp(self): + super().setUp() + self.builds_root = tempfile.mkdtemp() + self.project_root = tempfile.mkdtemp() + self.addCleanup(shutil.rmtree, self.builds_root, ignore_errors=True) + self.addCleanup(shutil.rmtree, self.project_root, ignore_errors=True) + pp_dir = os.path.join(self.project_root, "post_processing") + os.makedirs(pp_dir) + with open(os.path.join(pp_dir, "first_three.py"), "w") as f: + f.write("def process(expr):\n return expr.limit(3)\n") + self.build_path = _build_stats_wire_dir(self.builds_root) + self.clients = [] + + def tearDown(self): + for ws in self.clients: + ws.close() + super().tearDown() + + def _session(self, sid): + return self._app.settings["sessions"].get(sid) + + async def _load(self, sid, **body): + resp = await _post(self.get_http_port(), "/load_expr", + {"session": sid, "build_dir": self.build_path, "project_root": self.project_root, **body}) + self.assertEqual(resp.code, 200, resp.body) + + async def _connect(self, sid, caps=None): + """Open a WebSocket, with ``caps`` as ``?caps=``. Returns it with its + first ``initial_state``.""" + suffix = f"?caps={caps}" if caps else "" + ws = await tornado.websocket.websocket_connect( + f"ws://localhost:{self.get_http_port()}/ws/{sid}{suffix}") + self.clients.append(ws) + return ws, await _read_json(ws) + + def _stats(self, frame): + stats = frame["df_meta"].get("stats") + self.assertIsNotNone(stats, "initial_state carries no df_meta.stats") + return stats + + async def _inline_frame(self, sid, **changes): + """What an inline session of the same build sends a client: the complete + message every complete message here must match. ``changes`` are sent as + a state change first, and the broadcast that follows is returned.""" + await self._load(sid) + ws, frame = await self._connect(sid) + if changes: + ws.write_message(_state_change(**changes)) + frame = await _read_json(ws) + return frame + + def _assert_pending(self, frame, gen): + stats = self._stats(frame) + self.assertEqual((stats["status"], stats["tier"], stats["gen"]), ("pending", "schema", gen)) + self.assertEqual(list(_rows_by_stat(frame["df_data_dict"]["all_stats"])), ["dtype"], + "a pending frame must carry the schema tier only") + + def _assert_complete(self, frame, gen, inline_frame): + stats = self._stats(frame) + self.assertEqual((stats["status"], stats["tier"], stats["gen"]), ("complete", "full", gen)) + self.assertIn("histogram_bins", _rows_by_stat(frame["df_data_dict"]["all_stats"])) + self.assertEqual(_comparable(frame), _comparable(inline_frame)) + + async def _pair(self, sid): + """A caps client, then a legacy client, on a deferred session. The + legacy client's connect completes the stats, so a later push is the + first thing to put the session back to pending.""" + await self._load(sid, stats_delivery="deferred") + a, a_open = await self._connect(sid, caps="stats_update") + b, b_open = await self._connect(sid) + return a, a_open, b, b_open + + @tornado.testing.gen_test + async def test_stats_request_returns_a_stats_update_that_completes_the_stats(self): + await self._load("sw-complete", stats_delivery="deferred") + inline = await self._inline_frame("sw-complete-inline") + ws, first = await self._connect("sw-complete", caps="stats_update") + gen = self._stats(first)["gen"] + self._assert_pending(first, gen) + self.assertEqual(first["protocol_version"], 1) + + ws.write_message(_stats_request(gen)) + update = await _read_json(ws) + + self.assertEqual(update["type"], "stats_update") + self.assertEqual(update["stats_gen"], gen) + self.assertEqual((update["scope"], update["tier"], update["final"]), ("raw", "full", True)) + self.assertEqual((update["payload"]["format"], update["payload"]["layout"]), + ("parquet_b64", "wide"), "the payload must be inline, so binary pairing stays single-slot") + self.assertEqual(_rows_by_stat(update["payload"]), + _rows_by_stat(inline["df_data_dict"]["all_stats"])) + self.assertIsInstance(update["elapsed_ms"], (int, float)) + + @tornado.testing.gen_test + async def test_rows_are_served_before_and_after_the_stats_request(self): + await self._load("sw-rows", stats_delivery="deferred") + ws, first = await self._connect("sw-rows", caps="stats_update") + window = {"start": 0, "end": 5, "sourceName": "default", "origEnd": 5} + ws.write_message(json.dumps({"type": "infinite_request", "payload_args": window})) + self.assertEqual((await _read_json(ws))["length"], 5) + pq.read_table(io.BytesIO(await ws.read_message())) + self.assertEqual((await _read_json(ws))["type"], "stats_update", "the push follows the rows") + ws.write_message(_stats_request(self._stats(first)["gen"])) + self.assertEqual((await _read_json(ws))["type"], "stats_update") + ws.write_message(json.dumps({"type": "infinite_request", "payload_args": window})) + self.assertEqual((await _read_json(ws))["type"], "infinite_resp", + "a stats_update must leave no stray binary frame in the stream") + + async def _first_rows(self, ws): + """Ask for the first window and read its two frames.""" + ws.write_message(json.dumps({"type": "infinite_request", + "payload_args": {"start": 0, "end": 5, "sourceName": "default", "origEnd": 5}})) + resp = await _read_json(ws) + self.assertEqual(resp["type"], "infinite_resp") + pq.read_table(io.BytesIO(await ws.read_message())) + + @tornado.testing.gen_test + async def test_stats_are_pushed_after_the_first_row_reply_without_a_request(self): + await self._load("sw-push", stats_delivery="deferred") + inline = await self._inline_frame("sw-push-inline") + ws, first = await self._connect("sw-push", caps="stats_update") + gen = self._stats(first)["gen"] + + await self._first_rows(ws) + update = await _read_json(ws) + + self.assertEqual((update["type"], update["stats_gen"], update["final"]), ("stats_update", gen, True)) + self.assertEqual(_rows_by_stat(update["payload"]), _rows_by_stat(inline["df_data_dict"]["all_stats"])) + await self._first_rows(ws) # the push is owed once: a second row reply is followed by no second update + + @tornado.testing.gen_test + async def test_a_second_connection_is_pushed_the_stats_the_first_connection_ran(self): + await self._load("sw-push-two", stats_delivery="deferred") + a, _ = await self._connect("sw-push-two", caps="stats_update") + b, _ = await self._connect("sw-push-two", caps="stats_update") + await self._first_rows(a) + self.assertEqual((await _read_json(a))["type"], "stats_update") + + with _count_stat_queries() as queries: + await self._first_rows(b) + self.assertEqual((await _read_json(b))["type"], "stats_update") + self.assertEqual(queries, [], "stats are computed once per session, at the first connection's push") + + @tornado.testing.gen_test + async def test_a_push_emits_a_stats_push_span(self): + captured: list = [] + with patch.object(telemetry, "make_http_sink", lambda url, **kw: captured.append): + await self._load("sw-push-span", stats_delivery="deferred", + telemetry_url="http://companion.invalid/internal/telemetry") + ws, first = await self._connect("sw-push-span", caps="stats_update") + await self._first_rows(ws) + await _read_json(ws) + + (push,) = [r for r in captured if r["name"] == "stats.push"] + self.assertEqual(push["trace"], "sw-push-span") + self.assertEqual((push["attrs"]["stats_gen"], push["attrs"]["outcome"], push["attrs"]["tier"]), + (self._stats(first)["gen"], "update", "full")) + self.assertGreaterEqual(push["attrs"]["gap_ms"], 0) + + @tornado.testing.gen_test + async def test_a_stale_stats_gen_gets_stats_aborted_and_runs_nothing(self): + await self._load("sw-stale", stats_delivery="deferred") + ws, first = await self._connect("sw-stale", caps="stats_update") + gen = self._stats(first)["gen"] + + with _count_stat_queries() as queries: + ws.write_message(_stats_request(gen - 1)) + aborted = await _read_json(ws) + self.assertEqual(aborted["type"], "stats_aborted") + self.assertEqual((aborted["stats_gen"], aborted["current_gen"], aborted["reason"]), + (gen - 1, gen, "stale")) + self.assertEqual(queries, [], "a stale request must run no stat query") + + ws.write_message(_stats_request(gen)) + self.assertEqual((await _read_json(ws))["type"], "stats_update", + "the stale request must not have consumed or failed the session's stats") + + @tornado.testing.gen_test + async def test_an_unsupported_scope_gets_stats_aborted(self): + await self._load("sw-scope", stats_delivery="deferred") + ws, first = await self._connect("sw-scope", caps="stats_update") + ws.write_message(_stats_request(self._stats(first)["gen"], scope="filt")) + aborted = await _read_json(ws) + self.assertEqual((aborted["type"], aborted["reason"], aborted["scope"]), + ("stats_aborted", "unsupported_scope", "filt")) + + @tornado.testing.gen_test + async def test_load_expr_and_reload_expr_bump_stats_gen(self): + await self._load("sw-gen", stats_delivery="deferred") + ws, first = await self._connect("sw-gen", caps="stats_update") + gen = self._stats(first)["gen"] + + await self._load("sw-gen", stats_delivery="deferred", force_reload=True) + pushed = await _read_json(ws) + self.assertGreater(self._stats(pushed)["gen"], gen, "/load_expr must bump stats_gen") + gen = self._stats(pushed)["gen"] + + resp = await _post(self.get_http_port(), "/reload_expr/sw-gen", {}) + self.assertEqual(resp.code, 200, resp.body) + pushed = await _read_json(ws) + self.assertGreater(self._stats(pushed)["gen"], gen, "/reload_expr must bump stats_gen") + + @tornado.testing.gen_test + async def test_df_meta_stats_survives_a_dataflow_field_change(self): + await self._load("sw-change", stats_delivery="deferred") + ws, first = await self._connect("sw-change", caps="stats_update") + gen = self._stats(first)["gen"] + ws.write_message(_stats_request(gen)) + self.assertEqual((await _read_json(ws))["type"], "stats_update") + + with _count_stat_queries() as queries: + ws.write_message(_state_change(quick_command_args={"search": ["a"]})) + changed = await _read_json(ws) + self._assert_pending(changed, gen + 1) + self.assertEqual(queries, [], "a dataflow-field change on a deferred session must run no stats") + # The dataflow rebuilt df_meta for the filtered state; its keys are still there. + self.assertEqual((changed["df_meta"]["total_rows"], changed["df_meta"]["filtered_rows"]), (5, 2)) + + ws.write_message(_stats_request(gen)) + aborted = await _read_json(ws) + self.assertEqual((aborted["type"], aborted["current_gen"], aborted["reason"]), + ("stats_aborted", gen + 1, "stale")) + ws.write_message(_stats_request(gen + 1)) + update = await _read_json(ws) + self.assertEqual((update["type"], update["stats_gen"], update["final"]), + ("stats_update", gen + 1, True)) + + @tornado.testing.gen_test + async def test_a_legacy_client_stays_complete_while_a_caps_client_gets_a_stats_free_frame(self): + sid = "sw-ab" + inline = await self._inline_frame("sw-ab-inline", post_processing="first_three") + await self._load(sid, stats_delivery="deferred") + a1, a1_open = await self._connect(sid, caps="other,stats_update") + a2, _ = await self._connect(sid, caps="stats_update") + # Connecting runs the missing stats for a legacy client, so the session + # is complete when B1 changes it. + b1, b1_open = await self._connect(sid, caps="unknown") + b2, _ = await self._connect(sid) + gen = self._stats(a1_open)["gen"] + self.assertEqual(self._stats(b1_open)["status"], "complete") + + with _count_stat_queries() as queries: + b1.write_message(_state_change(post_processing="first_three")) + frames = {name: await _read_json(ws) + for name, ws in (("a1", a1), ("a2", a2), ("b1", b1), ("b2", b2))} + for name in ("a1", "a2"): + self._assert_pending(frames[name], gen + 1) + for name in ("b1", "b2"): + self._assert_complete(frames[name], gen + 1, inline) + ran = len(queries) + self.assertGreater(ran, 0, "B's frame must have run the stats it carries") + + a1.write_message(_stats_request(gen + 1)) + update = await _read_json(a1) + self.assertEqual((update["type"], update["stats_gen"]), ("stats_update", gen + 1)) + self.assertEqual(_rows_by_stat(update["payload"]), + _rows_by_stat(frames["b1"]["df_data_dict"]["all_stats"])) + self.assertEqual(len(queries), ran, + "A's request must be answered from the stats B's frame computed") + + @tornado.testing.gen_test + async def test_load_expr_push_keeps_a_legacy_client_complete(self): + sid = "sw-push-load-expr" + inline = await self._inline_frame("sw-push-load-expr-inline") + a, a_open, b, _ = await self._pair(sid) + gen = self._stats(a_open)["gen"] + await self._load(sid, stats_delivery="deferred", force_reload=True) + self._assert_pending(await _read_json(a), gen + 1) + self._assert_complete(await _read_json(b), gen + 1, inline) + + @tornado.testing.gen_test + async def test_reload_expr_push_keeps_a_legacy_client_complete(self): + sid = "sw-push-reload" + inline = await self._inline_frame("sw-push-reload-inline") + a, a_open, b, _ = await self._pair(sid) + gen = self._stats(a_open)["gen"] + resp = await _post(self.get_http_port(), f"/reload_expr/{sid}", {}) + self.assertEqual(resp.code, 200, resp.body) + self._assert_pending(await _read_json(a), gen + 1) + self._assert_complete(await _read_json(b), gen + 1, inline) + + @tornado.testing.gen_test + async def test_load_push_leaves_both_clients_complete(self): + """/load swaps the session to pandas, which has no deferred stats: the + stored policy must not leave it looking pending.""" + sid = "sw-push-load" + a, _, b, _ = await self._pair(sid) + csv_fd, csv_path = tempfile.mkstemp(suffix=".csv") + os.close(csv_fd) + try: + pd.DataFrame({"x": [1, 2, 3], "y": ["p", "q", "r"]}).to_csv(csv_path, index=False) + resp = await _post(self.get_http_port(), "/load", + {"session": sid, "path": csv_path, "mode": "buckaroo"}) + self.assertEqual(resp.code, 200, resp.body) + finally: + os.unlink(csv_path) + frames = [await _read_json(a), await _read_json(b)] + for frame in frames: + self.assertNotIn("stats", frame["df_meta"]) + self.assertIn("histogram_bins", _rows_by_stat(frame["df_data_dict"]["all_stats"])) + self.assertEqual(_comparable(frames[0]), _comparable(frames[1])) + + @tornado.testing.gen_test + async def test_load_compare_push_leaves_both_clients_complete(self): + sid = "sw-push-compare" + a, _, b, _ = await self._pair(sid) + paths = [] + try: + for frame in (pd.DataFrame({"id": [1, 2], "v": [10, 20]}), + pd.DataFrame({"id": [1, 3], "v": [10, 30]})): + fd, path = tempfile.mkstemp(suffix=".csv") + os.close(fd) + frame.to_csv(path, index=False) + paths.append(path) + resp = await _post(self.get_http_port(), "/load_compare", + {"session": sid, "path1": paths[0], "path2": paths[1], "join_columns": ["id"]}) + self.assertEqual(resp.code, 200, resp.body) + finally: + for path in paths: + os.unlink(path) + frames = [await _read_json(a), await _read_json(b)] + for frame in frames: + self.assertNotIn("stats", frame["df_meta"]) + self.assertEqual(_comparable(frames[0]), _comparable(frames[1])) + + @tornado.testing.gen_test + async def test_highlight_overlay_is_complete_for_a_legacy_client_and_stats_free_for_a_caps_client(self): + sid = "sw-overlay" + await self._load(sid, stats_delivery="deferred") + a, a_open = await self._connect(sid, caps="stats_update") + gen = self._stats(a_open)["gen"] + + a.write_message(_state_change(search_string="ca")) + a_overlay = await _read_json(a) + self._assert_pending(a_overlay, gen) + + b, _ = await self._connect(sid) + b.write_message(_state_change(search_string="ca")) + b_overlay = await _read_json(b) + self.assertEqual(self._stats(b_overlay)["status"], "complete") + self.assertIn("histogram_bins", _rows_by_stat(b_overlay["df_data_dict"]["all_stats"])) + # The overlay's display config is the complete one with the highlight on top. + session_args = self._session(sid).df_display_args["main"]["df_viewer_config"]["column_config"] + by_col = {cc["col_name"]: cc for cc in b_overlay["df_display_args"]["main"]["df_viewer_config"]["column_config"]} + for expected in session_args: + self.assertEqual(by_col[expected["col_name"]]["ag_grid_specs"], expected["ag_grid_specs"]) + highlighted = [cc for cc in by_col.values() + if cc.get("displayer_args", {}).get("highlight_phrase") == ["ca"]] + self.assertTrue(highlighted, "the overlay must still carry the highlight") + + @tornado.testing.gen_test + async def test_overlay_for_a_legacy_client_is_built_after_its_stats_are_completed(self): + """Completing the stats replaces the session's display config, so an + overlay that copied it first would send a legacy client the schema + tier's. Every send completes a session that has a legacy client + connected, so only a direct call reaches a pending session here.""" + await self._load("sw-overlay-order", stats_delivery="deferred") + session = self._session("sw-overlay-order") + sent: list = [] + legacy = SimpleNamespace(search_string="ca", caps=frozenset(), session_id="sw-overlay-order", + write_message=sent.append, _with_highlight=DataStreamHandler._with_highlight) + self.assertEqual(getattr(session, "stats_status", None), "pending") + + DataStreamHandler._send_client_state(legacy, session, session.buckaroo_state) + + self.assertEqual(session.stats_status, "complete") + overlay = json.loads(sent[0])["df_display_args"]["main"]["df_viewer_config"]["column_config"] + complete = session.df_display_args["main"]["df_viewer_config"]["column_config"] + self.assertEqual({cc["col_name"]: cc["ag_grid_specs"] for cc in overlay}, + {cc["col_name"]: cc["ag_grid_specs"] for cc in complete}) + + @tornado.testing.gen_test + async def test_a_failed_state_change_leaves_the_dataflow_at_the_tier_the_session_describes(self): + """The change resets a completed session's dataflow to the schema tier + before it applies the change. When applying it raises, the session still + describes the full stats of the snapshot it has, so the dataflow must + stay at the tier that matches.""" + await self._load("sw-failed-change", stats_delivery="deferred") + ws, _ = await self._connect("sw-failed-change") + session = self._session("sw-failed-change") + self.assertEqual((session.stats_status, session.xorq_dataflow.stats_tier), ("complete", "full")) + + with patch("buckaroo.server.websocket_handler.refresh_session_snapshot", side_effect=RuntimeError("boom")): + ws.write_message(_state_change(quick_command_args={"search": ["a"]})) + self.assertEqual((await _read_json(ws))["error_code"], "state_change_error") + + self.assertEqual((session.stats_status, session.xorq_dataflow.stats_tier), ("complete", "full")) + + @tornado.testing.gen_test + async def test_a_client_connecting_after_completion_gets_the_complete_state(self): + inline = await self._inline_frame("sw-late-inline") + await self._load("sw-late", stats_delivery="deferred") + a, a_open = await self._connect("sw-late", caps="stats_update") + gen = self._stats(a_open)["gen"] + a.write_message(_stats_request(gen)) + self.assertEqual((await _read_json(a))["type"], "stats_update") + + with _count_stat_queries() as queries: + _, late = await self._connect("sw-late", caps="stats_update") + _, legacy = await self._connect("sw-late") + self._assert_complete(late, gen, inline) + self._assert_complete(legacy, gen, inline) + self.assertEqual(queries, [], "stats computed once must not be computed again for a later client") + + @tornado.testing.gen_test + async def test_a_legacy_client_connecting_to_a_pending_session_completes_it(self): + inline = await self._inline_frame("sw-open-inline") + await self._load("sw-open", stats_delivery="deferred") + _, legacy = await self._connect("sw-open") + self._assert_complete(legacy, self._stats(legacy)["gen"], inline) + + @tornado.testing.gen_test + async def test_a_session_targeting_the_schema_tier_is_not_computed(self): + await self._load("sw-schema", stats_tier="schema", stats_delivery="deferred") + ws, first = await self._connect("sw-schema", caps="stats_update") + stats = self._stats(first) + self.assertEqual((stats["status"], stats["tier"], stats["reason"]), ("not_computed", "schema", "host")) + ws.write_message(_stats_request(stats["gen"])) + aborted = await _read_json(ws) + self.assertEqual((aborted["type"], aborted["reason"]), ("stats_aborted", "not_requestable")) + # Nothing is missing relative to the schema tier, so a legacy client + # gets the schema tier, as it does from an inline schema session. + _, legacy = await self._connect("sw-schema") + self.assertEqual(list(_rows_by_stat(legacy["df_data_dict"]["all_stats"])), ["dtype"]) + await self._load("sw-schema-inline", stats_tier="schema") + _, inline_schema = await self._connect("sw-schema-inline") + self.assertEqual(self._stats(inline_schema)["status"], "not_computed") + + @tornado.testing.gen_test + async def test_a_failed_stats_run_reports_error_until_the_next_generation(self): + await self._load("sw-error", stats_delivery="deferred") + ws, first = await self._connect("sw-error", caps="stats_update") + gen = self._stats(first)["gen"] + + with patch.object(xorq_loading.XorqServerDataflow, "_get_summary_sd", + side_effect=RuntimeError("stats query failed")): + ws.write_message(_stats_request(gen)) + aborted = await _read_json(ws) + self.assertEqual((aborted["type"], aborted["stats_gen"], aborted["reason"]), ("stats_aborted", gen, "error")) + + # The failure is the session's state for this generation: a client that + # connects now is told so, and a legacy client still gets its frame. + _, caps_frame = await self._connect("sw-error", caps="stats_update") + stats = self._stats(caps_frame) + self.assertEqual((stats["status"], stats["tier"], stats["gen"]), ("error", "schema", gen)) + self.assertIn("reason", stats) + _, legacy = await self._connect("sw-error") + self.assertEqual(legacy["type"], "initial_state") + self.assertEqual(list(_rows_by_stat(legacy["df_data_dict"]["all_stats"])), ["dtype"]) + + with _count_stat_queries() as queries: + ws.write_message(_stats_request(gen)) + self.assertEqual((await _read_json(ws))["reason"], "error") + self.assertEqual(queries, [], "a failed generation must not be retried by every request") + + ws.write_message(_state_change(quick_command_args={"search": ["a"]})) + self._assert_pending(await _read_json(ws), gen + 1) + ws.write_message(_stats_request(gen + 1)) + self.assertEqual((await _read_json(ws))["type"], "stats_update") + + @tornado.testing.gen_test + async def test_stats_request_emits_a_stats_request_span(self): + captured: list = [] + with patch.object(telemetry, "make_http_sink", lambda url, **kw: captured.append): + await self._load("sw-span", stats_delivery="deferred", + telemetry_url="http://companion.invalid/internal/telemetry") + ws, first = await self._connect("sw-span", caps="stats_update") + gen = self._stats(first)["gen"] + ws.write_message(_stats_request(gen - 1)) + await _read_json(ws) + ws.write_message(_stats_request(gen, columns=["price"])) + await _read_json(ws) + + stale, served = [r for r in captured if r["name"] == "stats.request"] + self.assertEqual(served["trace"], "sw-span") + self.assertEqual((served["attrs"]["stats_gen"], served["attrs"]["scope"], served["attrs"]["tier"]), + (gen, "raw", "full")) + self.assertEqual((served["attrs"]["outcome"], served["attrs"]["columns"]), ("update", 1)) + self.assertEqual((stale["attrs"]["outcome"], stale["attrs"]["stats_gen"]), ("stale", gen - 1)) + self.assertIn("stats.complete", [r["name"] for r in captured]) + + @tornado.testing.gen_test + async def test_returning_to_a_completed_state_is_answered_from_the_cache(self): + await self._load("sw-revisit", stats_delivery="deferred") + ws, first = await self._connect("sw-revisit", caps="stats_update") + gen = self._stats(first)["gen"] + ws.write_message(_stats_request(gen)) + original = await _read_json(ws) + ws.write_message(_state_change(quick_command_args={"search": ["a"]})) + self._assert_pending(await _read_json(ws), gen + 1) + ws.write_message(_stats_request(gen + 1)) + self.assertEqual((await _read_json(ws))["type"], "stats_update") + + ws.write_message(_state_change(quick_command_args={})) + self._assert_pending(await _read_json(ws), gen + 2) + with _count_stat_queries() as queries: + ws.write_message(_stats_request(gen + 2)) + again = await _read_json(ws) + self.assertEqual(queries, [], "the first state's stats are in summary_stats_cache") + self.assertEqual(_rows_by_stat(again["payload"]), _rows_by_stat(original["payload"])) + + @tornado.testing.gen_test + async def test_completing_the_stats_keeps_component_config(self): + await self._load("sw-theme", stats_delivery="deferred", component_config={"className": "sw-themed"}) + ws, first = await self._connect("sw-theme", caps="stats_update") + ws.write_message(_stats_request(self._stats(first)["gen"])) + self.assertEqual((await _read_json(ws))["type"], "stats_update") + _, late = await self._connect("sw-theme", caps="stats_update") + dvc = late["df_display_args"]["main"]["df_viewer_config"] + self.assertEqual(dvc["component_config"]["className"], "sw-themed") + + @tornado.testing.gen_test + async def test_a_warm_load_expr_keeps_the_generation(self): + await self._load("sw-warm", stats_delivery="deferred") + _, first = await self._connect("sw-warm", caps="stats_update") + gen = self._stats(first)["gen"] + await self._load("sw-warm", stats_delivery="deferred") + _, again = await self._connect("sw-warm", caps="stats_update") + self.assertEqual(self._stats(again)["gen"], gen, "a warm exit rebuilds nothing, so the generation stands") + + @tornado.testing.gen_test + async def test_a_stats_request_with_no_data_loaded_is_aborted(self): + ws = await tornado.websocket.websocket_connect( + f"ws://localhost:{self.get_http_port()}/ws/sw-no-data?caps=stats_update") + self.clients.append(ws) + ws.write_message(_stats_request(1)) + aborted = await _read_json(ws) + self.assertEqual((aborted["type"], aborted["reason"]), ("stats_aborted", "no_data")) + + @tornado.testing.gen_test + async def test_inline_sessions_send_no_df_meta_stats(self): + """Default behaviour: a session with the default policy sends the + message it always has, and a client reads the absence as complete.""" + await self._load("sw-inline") + _, frame = await self._connect("sw-inline", caps="stats_update") + self.assertNotIn("stats", frame["df_meta"]) + self.assertIn("histogram_bins", _rows_by_stat(frame["df_data_dict"]["all_stats"])) + + class TestReloadExpr(tornado.testing.AsyncHTTPTestCase): def get_app(self): return make_app()