From 80d79d23c291230776ea8728499ffadc6088707e Mon Sep 17 00:00:00 2001 From: Aryan Dhawan Date: Fri, 14 Aug 2026 19:44:05 -0400 Subject: [PATCH 1/6] Add paginated filtered run iteration --- metaflow/client/core.py | 77 ++++++- metaflow/metadata_provider/metadata.py | 33 ++- metaflow/metaflow_config.py | 1 + .../plugins/metadata_providers/service.py | 88 +++++++- test/unit/test_client_run_listing.py | 201 ++++++++++++++++++ 5 files changed, 394 insertions(+), 6 deletions(-) create mode 100644 test/unit/test_client_run_listing.py diff --git a/metaflow/client/core.py b/metaflow/client/core.py index 76b4a472316..ecc00330ab0 100644 --- a/metaflow/client/core.py +++ b/metaflow/client/core.py @@ -7,7 +7,7 @@ from datetime import datetime from tempfile import TemporaryDirectory from io import BytesIO -from itertools import chain +from itertools import chain, islice from typing import ( Any, Dict, @@ -435,6 +435,37 @@ def _filtered_children(self, *tags): if all(tag in child.tags for tag in tags): yield child + def _iter_children(self, query_filters=None, page_size=None, required_tags=()): + """Stream child records through the active metadata provider.""" + namespace_filter = {} + if self._namespace_check and self._current_namespace: + namespace_filter = {"any_tags": self._current_namespace} + + objects = self._metaflow.metadata.iter_objects( + self._NAME, + _CLASSES[self._CHILD_CLASS]._NAME, + namespace_filter, + self._attempt, + *self.path_components, + query_filters=query_filters, + page_size=page_size, + ) + for obj in objects: + child = _CLASSES[self._CHILD_CLASS]( + attempt=self._attempt, + _object=obj, + _parent=self, + _metaflow=self._metaflow, + _namespace_check=self._namespace_check, + _current_namespace=( + self._current_namespace if self._namespace_check else None + ), + ) + if self._iter_filter(child) and all( + tag in child.tags for tag in required_tags + ): + yield child + def _ipython_key_completions_(self): """Returns available options for ipython auto-complete.""" return [child.id for child in self._filtered_children()] @@ -2585,7 +2616,7 @@ def latest_successful_run(self) -> Optional[Run]: if run.successful: return run - def runs(self, *tags: str) -> Iterator[Run]: + def runs(self, *tags: str, **kwargs) -> Iterator[Run]: """ Returns an iterator over all `Run`s of this flow. @@ -2597,13 +2628,53 @@ def runs(self, *tags: str) -> Iterator[Run]: ---------- tags : str Tags to match. + filters : dict, optional + Server-side run filters using the metadata service's ``field:operator`` + grammar, for example ``{"status:eq": "failed"}``. + page_size : int, optional + Number of records requested from the metadata service per page. + max_runs : int, optional + Maximum number of runs to yield. Yields ------ Run `Run` objects in this flow. """ - return self._filtered_children(*tags) + filters = kwargs.pop("filters", None) + page_size = kwargs.pop("page_size", None) + max_runs = kwargs.pop("max_runs", None) + if kwargs: + raise TypeError("Unexpected Flow.runs options: %s" % ", ".join(kwargs)) + if filters is not None and not hasattr(filters, "items"): + raise TypeError("filters must be a mapping") + if max_runs is not None: + if isinstance(max_runs, bool) or not isinstance(max_runs, int): + raise TypeError("max_runs must be an integer") + if max_runs < 0: + raise ValueError("max_runs must be non-negative") + if max_runs == 0: + return iter(()) + + if filters is None and page_size is None and max_runs is None: + return self._filtered_children(*tags) + + query_filters = dict(filters or {}) + server_tags = list(tags) + if server_tags: + existing = query_filters.get("_tags:all") + if existing: + server_tags.insert(0, str(existing)) + query_filters["_tags:all"] = ",".join(server_tags) + + runs = self._iter_children( + query_filters=query_filters, + page_size=page_size, + required_tags=tags, + ) + if max_runs is None: + return runs + return islice(runs, max_runs) def __iter__(self) -> Iterator[Task]: """ diff --git a/metaflow/metadata_provider/metadata.py b/metaflow/metadata_provider/metadata.py index 9075f7e7f4b..7e5339f4c3f 100644 --- a/metaflow/metadata_provider/metadata.py +++ b/metaflow/metadata_provider/metadata.py @@ -6,7 +6,11 @@ from itertools import chain from typing import List -from metaflow.exception import MetaflowInternalError, MetaflowTaggingError +from metaflow.exception import ( + MetaflowException, + MetaflowInternalError, + MetaflowTaggingError, +) from metaflow.tagging_util import validate_tag from metaflow.util import get_username, resolve_identity_as_tuple, is_stringish @@ -439,6 +443,33 @@ def get_object(cls, obj_type, sub_type, filters, attempt, *args): pre_filter, attempt_int ) + @classmethod + def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): + """Iterate over a collection returned by ``get_object``. + + Providers can override this method to stream records without materializing the + complete collection. ``query_filters`` and ``page_size`` are optional provider + hints; the default implementation supports only unfiltered iteration. + """ + query_filters = kwargs.pop("query_filters", None) + page_size = kwargs.pop("page_size", None) + if kwargs: + raise TypeError("Unexpected iterator options: %s" % ", ".join(kwargs)) + if query_filters: + raise MetaflowException( + "Server-side metadata filters are not supported by the %s provider" + % cls.TYPE + ) + if page_size is not None: + if isinstance(page_size, bool) or not isinstance(page_size, int): + raise TypeError("page_size must be an integer") + if page_size <= 0: + raise ValueError("page_size must be positive") + + objects = cls.get_object(obj_type, sub_type, filters, attempt, *args) + for obj in objects or []: + yield obj + @classmethod def mutate_user_tags_for_run( cls, flow_id, run_id, tags_to_remove=None, tags_to_add=None diff --git a/metaflow/metaflow_config.py b/metaflow/metaflow_config.py index bf39015de4e..fbb1f47f8b8 100644 --- a/metaflow/metaflow_config.py +++ b/metaflow/metaflow_config.py @@ -287,6 +287,7 @@ ### SERVICE_URL = from_conf("SERVICE_URL") SERVICE_RETRY_COUNT = from_conf("SERVICE_RETRY_COUNT", 5) +SERVICE_PAGE_SIZE = from_conf("SERVICE_PAGE_SIZE", 100) SERVICE_AUTH_KEY = from_conf("SERVICE_AUTH_KEY") SERVICE_HEADERS = from_conf("SERVICE_HEADERS", {}) if SERVICE_AUTH_KEY is not None: diff --git a/metaflow/plugins/metadata_providers/service.py b/metaflow/plugins/metadata_providers/service.py index c9db2a31822..6062f7fcba2 100644 --- a/metaflow/plugins/metadata_providers/service.py +++ b/metaflow/plugins/metadata_providers/service.py @@ -12,7 +12,12 @@ ) from metaflow.metadata_provider import MetadataProvider from metaflow.metadata_provider.heartbeat import HB_URL_KEY -from metaflow.metaflow_config import SERVICE_HEADERS, SERVICE_RETRY_COUNT, SERVICE_URL +from metaflow.metaflow_config import ( + SERVICE_HEADERS, + SERVICE_PAGE_SIZE, + SERVICE_RETRY_COUNT, + SERVICE_URL, +) from metaflow.sidecar import Message, MessageTypes, Sidecar from urllib.parse import urlencode from metaflow.util import version_parse @@ -56,6 +61,8 @@ class ServiceMetadataProvider(MetadataProvider): _supports_attempt_gets = None _supports_tag_mutation = None + _NEXT_CURSOR_HEADER = "X-Next-Cursor" + def __init__(self, environment, flow, event_logger, monitor): super(ServiceMetadataProvider, self).__init__( environment, flow, event_logger, monitor @@ -304,6 +311,79 @@ def _get_object_internal( return None raise + @classmethod + def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): + """Stream run listings using the service's cursor pagination API.""" + query_filters = kwargs.pop("query_filters", None) or {} + page_size = kwargs.pop("page_size", None) + if kwargs: + raise TypeError("Unexpected iterator options: %s" % ", ".join(kwargs)) + + if (obj_type, sub_type) != ("flow", "run"): + for obj in super(ServiceMetadataProvider, cls).iter_objects( + obj_type, + sub_type, + filters, + attempt, + *args, + query_filters=query_filters, + page_size=page_size, + ): + yield obj + return + if attempt is not None: + raise ValueError("Run listings do not support attempts") + if not args: + raise MetaflowInternalError("A flow name is required to list runs") + + page_size = SERVICE_PAGE_SIZE if page_size is None else page_size + if isinstance(page_size, bool) or not isinstance(page_size, int): + raise TypeError("page_size must be an integer") + if page_size <= 0: + raise ValueError("page_size must be positive") + + query = {} + for key, value in query_filters.items(): + key = str(key) + if key in ("_cursor", "_limit"): + raise ValueError("%s is controlled by the client" % key) + if isinstance(value, (list, tuple, set, frozenset)): + value = ",".join(str(v) for v in value) + else: + value = str(value) + query[key] = value + + tag_filters = [] + for value in (filters or {}).values(): + if isinstance(value, (list, tuple, set, frozenset)): + tag_filters.extend(str(v) for v in value) + else: + tag_filters.append(str(value)) + if tag_filters: + current_tags = query.get("_tags:all") + if current_tags: + tag_filters.insert(0, current_tags) + query["_tags:all"] = ",".join(tag_filters) + + path = "%s/runs" % cls._obj_path(*args[:1]) + cursor = None + seen_cursors = set() + while True: + page_query = dict(query) + page_query["_limit"] = page_size + if cursor is not None: + page_query["_cursor"] = cursor + page_path = "%s?%s" % (path, urlencode(page_query, doseq=True)) + records, headers = cls._request(None, page_path, "GET", return_headers=True) + for record in MetadataProvider._apply_filter(records, filters): + yield record + + next_cursor = headers.get(cls._NEXT_CURSOR_HEADER) + if not next_cursor or next_cursor in seen_cursors: + return + seen_cursors.add(next_cursor) + cursor = next_cursor + def _new_run(self, run_id=None, tags=None, sys_tags=None): # first ensure that the flow exists self._get_or_create("flow") @@ -472,6 +552,7 @@ def _request( data=None, retry_409_path=None, return_raw_resp=False, + return_headers=False, ): if cls.INFO is None: raise MetaflowException( @@ -530,7 +611,10 @@ def _request( if return_raw_resp: return resp, True if resp.status_code < 300: - return resp.json(), True + body = resp.json() + if return_headers: + return body, resp.headers + return body, True elif resp.status_code == 409 and data is not None: # a special case: the post fails due to a conflict # this could occur when we missed a success response diff --git a/test/unit/test_client_run_listing.py b/test/unit/test_client_run_listing.py new file mode 100644 index 00000000000..fb01b9bda1c --- /dev/null +++ b/test/unit/test_client_run_listing.py @@ -0,0 +1,201 @@ +from urllib.parse import parse_qs, urlsplit + +import pytest + +from metaflow.client.core import Flow +from metaflow.exception import MetaflowException +from metaflow.metadata_provider import MetadataProvider +from metaflow.plugins.metadata_providers.service import ServiceMetadataProvider + + +def _record(run_number, tags=None, system_tags=None): + return { + "flow_id": "ExampleFlow", + "run_number": run_number, + "ts_epoch": run_number * 1000, + "tags": tags or [], + "system_tags": system_tags or [], + } + + +def test_service_run_iterator_follows_cursor_and_preserves_filters(monkeypatch): + responses = [ + ( + [ + _record(3, system_tags=["user:aryan"]), + _record(2, system_tags=["user:aryan"]), + ], + {"X-Next-Cursor": "next-page"}, + ), + ([_record(1, system_tags=["user:aryan"])], {}), + ] + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append((path, method, kwargs)) + return responses[len(calls) - 1] + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + records = list( + ServiceMetadataProvider.iter_objects( + "flow", + "run", + {"any_tags": "user:aryan"}, + None, + "ExampleFlow", + query_filters={"status:eq": "failed", "_tags:all": "prod"}, + page_size=2, + ) + ) + + assert [record["run_number"] for record in records] == [3, 2, 1] + first_query = parse_qs(urlsplit(calls[0][0]).query) + second_query = parse_qs(urlsplit(calls[1][0]).query) + assert first_query == { + "_limit": ["2"], + "_tags:all": ["prod,user:aryan"], + "status:eq": ["failed"], + } + assert second_query["_cursor"] == ["next-page"] + assert second_query["status:eq"] == ["failed"] + assert all(method == "GET" for _, method, _ in calls) + assert all(options == {"return_headers": True} for _, _, options in calls) + + +def test_service_run_iterator_stops_on_repeated_cursor(monkeypatch): + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append(path) + return [_record(len(calls))], {"X-Next-Cursor": "same-cursor"} + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + records = list( + ServiceMetadataProvider.iter_objects( + "flow", "run", None, None, "ExampleFlow", page_size=1 + ) + ) + + assert len(records) == 2 + assert len(calls) == 2 + + +@pytest.mark.parametrize("page_size", [0, -1, True, "10"]) +def test_service_run_iterator_rejects_invalid_page_size(page_size): + with pytest.raises((TypeError, ValueError), match="page_size"): + list( + ServiceMetadataProvider.iter_objects( + "flow", "run", None, None, "ExampleFlow", page_size=page_size + ) + ) + + +@pytest.mark.parametrize("reserved", ["_limit", "_cursor"]) +def test_service_run_iterator_owns_pagination_parameters(reserved): + with pytest.raises(ValueError, match=reserved): + list( + ServiceMetadataProvider.iter_objects( + "flow", + "run", + None, + None, + "ExampleFlow", + query_filters={reserved: "value"}, + ) + ) + + +def test_default_provider_iterator_preserves_existing_listing_behavior(): + class FakeProvider(MetadataProvider): + TYPE = "fake" + + @classmethod + def get_object(cls, obj_type, sub_type, filters, attempt, *args): + return [_record(2), _record(1)] + + assert [ + record["run_number"] + for record in FakeProvider.iter_objects("flow", "run", None, None, "Flow") + ] == [2, 1] + with pytest.raises(MetaflowException, match="not supported"): + list( + FakeProvider.iter_objects( + "flow", + "run", + None, + None, + "Flow", + query_filters={"status:eq": "failed"}, + ) + ) + + +def test_service_request_can_return_success_headers(monkeypatch): + class FakeResponse: + status_code = 200 + headers = {"X-Next-Cursor": "cursor"} + + @staticmethod + def json(): + return [_record(1)] + + class FakeSession: + @staticmethod + def get(url, headers): + return FakeResponse() + + monkeypatch.setattr(ServiceMetadataProvider, "_INFO", "http://metadata") + monkeypatch.setattr(ServiceMetadataProvider, "_session", FakeSession()) + + body, headers = ServiceMetadataProvider._request( + None, "/flows/ExampleFlow/runs", "GET", return_headers=True + ) + + assert body == [_record(1)] + assert headers["X-Next-Cursor"] == "cursor" + + +def test_flow_runs_forwards_filters_tags_and_bounds_results(monkeypatch): + flow = object.__new__(Flow) + flow._namespace_check = True + flow._current_namespace = "user:aryan" + captured = {} + + def fake_iter_children(self, query_filters=None, page_size=None, required_tags=()): + captured.update( + query_filters=query_filters, + page_size=page_size, + required_tags=required_tags, + ) + yield from range(4) + + monkeypatch.setattr(Flow, "_iter_children", fake_iter_children) + + runs = list( + flow.runs( + "prod", + filters={"status:eq": "failed"}, + page_size=2, + max_runs=2, + ) + ) + + assert runs == [0, 1] + assert captured == { + "query_filters": {"status:eq": "failed", "_tags:all": "prod"}, + "page_size": 2, + "required_tags": ("prod",), + } + + +def test_flow_runs_zero_limit_avoids_starting_iterator(monkeypatch): + flow = object.__new__(Flow) + + def unexpected_iterator(*args, **kwargs): + raise AssertionError("iterator should not be started") + + monkeypatch.setattr(Flow, "_iter_children", unexpected_iterator) + + assert list(flow.runs(max_runs=0)) == [] From 20583bf82dd10ed964f936fef3b7fccba605c47e Mon Sep 17 00:00:00 2001 From: Aryan Dhawan Date: Fri, 21 Aug 2026 16:19:39 -0400 Subject: [PATCH 2/6] Address review feedback: version-gate pagination, sort newest-first, explicit kwargs - Gate the cursor-paginated path on metadata-service version (>= 2.5.1) with an X-Limit capability-header check; fall back to legacy listing for older services and reject server-side filters when the service cannot honor them. - Generalize pagination to all collection types, not just flow/run listings. - Sort iter_objects results newest-first by ts_epoch so max_runs returns the newest runs regardless of the provider's ordering. - Replace Flow.runs(**kwargs) with explicit keyword-only filters/page_size/ max_runs, and apply positional tags locally so flow.runs("prod") works with local metadata as well as the service. - Extract collection listing into private helpers for readability. - Use mocks instead of fake providers in the run-listing tests; add coverage for legacy fallback, capability checks, and newest-first ordering. --- metaflow/client/core.py | 31 +- metaflow/metadata_provider/metadata.py | 9 +- .../plugins/metadata_providers/service.py | 266 ++++++++++++----- test/unit/test_client_run_listing.py | 267 +++++++++++++++--- 4 files changed, 441 insertions(+), 132 deletions(-) diff --git a/metaflow/client/core.py b/metaflow/client/core.py index ecc00330ab0..a0b0088cdf8 100644 --- a/metaflow/client/core.py +++ b/metaflow/client/core.py @@ -2616,7 +2616,13 @@ def latest_successful_run(self) -> Optional[Run]: if run.successful: return run - def runs(self, *tags: str, **kwargs) -> Iterator[Run]: + def runs( + self, + *tags: str, + filters: Optional[Dict[str, Any]] = None, + page_size: Optional[int] = None, + max_runs: Optional[int] = None, + ) -> Iterator[Run]: """ Returns an iterator over all `Run`s of this flow. @@ -2627,25 +2633,22 @@ def runs(self, *tags: str, **kwargs) -> Iterator[Run]: Parameters ---------- tags : str - Tags to match. + Tags to match. Applied locally after listing, so this works with + local metadata as well as the metadata service. filters : dict, optional Server-side run filters using the metadata service's ``field:operator`` - grammar, for example ``{"status:eq": "failed"}``. + grammar, for example ``{"status:eq": "failed"}``. Requires a metadata + service with pagination and filtering support. page_size : int, optional Number of records requested from the metadata service per page. max_runs : int, optional - Maximum number of runs to yield. + Maximum number of runs to yield, newest first. Yields ------ Run `Run` objects in this flow. """ - filters = kwargs.pop("filters", None) - page_size = kwargs.pop("page_size", None) - max_runs = kwargs.pop("max_runs", None) - if kwargs: - raise TypeError("Unexpected Flow.runs options: %s" % ", ".join(kwargs)) if filters is not None and not hasattr(filters, "items"): raise TypeError("filters must be a mapping") if max_runs is not None: @@ -2659,16 +2662,8 @@ def runs(self, *tags: str, **kwargs) -> Iterator[Run]: if filters is None and page_size is None and max_runs is None: return self._filtered_children(*tags) - query_filters = dict(filters or {}) - server_tags = list(tags) - if server_tags: - existing = query_filters.get("_tags:all") - if existing: - server_tags.insert(0, str(existing)) - query_filters["_tags:all"] = ",".join(server_tags) - runs = self._iter_children( - query_filters=query_filters, + query_filters=dict(filters) if filters else None, page_size=page_size, required_tags=tags, ) diff --git a/metaflow/metadata_provider/metadata.py b/metaflow/metadata_provider/metadata.py index 7e5339f4c3f..2e108c34dab 100644 --- a/metaflow/metadata_provider/metadata.py +++ b/metaflow/metadata_provider/metadata.py @@ -467,7 +467,14 @@ def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): raise ValueError("page_size must be positive") objects = cls.get_object(obj_type, sub_type, filters, attempt, *args) - for obj in objects or []: + if isinstance(objects, dict): + objects = [objects] + objects = sorted( + objects or [], + key=lambda obj: obj.get("ts_epoch") or 0, + reverse=True, + ) + for obj in objects: yield obj @classmethod diff --git a/metaflow/plugins/metadata_providers/service.py b/metaflow/plugins/metadata_providers/service.py index 6062f7fcba2..17de5e6e8d4 100644 --- a/metaflow/plugins/metadata_providers/service.py +++ b/metaflow/plugins/metadata_providers/service.py @@ -12,6 +12,7 @@ ) from metaflow.metadata_provider import MetadataProvider from metaflow.metadata_provider.heartbeat import HB_URL_KEY +from metaflow.metadata_provider.metadata import ObjectOrder from metaflow.metaflow_config import ( SERVICE_HEADERS, SERVICE_PAGE_SIZE, @@ -60,8 +61,11 @@ class ServiceMetadataProvider(MetadataProvider): _supports_attempt_gets = None _supports_tag_mutation = None + _supports_cursor_pagination = None _NEXT_CURSOR_HEADER = "X-Next-Cursor" + _LIMIT_HEADER = "X-Limit" + _MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION = "2.5.1" def __init__(self, environment, flow, event_logger, monitor): super(ServiceMetadataProvider, self).__init__( @@ -259,6 +263,147 @@ def _mutate_user_tags_for_run( time.sleep(0.3 * random.uniform(1.4, 1.6) ** tries) tries += 1 + @classmethod + def _service_supports_cursor_pagination(cls): + if cls._supports_cursor_pagination is None: + version = cls._version(None) + cls._supports_cursor_pagination = version is not None and version_parse( + version + ) >= version_parse(cls._MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION) + return cls._supports_cursor_pagination + + @staticmethod + def _header_value(headers, name): + if not headers: + return None + value = headers.get(name) + if value is not None: + return value + for key, val in headers.items(): + if str(key).lower() == name.lower(): + return val + return None + + @classmethod + def _collection_path(cls, obj_type, obj_order, sub_type, attempt, *args): + if obj_type != "root": + url = cls._obj_path(*args[:obj_order]) + else: + url = "" + if sub_type == "metadata": + url += "/metadata" + elif sub_type == "artifact" and obj_type == "task" and attempt is not None: + url += "/attempt/%s/artifacts" % attempt + else: + url += "/%ss" % sub_type + return url + + @classmethod + def _can_paginate_collection(cls, sub_type, attempt): + return sub_type != "self" and attempt is None + + @classmethod + def _listing_query(cls, query_filters, filters, page_size, cursor): + query = {} + for key, value in (query_filters or {}).items(): + key = str(key) + if key in ("_cursor", "_limit"): + raise ValueError("%s is controlled by the client" % key) + if isinstance(value, (list, tuple, set, frozenset)): + value = ",".join(str(v) for v in value) + else: + value = str(value) + query[key] = value + + tag_filters = [] + for value in (filters or {}).values(): + if isinstance(value, (list, tuple, set, frozenset)): + tag_filters.extend(str(v) for v in value) + else: + tag_filters.append(str(value)) + if tag_filters: + current_tags = query.get("_tags:all") + if current_tags: + tag_filters.insert(0, current_tags) + query["_tags:all"] = ",".join(tag_filters) + + query["_limit"] = page_size + if cursor is not None: + query["_cursor"] = cursor + return query + + @classmethod + def _legacy_get_collection( + cls, obj_type, obj_order, sub_type, filters, attempt, *args + ): + url = cls._collection_path(obj_type, obj_order, sub_type, attempt, *args) + try: + v, _ = cls._request(None, url, "GET") + return MetadataProvider._apply_filter(v, filters) + except ServiceException as ex: + if ex.http_code == 404: + return None + raise + + @classmethod + def _iter_paginated_records( + cls, + obj_type, + obj_order, + sub_type, + filters, + attempt, + *args, + query_filters=None, + page_size=None, + ): + page_size = SERVICE_PAGE_SIZE if page_size is None else page_size + if isinstance(page_size, bool) or not isinstance(page_size, int): + raise TypeError("page_size must be an integer") + if page_size <= 0: + raise ValueError("page_size must be positive") + + path = cls._collection_path(obj_type, obj_order, sub_type, attempt, *args) + cursor = None + seen_cursors = set() + first_page = True + query_filters = query_filters or {} + while True: + page_query = cls._listing_query(query_filters, filters, page_size, cursor) + page_path = "%s?%s" % (path, urlencode(page_query, doseq=True)) + try: + records, headers = cls._request( + None, page_path, "GET", return_headers=True + ) + except ServiceException as ex: + if ex.http_code == 404: + return + raise + if first_page: + first_page = False + if cls._header_value(headers, cls._LIMIT_HEADER) is None: + if query_filters: + raise ServiceException( + "Filtering requires a metadata service with pagination " + "and filtering support (at least version %s). Please " + "upgrade your service." + % cls._MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION + ) + legacy = cls._legacy_get_collection( + obj_type, obj_order, sub_type, filters, attempt, *args + ) + for record in legacy or []: + yield record + return + for record in MetadataProvider._apply_filter(records, filters): + yield record + + next_cursor = cls._header_value(headers, cls._NEXT_CURSOR_HEADER) + if not next_cursor or next_cursor in seen_cursors: + return + seen_cursors.add(next_cursor) + cursor = next_cursor + @classmethod def _get_object_internal( cls, obj_type, obj_order, sub_type, sub_order, filters, attempt, *args @@ -292,34 +437,50 @@ def _get_object_internal( return None raise - # For the other types, we locate all the objects we need to find and return them - if obj_type != "root": - url = ServiceMetadataProvider._obj_path(*args[:obj_order]) - else: - url = "" - if sub_type == "metadata": - url += "/metadata" - elif sub_type == "artifact" and obj_type == "task" and attempt is not None: - url += "/attempt/%s/artifacts" % attempt - else: - url += "/%ss" % sub_type - try: - v, _ = cls._request(None, url, "GET") - return MetadataProvider._apply_filter(v, filters) - except ServiceException as ex: - if ex.http_code == 404: - return None - raise + # Newer services can stream collections; keep returning a list here so + # get_object callers are unchanged. Older services stay on the bulk GET. + if ( + cls._can_paginate_collection(sub_type, attempt) + and cls._service_supports_cursor_pagination() + ): + try: + return list( + cls._iter_paginated_records( + obj_type, obj_order, sub_type, filters, attempt, *args + ) + ) + except ServiceException as ex: + if ex.http_code == 404: + return None + raise + + return cls._legacy_get_collection( + obj_type, obj_order, sub_type, filters, attempt, *args + ) @classmethod def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): - """Stream run listings using the service's cursor pagination API.""" + """Stream collection listings using cursor pagination when the service supports it.""" query_filters = kwargs.pop("query_filters", None) or {} page_size = kwargs.pop("page_size", None) if kwargs: raise TypeError("Unexpected iterator options: %s" % ", ".join(kwargs)) - if (obj_type, sub_type) != ("flow", "run"): + if not cls._service_supports_cursor_pagination(): + if query_filters: + raise ServiceException( + "Filtering requires a metadata service with pagination " + "and filtering support (at least version %s). Please " + "upgrade your service." + % cls._MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION + ) + for obj in super(ServiceMetadataProvider, cls).iter_objects( + obj_type, sub_type, filters, attempt, *args, page_size=page_size + ): + yield obj + return + + if not cls._can_paginate_collection(sub_type, attempt): for obj in super(ServiceMetadataProvider, cls).iter_objects( obj_type, sub_type, @@ -331,58 +492,21 @@ def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): ): yield obj return - if attempt is not None: - raise ValueError("Run listings do not support attempts") - if not args: - raise MetaflowInternalError("A flow name is required to list runs") - - page_size = SERVICE_PAGE_SIZE if page_size is None else page_size - if isinstance(page_size, bool) or not isinstance(page_size, int): - raise TypeError("page_size must be an integer") - if page_size <= 0: - raise ValueError("page_size must be positive") - - query = {} - for key, value in query_filters.items(): - key = str(key) - if key in ("_cursor", "_limit"): - raise ValueError("%s is controlled by the client" % key) - if isinstance(value, (list, tuple, set, frozenset)): - value = ",".join(str(v) for v in value) - else: - value = str(value) - query[key] = value - tag_filters = [] - for value in (filters or {}).values(): - if isinstance(value, (list, tuple, set, frozenset)): - tag_filters.extend(str(v) for v in value) - else: - tag_filters.append(str(value)) - if tag_filters: - current_tags = query.get("_tags:all") - if current_tags: - tag_filters.insert(0, current_tags) - query["_tags:all"] = ",".join(tag_filters) - - path = "%s/runs" % cls._obj_path(*args[:1]) - cursor = None - seen_cursors = set() - while True: - page_query = dict(query) - page_query["_limit"] = page_size - if cursor is not None: - page_query["_cursor"] = cursor - page_path = "%s?%s" % (path, urlencode(page_query, doseq=True)) - records, headers = cls._request(None, page_path, "GET", return_headers=True) - for record in MetadataProvider._apply_filter(records, filters): - yield record - - next_cursor = headers.get(cls._NEXT_CURSOR_HEADER) - if not next_cursor or next_cursor in seen_cursors: - return - seen_cursors.add(next_cursor) - cursor = next_cursor + obj_order = ObjectOrder.type_to_order(obj_type) + if obj_order is None: + raise MetaflowInternalError("Cannot find type %s" % obj_type) + for obj in cls._iter_paginated_records( + obj_type, + obj_order, + sub_type, + filters, + attempt, + *args, + query_filters=query_filters, + page_size=page_size, + ): + yield obj def _new_run(self, run_id=None, tags=None, sys_tags=None): # first ensure that the flow exists diff --git a/test/unit/test_client_run_listing.py b/test/unit/test_client_run_listing.py index fb01b9bda1c..04f61cc1862 100644 --- a/test/unit/test_client_run_listing.py +++ b/test/unit/test_client_run_listing.py @@ -1,11 +1,19 @@ +from unittest.mock import Mock from urllib.parse import parse_qs, urlsplit import pytest from metaflow.client.core import Flow from metaflow.exception import MetaflowException -from metaflow.metadata_provider import MetadataProvider -from metaflow.plugins.metadata_providers.service import ServiceMetadataProvider +from metaflow.plugins.metadata_providers.local import LocalMetadataProvider +from metaflow.plugins.metadata_providers.service import ( + ServiceException, + ServiceMetadataProvider, +) + +PAGINATING_SERVICE_VERSION = ( + ServiceMetadataProvider._MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION +) def _record(run_number, tags=None, system_tags=None): @@ -18,16 +26,34 @@ def _record(run_number, tags=None, system_tags=None): } +@pytest.fixture(autouse=True) +def reset_service_capability_cache(): + ServiceMetadataProvider._supports_cursor_pagination = None + yield + ServiceMetadataProvider._supports_cursor_pagination = None + + +def _paginating_version(cls, monitor): + return PAGINATING_SERVICE_VERSION + + +def _legacy_version(cls, monitor): + return "2.4.0" + + def test_service_run_iterator_follows_cursor_and_preserves_filters(monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) responses = [ ( [ _record(3, system_tags=["user:aryan"]), _record(2, system_tags=["user:aryan"]), ], - {"X-Next-Cursor": "next-page"}, + {"X-Next-Cursor": "next-page", "X-Limit": "2"}, ), - ([_record(1, system_tags=["user:aryan"])], {}), + ([_record(1, system_tags=["user:aryan"])], {"X-Limit": "2"}), ] calls = [] @@ -63,12 +89,37 @@ def fake_request(cls, monitor, path, method, **kwargs): assert all(options == {"return_headers": True} for _, _, options in calls) +def test_service_iterator_paginates_all_collection_types(monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append(path) + return [_record(1)], {"X-Limit": "1"} + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + list( + ServiceMetadataProvider.iter_objects( + "run", "step", None, None, "ExampleFlow", "12", page_size=1 + ) + ) + + assert calls[0].startswith("/flows/ExampleFlow/runs/12/steps?") + assert parse_qs(urlsplit(calls[0]).query) == {"_limit": ["1"]} + + def test_service_run_iterator_stops_on_repeated_cursor(monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) calls = [] def fake_request(cls, monitor, path, method, **kwargs): calls.append(path) - return [_record(len(calls))], {"X-Next-Cursor": "same-cursor"} + return [_record(len(calls))], {"X-Next-Cursor": "same-cursor", "X-Limit": "1"} monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) @@ -82,8 +133,110 @@ def fake_request(cls, monitor, path, method, **kwargs): assert len(calls) == 2 +def test_old_service_uses_legacy_listing_without_query_params(monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_legacy_version) + ) + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append(path) + return [_record(2), _record(1)], {} + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + records = list( + ServiceMetadataProvider.iter_objects("flow", "run", None, None, "ExampleFlow") + ) + + assert [record["run_number"] for record in records] == [2, 1] + assert calls == ["/flows/ExampleFlow/runs"] + assert "?" not in calls[0] + + +def test_old_service_rejects_server_filters_without_listing(monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_legacy_version) + ) + + def unexpected_request(*args, **kwargs): + raise AssertionError( + "legacy services must not receive filtered listing requests" + ) + + monkeypatch.setattr( + ServiceMetadataProvider, "_request", classmethod(unexpected_request) + ) + + with pytest.raises(ServiceException, match="Filtering requires"): + list( + ServiceMetadataProvider.iter_objects( + "flow", + "run", + None, + None, + "ExampleFlow", + query_filters={"status:eq": "failed"}, + ) + ) + + +def test_new_service_without_limit_header_does_not_yield_unfiltered_records( + monkeypatch, +): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + yielded = [] + + def fake_request(cls, monitor, path, method, **kwargs): + return [_record(1)], {} + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + with pytest.raises(ServiceException, match="Filtering requires"): + for record in ServiceMetadataProvider.iter_objects( + "flow", + "run", + None, + None, + "ExampleFlow", + query_filters={"status:eq": "failed"}, + ): + yielded.append(record) + + assert yielded == [] + + +def test_new_service_without_limit_header_falls_back_to_legacy_listing(monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append((path, kwargs)) + if "return_headers" in kwargs: + return [_record(9)], {} + return [_record(2), _record(1)], True + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + records = list( + ServiceMetadataProvider.iter_objects("flow", "run", None, None, "ExampleFlow") + ) + + assert [record["run_number"] for record in records] == [2, 1] + assert calls[0][1] == {"return_headers": True} + assert "?" in calls[0][0] + assert calls[1][0] == "/flows/ExampleFlow/runs" + + @pytest.mark.parametrize("page_size", [0, -1, True, "10"]) -def test_service_run_iterator_rejects_invalid_page_size(page_size): +def test_service_run_iterator_rejects_invalid_page_size(page_size, monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) with pytest.raises((TypeError, ValueError), match="page_size"): list( ServiceMetadataProvider.iter_objects( @@ -93,7 +246,10 @@ def test_service_run_iterator_rejects_invalid_page_size(page_size): @pytest.mark.parametrize("reserved", ["_limit", "_cursor"]) -def test_service_run_iterator_owns_pagination_parameters(reserved): +def test_service_run_iterator_owns_pagination_parameters(reserved, monkeypatch): + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) with pytest.raises(ValueError, match=reserved): list( ServiceMetadataProvider.iter_objects( @@ -107,21 +263,24 @@ def test_service_run_iterator_owns_pagination_parameters(reserved): ) -def test_default_provider_iterator_preserves_existing_listing_behavior(): - class FakeProvider(MetadataProvider): - TYPE = "fake" - - @classmethod - def get_object(cls, obj_type, sub_type, filters, attempt, *args): - return [_record(2), _record(1)] +def test_default_provider_iterator_sorts_newest_first_and_rejects_server_filters( + monkeypatch, +): + monkeypatch.setattr( + LocalMetadataProvider, + "get_object", + classmethod(lambda cls, *args, **kwargs: [_record(1), _record(2)]), + ) assert [ record["run_number"] - for record in FakeProvider.iter_objects("flow", "run", None, None, "Flow") + for record in LocalMetadataProvider.iter_objects( + "flow", "run", None, None, "Flow" + ) ] == [2, 1] with pytest.raises(MetaflowException, match="not supported"): list( - FakeProvider.iter_objects( + LocalMetadataProvider.iter_objects( "flow", "run", None, @@ -133,21 +292,16 @@ def get_object(cls, obj_type, sub_type, filters, attempt, *args): def test_service_request_can_return_success_headers(monkeypatch): - class FakeResponse: - status_code = 200 - headers = {"X-Next-Cursor": "cursor"} + response = Mock() + response.status_code = 200 + response.headers = {"X-Next-Cursor": "cursor"} + response.json.return_value = [_record(1)] - @staticmethod - def json(): - return [_record(1)] - - class FakeSession: - @staticmethod - def get(url, headers): - return FakeResponse() + session = Mock() + session.get.return_value = response monkeypatch.setattr(ServiceMetadataProvider, "_INFO", "http://metadata") - monkeypatch.setattr(ServiceMetadataProvider, "_session", FakeSession()) + monkeypatch.setattr(ServiceMetadataProvider, "_session", session) body, headers = ServiceMetadataProvider._request( None, "/flows/ExampleFlow/runs", "GET", return_headers=True @@ -157,13 +311,10 @@ def get(url, headers): assert headers["X-Next-Cursor"] == "cursor" -def test_flow_runs_forwards_filters_tags_and_bounds_results(monkeypatch): - flow = object.__new__(Flow) - flow._namespace_check = True - flow._current_namespace = "user:aryan" +def test_flow_runs_forwards_filters_and_bounds_results(): captured = {} - def fake_iter_children(self, query_filters=None, page_size=None, required_tags=()): + def fake_iter_children(query_filters=None, page_size=None, required_tags=()): captured.update( query_filters=query_filters, page_size=page_size, @@ -171,10 +322,11 @@ def fake_iter_children(self, query_filters=None, page_size=None, required_tags=( ) yield from range(4) - monkeypatch.setattr(Flow, "_iter_children", fake_iter_children) + flow = Mock() + flow._iter_children = fake_iter_children runs = list( - flow.runs( + Flow.runs.__get__(flow, Flow)( "prod", filters={"status:eq": "failed"}, page_size=2, @@ -184,18 +336,49 @@ def fake_iter_children(self, query_filters=None, page_size=None, required_tags=( assert runs == [0, 1] assert captured == { - "query_filters": {"status:eq": "failed", "_tags:all": "prod"}, + "query_filters": {"status:eq": "failed"}, "page_size": 2, "required_tags": ("prod",), } -def test_flow_runs_zero_limit_avoids_starting_iterator(monkeypatch): - flow = object.__new__(Flow) +def test_flow_runs_tag_and_max_runs_work_without_server_filters(): + captured = {} + + def fake_iter_children(query_filters=None, page_size=None, required_tags=()): + captured.update( + query_filters=query_filters, + page_size=page_size, + required_tags=required_tags, + ) + yield from ("newest", "older") + + flow = Mock() + flow._iter_children = fake_iter_children + + assert list(Flow.runs.__get__(flow, Flow)("prod", max_runs=1)) == ["newest"] + assert captured == { + "query_filters": None, + "page_size": None, + "required_tags": ("prod",), + } + + +def test_flow_runs_max_runs_returns_newest_from_oldest_first_provider(monkeypatch): + monkeypatch.setattr( + LocalMetadataProvider, + "get_object", + classmethod(lambda cls, *args, **kwargs: [_record(1), _record(2)]), + ) + + records = list( + LocalMetadataProvider.iter_objects("flow", "run", None, None, "Flow") + ) + assert [record["run_number"] for record in records[:1]] == [2] - def unexpected_iterator(*args, **kwargs): - raise AssertionError("iterator should not be started") - monkeypatch.setattr(Flow, "_iter_children", unexpected_iterator) +def test_flow_runs_zero_limit_avoids_starting_iterator(): + flow = Mock() + flow._iter_children.side_effect = AssertionError("iterator should not be started") - assert list(flow.runs(max_runs=0)) == [] + assert list(Flow.runs.__get__(flow, Flow)(max_runs=0)) == [] From 15181fabeb1476777d9296085c43182431bf4581 Mon Sep 17 00:00:00 2001 From: Aryan Dhawan Date: Sun, 23 Aug 2026 01:38:24 -0400 Subject: [PATCH 3/6] Document pagination trade-offs raised in review Follow-up on the paginated listing path: - Explain why _get_object_internal materializes pages into a list: get_object's object-or-list contract must stay stable, the goal here is to relieve server-side pressure, and callers needing to stream large collections go through iter_objects()/_iter_paginated_records (e.g. Flow.runs()). - Document that result ordering is implicitly newest-first (descending ts_epoch) because there is no _order query param yet, in both the service iterator and the base MetadataProvider.iter_objects sort, so the two paths stay consistent. --- metaflow/metadata_provider/metadata.py | 5 +++++ .../plugins/metadata_providers/service.py | 21 +++++++++++++++++-- 2 files changed, 24 insertions(+), 2 deletions(-) diff --git a/metaflow/metadata_provider/metadata.py b/metaflow/metadata_provider/metadata.py index 2e108c34dab..634db34a723 100644 --- a/metaflow/metadata_provider/metadata.py +++ b/metaflow/metadata_provider/metadata.py @@ -469,6 +469,11 @@ def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): objects = cls.get_object(obj_type, sub_type, filters, attempt, *args) if isinstance(objects, dict): objects = [objects] + # Yield newest-first (descending ts_epoch). This mirrors the order the + # metadata service returns for paginated listings, so callers see a + # consistent order whether records came from the service or a legacy / + # local provider. ts_epoch may be missing on some records, so fall back + # to 0 to keep the sort total. objects = sorted( objects or [], key=lambda obj: obj.get("ts_epoch") or 0, diff --git a/metaflow/plugins/metadata_providers/service.py b/metaflow/plugins/metadata_providers/service.py index 17de5e6e8d4..d378b65bb30 100644 --- a/metaflow/plugins/metadata_providers/service.py +++ b/metaflow/plugins/metadata_providers/service.py @@ -368,6 +368,11 @@ def _iter_paginated_records( seen_cursors = set() first_page = True query_filters = query_filters or {} + # Result ordering is implicit: we send no _order query param (none exists + # yet), so pages arrive in the service's default order, which is + # newest-first (descending ts_epoch). MetadataProvider.iter_objects + # applies the same descending ts_epoch sort for legacy/local providers, + # so ordering stays consistent regardless of which path served the records. while True: page_query = cls._listing_query(query_filters, filters, page_size, cursor) page_path = "%s?%s" % (path, urlencode(page_query, doseq=True)) @@ -437,8 +442,20 @@ def _get_object_internal( return None raise - # Newer services can stream collections; keep returning a list here so - # get_object callers are unchanged. Older services stay on the bulk GET. + # Newer services can stream collections, but get_object's contract is a + # concrete object-or-list, so we materialize the pages here to keep every + # existing caller unchanged. That trades some client-side memory for + # backward compatibility -- the goal of pagination here is to relieve + # server-side pressure, not client-side. Callers that need to stream + # large collections without materializing should go through + # iter_objects() / _iter_paginated_records (e.g. Flow.runs()), which + # yield page by page. Older services stay on the bulk GET. + # + # Possible follow-up (left for a future contributor): let the generic + # get_object / MetaflowObject.__iter__ path stream for new services too, + # so plain iteration is also memory-bounded. That requires changing + # get_object's object-or-list return contract, so it is intentionally + # out of scope here. if ( cls._can_paginate_collection(sub_type, attempt) and cls._service_supports_cursor_pagination() From 1e9763251db2b02d715fc4f5292842bcddc31a8a Mon Sep 17 00:00:00 2001 From: Aryan Dhawan Date: Sun, 23 Aug 2026 17:28:46 -0400 Subject: [PATCH 4/6] Address review: gate on 2.6.0, fix 404 semantics, share access validation - Bump _MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION to 2.6.0 -- the release that ships pagination + filtering (per maintainer). - The paginated get_object path returned [] where the legacy path returned None for a missing (404) collection. Add a raise_on_missing flag threaded through _iter_paginated_records and _legacy_get_collection so get_object keeps legacy's atomic contract: a 404 at any point in the listing (first page, mid-pagination, or the no-X-Limit legacy fallback) resolves to None, never an empty or silently truncated list. Streaming iter_objects is unchanged: a 404 just ends the stream. - Lift get_object's obj_type/sub_type validation guards into a shared _validate_object_query helper and call it from the paginated listing path, so streamed access rejects the same nonsensical combinations as materialized access. - Cover all of the above with tests. --- metaflow/metadata_provider/metadata.py | 53 ++++--- .../plugins/metadata_providers/service.py | 52 +++++-- test/unit/test_client_run_listing.py | 132 +++++++++++++++++- 3 files changed, 206 insertions(+), 31 deletions(-) diff --git a/metaflow/metadata_provider/metadata.py b/metaflow/metadata_provider/metadata.py index 634db34a723..8a1e7276c34 100644 --- a/metaflow/metadata_provider/metadata.py +++ b/metaflow/metadata_provider/metadata.py @@ -352,6 +352,37 @@ def add_sticky_tags(self, tags=None, sys_tags=None): if sys_tags: self.sticky_sys_tags.update(sys_tags) + @classmethod + def _validate_object_query(cls, obj_type, sub_type): + """Reject nonsensical obj_type/sub_type combinations. + + Shared by get_object and the streaming listing paths so every access, + materialized or paginated, enforces the same rules. Returns the + (type_order, sub_order) pair for callers that need it. + """ + type_order = ObjectOrder.type_to_order(obj_type) + sub_order = ObjectOrder.type_to_order(sub_type) + + if type_order is None: + raise MetaflowInternalError(msg="Cannot find type %s" % obj_type) + if type_order >= ObjectOrder.type_to_order("metadata"): + raise MetaflowInternalError(msg="Type %s is not allowed" % obj_type) + + if sub_order is None: + raise MetaflowInternalError(msg="Cannot find subtype %s" % sub_type) + + if type_order >= sub_order: + raise MetaflowInternalError( + msg="Subtype %s not allowed for %s" % (sub_type, obj_type) + ) + + # Metadata is always only at the task level + if sub_type == "metadata" and obj_type != "task": + raise MetaflowInternalError( + msg="Metadata can only be retrieved at the task level" + ) + return type_order, sub_order + @classmethod def get_object(cls, obj_type, sub_type, filters, attempt, *args): """Returns the requested object depending on obj_type and sub_type @@ -401,27 +432,7 @@ def get_object(cls, obj_type, sub_type, filters, attempt, *args): object or list : Depending on the call, the type of object return varies """ - type_order = ObjectOrder.type_to_order(obj_type) - sub_order = ObjectOrder.type_to_order(sub_type) - - if type_order is None: - raise MetaflowInternalError(msg="Cannot find type %s" % obj_type) - if type_order >= ObjectOrder.type_to_order("metadata"): - raise MetaflowInternalError(msg="Type %s is not allowed" % obj_type) - - if sub_order is None: - raise MetaflowInternalError(msg="Cannot find subtype %s" % sub_type) - - if type_order >= sub_order: - raise MetaflowInternalError( - msg="Subtype %s not allowed for %s" % (sub_type, obj_type) - ) - - # Metadata is always only at the task level - if sub_type == "metadata" and obj_type != "task": - raise MetaflowInternalError( - msg="Metadata can only be retrieved at the task level" - ) + type_order, sub_order = cls._validate_object_query(obj_type, sub_type) if attempt is not None: try: diff --git a/metaflow/plugins/metadata_providers/service.py b/metaflow/plugins/metadata_providers/service.py index d378b65bb30..2f754dcd359 100644 --- a/metaflow/plugins/metadata_providers/service.py +++ b/metaflow/plugins/metadata_providers/service.py @@ -12,7 +12,6 @@ ) from metaflow.metadata_provider import MetadataProvider from metaflow.metadata_provider.heartbeat import HB_URL_KEY -from metaflow.metadata_provider.metadata import ObjectOrder from metaflow.metaflow_config import ( SERVICE_HEADERS, SERVICE_PAGE_SIZE, @@ -65,7 +64,7 @@ class ServiceMetadataProvider(MetadataProvider): _NEXT_CURSOR_HEADER = "X-Next-Cursor" _LIMIT_HEADER = "X-Limit" - _MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION = "2.5.1" + _MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION = "2.6.0" def __init__(self, environment, flow, event_logger, monitor): super(ServiceMetadataProvider, self).__init__( @@ -334,7 +333,14 @@ def _listing_query(cls, query_filters, filters, page_size, cursor): @classmethod def _legacy_get_collection( - cls, obj_type, obj_order, sub_type, filters, attempt, *args + cls, + obj_type, + obj_order, + sub_type, + filters, + attempt, + *args, + raise_on_missing=False, ): url = cls._collection_path(obj_type, obj_order, sub_type, attempt, *args) try: @@ -342,6 +348,8 @@ def _legacy_get_collection( return MetadataProvider._apply_filter(v, filters) except ServiceException as ex: if ex.http_code == 404: + if raise_on_missing: + raise return None raise @@ -356,6 +364,7 @@ def _iter_paginated_records( *args, query_filters=None, page_size=None, + raise_on_missing=False, ): page_size = SERVICE_PAGE_SIZE if page_size is None else page_size if isinstance(page_size, bool) or not isinstance(page_size, int): @@ -382,6 +391,14 @@ def _iter_paginated_records( ) except ServiceException as ex: if ex.http_code == 404: + # A missing collection. When the caller needs to tell "not + # found" (None) apart from "empty" ([]) -- e.g. get_object -- + # re-raise no matter which page 404'd, so the whole listing + # resolves to None and keeps legacy's atomic full-result-or- + # None contract (never a silently truncated partial list). + # Streaming callers instead just end the stream here. + if raise_on_missing: + raise return raise if first_page: @@ -395,7 +412,13 @@ def _iter_paginated_records( % cls._MIN_SERVICE_VERSION_WITH_CURSOR_PAGINATION ) legacy = cls._legacy_get_collection( - obj_type, obj_order, sub_type, filters, attempt, *args + obj_type, + obj_order, + sub_type, + filters, + attempt, + *args, + raise_on_missing=raise_on_missing, ) for record in legacy or []: yield record @@ -463,7 +486,13 @@ def _get_object_internal( try: return list( cls._iter_paginated_records( - obj_type, obj_order, sub_type, filters, attempt, *args + obj_type, + obj_order, + sub_type, + filters, + attempt, + *args, + raise_on_missing=True, ) ) except ServiceException as ex: @@ -477,7 +506,11 @@ def _get_object_internal( @classmethod def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): - """Stream collection listings using cursor pagination when the service supports it.""" + """Stream collection listings using cursor pagination when the service supports it. + + Like the base implementation, this is a generator: argument validation + and capability errors surface when iteration starts, not at call time. + """ query_filters = kwargs.pop("query_filters", None) or {} page_size = kwargs.pop("page_size", None) if kwargs: @@ -510,9 +543,10 @@ def iter_objects(cls, obj_type, sub_type, filters, attempt, *args, **kwargs): yield obj return - obj_order = ObjectOrder.type_to_order(obj_type) - if obj_order is None: - raise MetaflowInternalError("Cannot find type %s" % obj_type) + # Same access-validation as get_object -- the paginated path must not + # accept nonsensical obj_type/sub_type combinations the materialized + # path would reject. + obj_order, _ = cls._validate_object_query(obj_type, sub_type) for obj in cls._iter_paginated_records( obj_type, obj_order, diff --git a/test/unit/test_client_run_listing.py b/test/unit/test_client_run_listing.py index 04f61cc1862..d65cbf44ac0 100644 --- a/test/unit/test_client_run_listing.py +++ b/test/unit/test_client_run_listing.py @@ -4,7 +4,7 @@ import pytest from metaflow.client.core import Flow -from metaflow.exception import MetaflowException +from metaflow.exception import MetaflowException, MetaflowInternalError from metaflow.plugins.metadata_providers.local import LocalMetadataProvider from metaflow.plugins.metadata_providers.service import ( ServiceException, @@ -382,3 +382,133 @@ def test_flow_runs_zero_limit_avoids_starting_iterator(): flow._iter_children.side_effect = AssertionError("iterator should not be started") assert list(Flow.runs.__get__(flow, Flow)(max_runs=0)) == [] + + +def test_get_object_internal_returns_none_not_empty_on_404(monkeypatch): + """A missing (404) collection must return None like the legacy path, not []. + + The paginated iterator swallows a first-page 404, so without care + list(...) would yield [] and mask "not found" as "empty". + """ + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + + def fake_request(cls, monitor, path, method, **kwargs): + raise ServiceException("collection not found", http_code=404) + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + result = ServiceMetadataProvider._get_object_internal( + "flow", 1, "run", 2, None, None, "ExampleFlow" + ) + assert result is None + + +def test_iter_objects_yields_empty_on_404(monkeypatch): + """Streaming a missing collection yields nothing (no 404 leaking to callers).""" + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + + def fake_request(cls, monitor, path, method, **kwargs): + raise ServiceException("collection not found", http_code=404) + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + assert ( + list( + ServiceMetadataProvider.iter_objects( + "flow", "run", None, None, "ExampleFlow" + ) + ) + == [] + ) + + +def test_get_object_internal_mid_page_404_returns_none(monkeypatch): + """get_object keeps legacy's atomic contract: a 404 at ANY point in the + listing resolves to None, never a silently truncated partial list.""" + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append(path) + if len(calls) == 1: + return [_record(2)], {"X-Next-Cursor": "p2", "X-Limit": "1"} + raise ServiceException("collection deleted mid-scan", http_code=404) + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + result = ServiceMetadataProvider._get_object_internal( + "flow", 1, "run", 2, None, None, "ExampleFlow" + ) + assert result is None + assert len(calls) == 2 + + +def test_iter_objects_mid_page_404_ends_stream_with_first_page(monkeypatch): + """Streaming keeps what was fetched: a mid-pagination 404 just ends the + stream after the records already yielded.""" + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append(path) + if len(calls) == 1: + return [_record(2)], {"X-Next-Cursor": "p2", "X-Limit": "1"} + raise ServiceException("collection deleted mid-scan", http_code=404) + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + records = list( + ServiceMetadataProvider.iter_objects("flow", "run", None, None, "ExampleFlow") + ) + assert [record["run_number"] for record in records] == [2] + assert len(calls) == 2 + + +def test_get_object_internal_returns_none_when_legacy_fallback_404s(monkeypatch): + """The no-X-Limit fallback must not mask 'not found' as 'empty': if the + legacy GET 404s, get_object returns None, not [].""" + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + calls = [] + + def fake_request(cls, monitor, path, method, **kwargs): + calls.append(kwargs) + if "return_headers" in kwargs: + # Paginated probe: 200 but no X-Limit -> triggers legacy fallback. + return [_record(9)], {} + raise ServiceException("collection not found", http_code=404) + + monkeypatch.setattr(ServiceMetadataProvider, "_request", classmethod(fake_request)) + + result = ServiceMetadataProvider._get_object_internal( + "flow", 1, "run", 2, None, None, "ExampleFlow" + ) + assert result is None + assert len(calls) == 2 + + +def test_service_paginated_iterator_validates_obj_subtype(monkeypatch): + """The paginated path enforces the same obj/sub_type guards as get_object.""" + monkeypatch.setattr( + ServiceMetadataProvider, "_version", classmethod(_paginating_version) + ) + + def unexpected_request(*args, **kwargs): + raise AssertionError("validation must fail before any request") + + monkeypatch.setattr( + ServiceMetadataProvider, "_request", classmethod(unexpected_request) + ) + + # 'flow' is not slotted below 'run' -> nonsensical; must raise before any request. + with pytest.raises(MetaflowInternalError, match="not allowed"): + list(ServiceMetadataProvider.iter_objects("run", "flow", None, None, "F", "1")) From 210eec8b521ff4c453b64b0bd946f2fbf8d8c986 Mon Sep 17 00:00:00 2001 From: Aryan Dhawan Date: Fri, 21 Aug 2026 16:38:19 -0400 Subject: [PATCH 5/6] Fold failure investigation into the core client Replaces the separate metaflow/agent module with core-client accessors, per the review that record iteration and normal client operations belong in core: - Flow.failed_runs(*, since=None, max_runs=None): iterate failed runs newest first via runs(filters={"status:eq": "failed"}); raises on local/older services like any other server-filtered listing. - Run.failed_task: the first unsuccessful task in the run (newest step first). - Task.failure_summary -> FailureSummary: normalized exception type/message/ stacktrace/attempt, or None when the task has no exception. Read errors propagate rather than being swallowed into diagnostic strings. Drops the injectable resolver/fetcher and the TaskRef indirection (uses real Task objects). Adds test/unit/test_client_failure_investigation.py. --- metaflow/client/core.py | 140 +++++++++++++++ .../unit/test_client_failure_investigation.py | 169 ++++++++++++++++++ 2 files changed, 309 insertions(+) create mode 100644 test/unit/test_client_failure_investigation.py diff --git a/metaflow/client/core.py b/metaflow/client/core.py index a0b0088cdf8..bb745cc623f 100644 --- a/metaflow/client/core.py +++ b/metaflow/client/core.py @@ -4,6 +4,8 @@ import os import tarfile from collections import namedtuple +from collections.abc import Mapping +from dataclasses import dataclass from datetime import datetime from tempfile import TemporaryDirectory from io import BytesIO @@ -61,6 +63,48 @@ current_metadata = False +@dataclass(frozen=True) +class FailureSummary: + """ + Normalized description of the exception that failed a `Task`. + + Returned by `Task.failure_summary`. Any field may be None when the stored + exception artifact does not carry that particular detail. + """ + + exception_type: Optional[str] + message: Optional[str] + stacktrace: Optional[str] + attempt: Optional[int] + + +def _normalize_exception(data: Any) -> Dict[str, Optional[str]]: + """ + Flatten a task's ``_exception`` artifact into type/message/stacktrace strings. + + Handles both mapping-shaped exception records and exception-like objects, + falling back to ``str(data)`` for the message so there is always something + human (or agent) readable. + """ + if isinstance(data, Mapping): + exception_type = data.get("type") + message = data.get("message") or data.get("exception") + stacktrace = data.get("stacktrace") + else: + exception_type = getattr(data, "type", None) + message = getattr(data, "message", None) or getattr(data, "exception", None) + stacktrace = getattr(data, "stacktrace", None) + if exception_type is None: + exception_type = "%s.%s" % (type(data).__module__, type(data).__name__) + if message is None: + message = str(data) + return { + "type": str(exception_type) if exception_type is not None else None, + "message": str(message), + "stacktrace": str(stacktrace) if stacktrace is not None else None, + } + + def metadata(ms: str) -> str: """ Switch Metadata provider. @@ -1634,6 +1678,36 @@ def exception(self) -> Optional[Any]: except KeyError: return None + @property + def failure_summary(self) -> Optional[FailureSummary]: + """ + Returns a normalized summary of the exception that failed this task. + + This is a convenience over `exception` that flattens the stored + exception artifact into ``exception_type``, ``message``, ``stacktrace``, + and ``attempt`` fields, so callers (including agents summarizing a + failure) do not need to know how the exception was serialized. + + Returns None when the task recorded no exception (for example a task + that succeeded or has not failed). Errors while reading the underlying + exception or attempt are allowed to propagate. + + Returns + ------- + FailureSummary, optional + Normalized failure details, or None if the task has no exception. + """ + exception = self.exception + if exception is None: + return None + normalized = _normalize_exception(exception) + return FailureSummary( + exception_type=normalized["type"], + message=normalized["message"], + stacktrace=normalized["stacktrace"], + attempt=self.current_attempt, + ) + @property def finished_at(self) -> Optional[datetime]: """ @@ -2336,6 +2410,39 @@ def successful(self) -> bool: else: return False + @property + def failed_task(self) -> Optional[Task]: + """ + Returns the first failed (unsuccessful) task in this run, if any. + + Steps and their tasks are scanned in iteration order, which is + newest-created first (see `MetaflowObject.__iter__`). The first task + whose `successful` is False is returned; for a failed run this is + typically the task that caused the failure. Returns None when every + task in the run is successful. + + Together with `Flow.failed_runs` and `Task.failure_summary` this lets + you investigate failures through the regular client: + + ``` + for run in Flow("MyFlow").failed_runs(max_runs=10): + task = run.failed_task + if task is not None: + summary = task.failure_summary + print(task.pathspec, summary and summary.exception_type) + ``` + + Returns + ------- + Task, optional + The first unsuccessful task, or None if the run has none. + """ + for step in self: + for task in step: + if not task.successful: + return task + return None + @property def finished(self) -> bool: """ @@ -2671,6 +2778,39 @@ def runs( return runs return islice(runs, max_runs) + def failed_runs( + self, + *, + since: Optional[int] = None, + max_runs: Optional[int] = None, + ) -> Iterator[Run]: + """ + Returns an iterator over the failed `Run`s of this flow, newest first. + + This is a convenience wrapper over `runs` that applies the metadata + service's ``status:eq`` = ``failed`` filter server-side. Because it + relies on server-side filtering, it requires a metadata service with + pagination and filtering support; against the local metadata provider + or an older service it raises, exactly like ``runs(filters=...)``. + + Parameters + ---------- + since : int, optional + Inclusive lower bound on run start time as epoch milliseconds + (``ts_epoch``). Only runs at or after this time are returned. + max_runs : int, optional + Maximum number of failed runs to yield, newest first. + + Yields + ------ + Run + Failed `Run` objects in this flow, newest first. + """ + filters = {"status:eq": "failed"} + if since is not None: + filters["ts_epoch:ge"] = int(since) + return self.runs(filters=filters, max_runs=max_runs) + def __iter__(self) -> Iterator[Task]: """ Iterate over all children Run of this Flow. diff --git a/test/unit/test_client_failure_investigation.py b/test/unit/test_client_failure_investigation.py new file mode 100644 index 00000000000..a66d9e9ae71 --- /dev/null +++ b/test/unit/test_client_failure_investigation.py @@ -0,0 +1,169 @@ +import pytest + +from metaflow.client.core import ( + FailureSummary, + Flow, + Run, + Task, + _normalize_exception, +) + + +class _ExceptionObject: + """Stand-in for a deserialized exception artifact carrying attributes.""" + + def __init__(self, **attributes): + for key, value in attributes.items(): + setattr(self, key, value) + + +# --------------------------------------------------------------------------- +# _normalize_exception -- pure function, tested with real inputs +# --------------------------------------------------------------------------- + + +def test_normalize_exception_from_mapping(): + assert _normalize_exception( + {"type": "ValueError", "message": "boom", "stacktrace": "line 1"} + ) == {"type": "ValueError", "message": "boom", "stacktrace": "line 1"} + + +def test_normalize_exception_mapping_falls_back_to_exception_key(): + result = _normalize_exception({"exception": "kaboom"}) + assert result["message"] == "kaboom" + assert result["type"] is None + assert result["stacktrace"] is None + + +def test_normalize_exception_from_object_attributes(): + assert _normalize_exception( + _ExceptionObject(type="RuntimeError", message="bad", stacktrace="tb") + ) == {"type": "RuntimeError", "message": "bad", "stacktrace": "tb"} + + +def test_normalize_exception_object_without_type_uses_qualified_name(): + result = _normalize_exception(_ExceptionObject()) + assert result["type"].endswith("._ExceptionObject") + assert result["message"] # falls back to str(data) + assert result["stacktrace"] is None + + +def test_normalize_exception_from_bare_string(): + result = _normalize_exception("just a string") + assert result["message"] == "just a string" + assert result["stacktrace"] is None + + +# --------------------------------------------------------------------------- +# Task.failure_summary +# --------------------------------------------------------------------------- + + +def test_task_failure_summary_none_when_no_exception(mocker): + task = mocker.Mock() + task.exception = None + assert Task.failure_summary.fget(task) is None + + +def test_task_failure_summary_builds_summary_from_exception(mocker): + task = mocker.Mock() + task.exception = {"type": "ValueError", "message": "boom", "stacktrace": "tb"} + task.current_attempt = 2 + + summary = Task.failure_summary.fget(task) + + assert isinstance(summary, FailureSummary) + assert summary.exception_type == "ValueError" + assert summary.message == "boom" + assert summary.stacktrace == "tb" + assert summary.attempt == 2 + + +def test_task_failure_summary_propagates_read_errors(): + class _Boom: + @property + def exception(self): + raise RuntimeError("exception artifact unavailable") + + with pytest.raises(RuntimeError, match="unavailable"): + Task.failure_summary.fget(_Boom()) + + +# --------------------------------------------------------------------------- +# Run.failed_task +# --------------------------------------------------------------------------- + + +def _step(mocker, tasks): + step = mocker.MagicMock() + step.__iter__.return_value = iter(tasks) + return step + + +def _run(mocker, steps): + run = mocker.MagicMock() + run.__iter__.return_value = iter(steps) + return run + + +def test_run_failed_task_returns_first_unsuccessful_in_iteration_order(mocker): + ok = mocker.Mock(successful=True) + bad = mocker.Mock(successful=False) + run = _run(mocker, [_step(mocker, [ok, bad])]) + + assert Run.failed_task.fget(run) is bad + + +def test_run_failed_task_scans_steps_in_order(mocker): + ok = mocker.Mock(successful=True) + bad = mocker.Mock(successful=False) + later = mocker.Mock(successful=False) + run = _run(mocker, [_step(mocker, [ok]), _step(mocker, [bad, later])]) + + assert Run.failed_task.fget(run) is bad + + +def test_run_failed_task_none_when_all_successful(mocker): + run = _run( + mocker, + [_step(mocker, [mocker.Mock(successful=True), mocker.Mock(successful=True)])], + ) + + assert Run.failed_task.fget(run) is None + + +# --------------------------------------------------------------------------- +# Flow.failed_runs +# --------------------------------------------------------------------------- + + +def test_flow_failed_runs_forwards_status_filter_and_bounds(mocker): + flow = mocker.Mock() + flow.runs.return_value = iter(["r3", "r2"]) + + result = list(Flow.failed_runs(flow, max_runs=2)) + + assert result == ["r3", "r2"] + flow.runs.assert_called_once_with(filters={"status:eq": "failed"}, max_runs=2) + + +def test_flow_failed_runs_since_adds_ts_epoch_filter(mocker): + flow = mocker.Mock() + flow.runs.return_value = iter([]) + + list(Flow.failed_runs(flow, since=1700000000000)) + + flow.runs.assert_called_once_with( + filters={"status:eq": "failed", "ts_epoch:ge": 1700000000000}, + max_runs=None, + ) + + +def test_flow_failed_runs_returns_the_runs_iterator_directly(mocker): + flow = mocker.Mock() + sentinel = iter(["r1"]) + flow.runs.return_value = sentinel + + # Ergonomics: failed_runs hands back exactly what runs() returns, so callers + # can do `for run in flow.failed_runs(): ...` without an extra wrapper. + assert Flow.failed_runs(flow) is sentinel From 7ea2f8c7e60fa237019a6b49a429a5f69e364714 Mon Sep 17 00:00:00 2001 From: Aryan Dhawan Date: Sun, 23 Aug 2026 18:11:40 -0400 Subject: [PATCH 6/6] Address review: failed_task returns the latest failed task, say so --- metaflow/client/core.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/metaflow/client/core.py b/metaflow/client/core.py index bb745cc623f..e36e1d05636 100644 --- a/metaflow/client/core.py +++ b/metaflow/client/core.py @@ -2413,13 +2413,13 @@ def successful(self) -> bool: @property def failed_task(self) -> Optional[Task]: """ - Returns the first failed (unsuccessful) task in this run, if any. + Returns the latest failed (unsuccessful) task in this run, if any. Steps and their tasks are scanned in iteration order, which is - newest-created first (see `MetaflowObject.__iter__`). The first task - whose `successful` is False is returned; for a failed run this is - typically the task that caused the failure. Returns None when every - task in the run is successful. + newest-created first (see `MetaflowObject.__iter__`), so the first + match is the latest failed task; for a failed run this is typically + the task that caused the failure. Returns None when every task in + the run is successful. Together with `Flow.failed_runs` and `Task.failure_summary` this lets you investigate failures through the regular client: @@ -2435,7 +2435,7 @@ def failed_task(self) -> Optional[Task]: Returns ------- Task, optional - The first unsuccessful task, or None if the run has none. + The latest unsuccessful task, or None if the run has none. """ for step in self: for task in step: