diff --git a/langfuse/_client/client.py b/langfuse/_client/client.py index 42d861fd4..889d3de1e 100644 --- a/langfuse/_client/client.py +++ b/langfuse/_client/client.py @@ -3272,7 +3272,7 @@ def run_batched_evaluation( resume_from: Optional[BatchEvaluationResumeToken] = None, verbose: bool = False, ) -> BatchEvaluationResult: - """Fetch traces or observations using legacy read APIs and evaluate each item. + """Fetch traces or observations using the v2 observations API and evaluate each item. This method provides a powerful way to evaluate existing data in Langfuse at scale. It fetches items based on filters, transforms them using a mapper function, runs @@ -3288,10 +3288,13 @@ def run_batched_evaluation( it memory-efficient for large datasets. It includes comprehensive error handling, retry logic, and resume capability for long-running evaluations. - Legacy platform compatibility: - This method reads traces from `GET /api/public/traces` and observations - from the legacy `GET /api/public/observations` endpoint. It is supported - with Langfuse platform v3 and is not yet supported with platform v4. + Data source: + Both scopes are read from `GET /api/public/v2/observations` with cursor + pagination. This works on Langfuse platform v4 events_only deployments + (where the v3 read endpoints are unavailable) and remains available on + v3. For `scope="traces"`, observations are collapsed to one + representative per trace (preferring the root observation), because + the v2 endpoint has no trace-level read. Args: scope: The type of items to evaluate. Must be one of: @@ -3310,7 +3313,12 @@ def run_batched_evaluation( Default: None (fetches all items). fetch_batch_size: Number of items to fetch per API call and hold in memory. Larger values may be faster but use more memory. Default: 50. - fetch_trace_fields: Comma-separated list of fields to include when fetching traces. Available field groups: 'core' (always included), 'io' (input, output, metadata), 'scores', 'observations', 'metrics'. If not specified, all fields are returned. Example: 'core,scores,metrics'. Note: Excluded 'observations' or 'scores' fields return empty arrays; excluded 'metrics' returns -1 for 'totalCost' and 'latency'. Only relevant if scope is 'traces'. + fetch_trace_fields: Comma-separated list of v2 observation field groups to + request (merged with the default set: 'core', 'basic', 'io', 'metadata', + 'model', 'usage', 'trace_context'). Legacy-only groups ('observations', + 'scores') are dropped because the v2 endpoint returns one observation + at a time. Example: 'io,metrics' to additionally fetch latency metrics. + Note: v2 metadata values are truncated to 200 characters unless expanded. max_items: Maximum total number of items to process. If None, processes all items matching the filter. Useful for testing or limiting evaluation runs. Default: None (process all). diff --git a/langfuse/batch_evaluation.py b/langfuse/batch_evaluation.py index 723b45757..63d4fec8b 100644 --- a/langfuse/batch_evaluation.py +++ b/langfuse/batch_evaluation.py @@ -25,11 +25,184 @@ from langfuse.api import ( ObservationsView, + ObservationV2, TraceWithFullDetails, ) from langfuse.experiment import Evaluation, EvaluatorFunction from langfuse.logger import langfuse_logger as logger +# v2 observations field groups. ``core`` carries id/trace_id/type/name and +# similar observation metadata, ``basic`` adds level/status/version/etc, +# ``io`` carries the input/output strings, ``metadata`` carries the +# observation metadata (separate from ``io``), ``usage`` carries token and +# cost details, ``model`` carries the model id, ``trace_context`` carries +# tags/release/trace_name. The legacy ``/api/public/traces`` ``io`` / +# ``scores`` / ``observations`` / ``metrics`` field groups are not selectable +# on the v2 read endpoint because it returns one observation at a time. +_DEFAULT_V2_OBSERVATION_FIELDS = "core,basic,io,metadata,model,usage,trace_context" +_LEGACY_ONLY_FIELDS = frozenset({"observations", "scores"}) + +# Trace-level filter columns are translated to their v2 observations-endpoint +# equivalents when ``scope='traces'`` so a v3-shaped filter still selects the +# intended traces via the v2 read path. +_TRACE_FILTER_COLUMN_REWRITES = { + "id": "traceId", + "name": "traceName", + "timestamp": "startTime", +} + +# ``scope='traces'`` evaluates one observation per trace, so the request is +# narrowed to root observations. The v2 endpoint pages by cursor over +# observations, not traces: a trace's root and its children can straddle a page +# boundary, so picking a representative per page would let whichever page +# arrives first decide the representative for the whole run. Asking the server +# for roots makes the selection global to the run. +# +# The server does not guarantee a single root per trace -- two sibling spans on +# the same trace have both been observed flagged as roots. The collapse is what +# reduces a page to one observation per trace, and ``seen_trace_ids`` keeps it +# to one across pages, so the guarantee comes from those two rather than from the +# filter. A trace with no flagged root is not returned at all. +_TRACE_ROOT_CONDITION = { + "type": "boolean", + "column": "isRootObservation", + "operator": "=", + "value": True, +} + + +def _validate_trace_filter(filter_json: Optional[str]) -> None: + """Reject a ``scope='traces'`` filter that is not a JSON array. + + ``_translate_trace_filter`` has to merge the root condition into the + caller's filter array. A filter that does not parse as a JSON array + therefore cannot be honoured: forwarding it would drop the caller's + constraints, and raising later -- inside the fetch loop -- would be + swallowed into a generic "failed to fetch batch" result whose resume token + carries an empty timestamp bound. + + Raises: + ValueError: If ``filter_json`` is set but is not a JSON array of + conditions. + """ + if not filter_json: + return + try: + conditions = json.loads(filter_json) + except (TypeError, ValueError) as exc: + message = ( + "batch_evaluation expects the filter to be a JSON array of conditions, " + f"but it could not be parsed: {filter_json!r}" + ) + raise ValueError(message) from exc + if not isinstance(conditions, list): + message = ( + "batch_evaluation expects the filter to be a JSON array of conditions, " + f"but it parsed as a {type(conditions).__name__}: {filter_json!r}" + ) + raise ValueError(message) + + +def _translate_trace_filter(filter_json: Optional[str]) -> str: + """Rewrite a v3-shaped trace filter into a v2 observations filter. + + The v2 endpoint cannot filter directly on trace-level columns like ``name`` + or ``timestamp``; it filters on observations, exposing those trace columns + as ``traceName`` and ``startTime``. We translate the JSON filter in place + for ``scope='traces'`` callers so that existing trace-level filters + continue to select the same set of traces. + + The root narrowing is applied here too, because the v2 endpoint ignores + query-parameter filters whenever a ``filter`` string is supplied -- sending + ``is_root_observation=True`` alongside a caller filter would be silently + ignored. A caller condition on the same column is dropped rather than + appended, because this function decides that column: a filter that excluded + roots would otherwise combine with the narrowing to return nothing. + + Returns: + A JSON array string that always carries the root condition, so the + result is never empty and never ``None``. + """ + _validate_trace_filter(filter_json) + + conditions: list[Any] = json.loads(filter_json) if filter_json else [] + + translated: list[Any] = [] + for cond in conditions: + if not isinstance(cond, dict): + translated.append(cond) + continue + # ``column`` is read before the rewrite below; the two candidate sets + # (_TRACE_FILTER_COLUMN_REWRITES keys and "isRootObservation") are + # disjoint, so the pre-rewrite value is the right one for both checks. + column = cond.get("column") + if column in _TRACE_FILTER_COLUMN_REWRITES: + cond["column"] = _TRACE_FILTER_COLUMN_REWRITES[column] + if column == "isRootObservation": + # Owned by the root narrowing below, whatever the caller sent. + continue + translated.append(cond) + translated.append(dict(_TRACE_ROOT_CONDITION)) + return json.dumps(translated) + + +def _v2_observations_fields(fetch_fields: Optional[str]) -> str: + """Translate the legacy ``fetch_trace_fields`` argument into the v2 field + groups understood by ``GET /api/public/v2/observations``. + + The default groups cover everything the existing evaluator contract reads. + Caller-supplied groups are merged in; legacy-only groups are dropped. + """ + if fetch_fields is None: + return _DEFAULT_V2_OBSERVATION_FIELDS + + user_groups = {item.strip() for item in fetch_fields.split(",") if item.strip()} + base_groups = {item.strip() for item in _DEFAULT_V2_OBSERVATION_FIELDS.split(",")} + merged = (user_groups | base_groups) - _LEGACY_ONLY_FIELDS + if not merged: + return _DEFAULT_V2_OBSERVATION_FIELDS + return ",".join(sorted(merged)) + + +def _collapse_observations_to_traces( + observations: List[ObservationV2], + seen_trace_ids: Optional[Set[str]] = None, +) -> List[ObservationV2]: + """Collapse a flat list of observations to one observation per trace_id. + + The v2 observations endpoint has no trace-level read; this helper takes + whatever observations the page returned and returns one representative + per trace, preferring the observation that the server already marked as + the root (``is_root_observation=True``) regardless of its position in the + page. + + ``seen_trace_ids``, when provided, holds the trace IDs already processed + earlier in the run. Observations for those traces are skipped so the same + trace is not evaluated twice. Now that ``scope='traces'`` asks the server + for root observations only, a well-behaved response already carries at most + one row per trace, so this is a safety net rather than the mechanism that + makes the collapse correct across pages. + """ + chosen: Dict[str, ObservationV2] = {} + for observation in observations: + trace_id = getattr(observation, "trace_id", None) + if not trace_id: + continue + if seen_trace_ids is not None and trace_id in seen_trace_ids: + continue + existing = chosen.get(trace_id) + if existing is None or ( + getattr(observation, "is_root_observation", False) + and not getattr(existing, "is_root_observation", False) + ): + chosen[trace_id] = observation + + if seen_trace_ids is not None: + seen_trace_ids.update(chosen.keys()) + + return list(chosen.values()) + + if TYPE_CHECKING: from langfuse._client.client import Langfuse @@ -137,7 +310,7 @@ class MapperFunction(Protocol): def __call__( self, *, - item: Union["TraceWithFullDetails", "ObservationsView"], + item: Union["TraceWithFullDetails", "ObservationsView", "ObservationV2"], **kwargs: Dict[str, Any], ) -> Union[EvaluatorInputs, Awaitable[EvaluatorInputs]]: """Transform an API response object into evaluator inputs. @@ -148,8 +321,11 @@ def __call__( Args: item: The API response object to transform. The type depends on the scope: - - TraceWithFullDetails: When evaluating traces - - ObservationsView: When evaluating observations + - TraceWithFullDetails: legacy v3 trace response + - ObservationsView: legacy v1 observations response + - ObservationV2: v2 observations response (used on Langfuse + platform v4 events_only deployments where the legacy endpoints + are unavailable) Returns: EvaluatorInputs: A structured container with: @@ -162,14 +338,17 @@ def __call__( (for async mappers that need to fetch additional data). Examples: - Basic trace mapper: + Basic mapper for ``scope='traces'``: ```python def map_trace(trace): + # For scope='traces' the item is the trace's root observation, + # not a whole trace: `id` is the observation ID and `trace_id` + # is the trace. Scores are written against `trace_id`. return EvaluatorInputs( input=trace.input, output=trace.output, expected_output=None, - metadata={"trace_id": trace.id, "user": trace.user_id} + metadata={"trace_id": trace.trace_id, "user": trace.user_id} ) ``` @@ -835,6 +1014,10 @@ def __init__(self, client: "Langfuse"): client: The Langfuse client instance. """ self.client = client + # Holds trace IDs already processed in earlier cursor pages so that + # ``scope='traces'`` never evaluates the same trace twice. Reset at + # the start of each ``run_async`` call. + self._seen_trace_ids: Set[str] = set() async def run_async( self, @@ -855,15 +1038,24 @@ async def run_async( verbose: bool = False, resume_from: Optional[BatchEvaluationResumeToken] = None, ) -> BatchEvaluationResult: - """Run batch evaluation asynchronously using legacy read APIs. + """Run batch evaluation asynchronously using the v2 observations API. This is the main implementation method that orchestrates the entire batch evaluation process: fetching items, mapping, evaluating, creating scores, and tracking statistics. - This runner reads traces from `GET /api/public/traces` and observations - from the legacy `GET /api/public/observations` endpoint. It is supported - with Langfuse platform v3 and is not yet supported with platform v4. + This runner reads both scopes from `GET /api/public/v2/observations` + with cursor pagination. That endpoint is the only read path available on + Langfuse platform v4 events_only deployments and remains available on + v3. For `scope='traces'`, the request itself is narrowed to root + observations, so the v2 endpoint returns at most one observation per + trace. The narrowing has to happen server-side: the endpoint pages by + cursor over observations rather than traces, so a trace's root and its + children can straddle a page boundary, and collapsing each page + independently would let whichever page arrived first fix the + representative for the whole run. A consequence worth knowing: a trace + with no observation the server marks as a root is not returned, and so + is not evaluated. Args: scope: The type of items to evaluate ("traces", "observations"). @@ -871,7 +1063,13 @@ async def run_async( evaluators: List of evaluation functions to run on each item. filter: JSON filter string for querying items. fetch_batch_size: Number of items to fetch per API call. - fetch_trace_fields: Comma-separated list of fields to include when fetching traces. Available field groups: 'core' (always included), 'io' (input, output, metadata), 'scores', 'observations', 'metrics'. If not specified, all fields are returned. Example: 'core,scores,metrics'. Note: Excluded 'observations' or 'scores' fields return empty arrays; excluded 'metrics' returns -1 for 'totalCost' and 'latency'. Only relevant if scope is 'traces'. Default: 'io' + fetch_trace_fields: Comma-separated list of v2 observation field groups to + request (merged with the default set: 'core', 'basic', 'io', 'metadata', + 'model', 'usage', 'trace_context'). Legacy-only groups ('observations', + 'scores') are dropped because the v2 endpoint returns one observation + at a time. Example: 'io,metrics' to additionally fetch latency metrics. + Note: v2 metadata values are truncated to 200 characters unless expanded. + Default: the full default set listed above. max_items: Maximum number of items to process (None = all). max_concurrency: Maximum number of concurrent evaluations. composite_evaluator: Optional function to create composite scores. @@ -889,6 +1087,9 @@ async def run_async( """ start_time = time.time() + # Reset cross-page trace-collapse state for this run. + self._seen_trace_ids.clear() + # Initialize tracking variables total_items_fetched = 0 total_items_processed = 0 @@ -909,6 +1110,14 @@ async def run_async( } # Handle resume token by modifying filter + if scope == "traces": + # Validate before the fetch loop: `_translate_trace_filter` merges + # the root condition into the caller's filter array, so a filter + # that is not a JSON array cannot be forwarded. Validating here + # rather than inside the fetch means the caller sees the bad filter + # instead of a swallowed fetch failure carrying an empty resume + # timestamp. + _validate_trace_filter(filter) effective_filter = self._build_timestamp_filter(filter, resume_from) normalized_additional_trace_tags = ( self._dedupe_tags(_additional_trace_tags) @@ -920,8 +1129,10 @@ async def run_async( # Create semaphore for concurrency control semaphore = asyncio.Semaphore(max_concurrency) - # Pagination state - page = 1 + # Pagination state. The v2 observations endpoint is cursor-based, so + # the runner walks the result set with ``next_cursor`` and stops when + # the server returns ``meta.cursor=None``. + cursor: Optional[str] = None has_more = True last_item_timestamp: Optional[str] = None last_item_id: Optional[str] = None @@ -948,19 +1159,22 @@ async def run_async( # Fetch next batch with retry logic try: - items = await self._fetch_batch_with_retry( + items, next_cursor = await self._fetch_batch_with_retry( scope=scope, filter=effective_filter, - page=page, + cursor=cursor, limit=fetch_batch_size, max_retries=max_retries, fields=fetch_trace_fields, ) except Exception as e: - # Failed after max_retries - create resume token and return + # Failed after max_retries - flush what has been processed so + # far, then create a resume token and return. error_msg = f"Failed to fetch batch after {max_retries} retries" logger.error("%s: %s", error_msg, e) + self.client.flush() + resume_token = BatchEvaluationResumeToken( scope=scope, filter=filter, # Original filter, not modified @@ -986,6 +1200,15 @@ async def run_async( item_evaluations=item_evaluations, ) + # Advance the cursor and stop when the server reports it is + # done. An empty page is also treated as the end of the stream: + # under v2 semantics an empty page typically coincides with + # ``cursor=None``, and breaking here matches the pre-v3 path's + # behaviour. + cursor = next_cursor + if cursor is None: + has_more = False + # Check if we got any items if not items: has_more = False @@ -996,7 +1219,7 @@ async def run_async( total_items_fetched += len(items) if verbose: - logger.info("Fetched batch %s (%s items)", page, len(items)) + logger.info("Fetched batch (%s items)", len(items)) # Limit items if max_items would be exceeded items_to_process = items @@ -1013,7 +1236,7 @@ async def run_async( # Process items concurrently async def process_item( - item: Union[TraceWithFullDetails, ObservationsView], + item: Union[TraceWithFullDetails, ObservationsView, ObservationV2], ) -> Tuple[str, Union[Tuple[int, int, int, List[Evaluation]], Exception]]: """Process a single item and return (item_id, result).""" async with semaphore: @@ -1094,17 +1317,15 @@ async def process_item( total_scores_created, ) - # Check if we should continue to next page - if len(items) < fetch_batch_size: - # Last page - no more items available - has_more = False - else: - page += 1 - - # Check max_items again before next fetch - if max_items is not None and total_items_fetched >= max_items: + # Check if we should continue to next page. ``has_more`` is set inside + # the fetch block above, so we only need to honour ``max_items`` + # here. If the server has already reported ``cursor=None`` the + # call set ``has_more = False``; do not flip it back to ``True`` + # just because ``max_items`` was reached on the last page. + if max_items is not None and total_items_fetched >= max_items: + if cursor is not None: has_more = True # More items exist but we're stopping - break + break # Flush all scores to Langfuse if verbose: @@ -1152,52 +1373,83 @@ async def _fetch_batch_with_retry( *, scope: str, filter: Optional[str], - page: int, + cursor: Optional[str], limit: int, max_retries: int, fields: Optional[str], - ) -> List[Union[TraceWithFullDetails, ObservationsView]]: - """Fetch a batch of items with retry logic. + ) -> Tuple[ + List[Union[TraceWithFullDetails, ObservationsView, ObservationV2]], + Optional[str], + ]: + """Fetch a batch of items via the v2 observations API with cursor pagination. + + Both ``scope='traces'`` and ``scope='observations'`` go through the + v2 ``GET /api/public/v2/observations`` endpoint, which is the only + read path that works on Langfuse platform v4 events_only deployments. + The legacy ``GET /api/public/traces`` and ``GET /api/public/observations`` + endpoints are unavailable there. Args: scope: The type of items ("traces", "observations"). filter: JSON filter string for querying. - page: Page number (1-indexed). - limit: Number of items per page. + cursor: Pagination cursor returned by the previous page; ``None`` on + the first call. + limit: Number of items to request per page. max_retries: Maximum number of retry attempts. - verbose: Whether to log retry attempts. - fields: Trace fields to fetch + fields: Comma-separated list of v2 field groups to include when + fetching traces. Maps to the ``fields`` query parameter on + ``/api/public/v2/observations``. Only used when + ``scope='traces'``. Returns: - List of items from the API. + Tuple of (items, next_cursor). ``items`` is a list of either + ``ObservationV2`` (when scope='observations') or one + ``ObservationV2`` per trace (when scope='traces'). ``next_cursor`` + is the cursor for the next page, or ``None`` when no more pages + remain. Raises: + ValueError: If ``scope`` is not "traces" or "observations". Exception: If all retry attempts fail. """ - if scope == "traces": - response = self.client.api.trace.list( - page=page, - limit=limit, - filter=filter, - request_options={"max_retries": max_retries}, - fields=fields, - ) # type: ignore - return list(response.data) # type: ignore - elif scope == "observations": - response = self.client.api.legacy.observations_v1.get_many( - page=page, - limit=limit, - filter=filter, - request_options={"max_retries": max_retries}, - ) # type: ignore - return list(response.data) # type: ignore - else: + if scope not in ("traces", "observations"): error_message = f"Invalid scope: {scope}" raise ValueError(error_message) + v2_fields = _v2_observations_fields(fields) + v2_filter = _translate_trace_filter(filter) if scope == "traces" else filter + + response = self.client.api.observations.get_many( # type: ignore[union-attr] + cursor=cursor, + limit=limit, + filter=v2_filter, + request_options={"max_retries": max_retries}, + fields=v2_fields, + ) + + next_cursor = cast( + Optional[str], response.meta.cursor if response.meta else None + ) + + if scope == "traces": + items = cast( + List[Union[TraceWithFullDetails, ObservationsView, ObservationV2]], + _collapse_observations_to_traces( + list(response.data), # type: ignore[arg-type] + seen_trace_ids=self._seen_trace_ids, + ), + ) + else: + items = cast( + List[Union[TraceWithFullDetails, ObservationsView, ObservationV2]], + list(response.data), # type: ignore[arg-type] + ) + + return items, next_cursor + async def _process_batch_evaluation_item( self, - item: Union[TraceWithFullDetails, ObservationsView], + item: Union[TraceWithFullDetails, ObservationsView, ObservationV2], scope: str, mapper: MapperFunction, evaluators: List[EvaluatorFunction], @@ -1352,7 +1604,7 @@ async def _run_evaluator_internal( async def _run_mapper( self, mapper: MapperFunction, - item: Union[TraceWithFullDetails, ObservationsView], + item: Union[TraceWithFullDetails, ObservationsView, ObservationV2], ) -> EvaluatorInputs: """Run mapper function (handles both sync and async mappers). @@ -1529,7 +1781,7 @@ def _build_timestamp_filter( @staticmethod def _get_item_id( - item: Union[TraceWithFullDetails, ObservationsView], + item: Union[TraceWithFullDetails, ObservationsView, ObservationV2], scope: str, ) -> str: """Extract ID from item based on scope. @@ -1539,13 +1791,19 @@ def _get_item_id( scope: The type of item. Returns: - The item's ID. + The item's ID. For ``scope='traces'`` the returned value is the + trace ID, not the observation ID, so downstream score-create calls + attach to the intended trace. """ + if scope == "traces": + trace_id = getattr(item, "trace_id", None) + if trace_id: + return trace_id # type: ignore[no-any-return,return-value] return item.id @staticmethod def _get_item_timestamp( - item: Union[TraceWithFullDetails, ObservationsView], + item: Union[TraceWithFullDetails, ObservationsView, ObservationV2], scope: str, ) -> str: """Extract timestamp from item based on scope. @@ -1555,33 +1813,25 @@ def _get_item_timestamp( scope: The type of item. Returns: - ISO 8601 timestamp string. + ISO 8601 timestamp string. For ``scope='traces'`` we use the root + observation's ``start_time`` as a proxy for the trace's timestamp + because the v2 endpoint has no trace-level read. """ - if scope == "traces": - # Type narrowing for traces - if hasattr(item, "timestamp"): - return item.timestamp.isoformat() # type: ignore[attr-defined] - elif scope == "observations": - # Type narrowing for observations - if hasattr(item, "start_time"): - return item.start_time.isoformat() # type: ignore[attr-defined] + start_time = getattr(item, "start_time", None) + if start_time is not None: + return start_time.isoformat() # type: ignore[attr-defined,no-any-return] + timestamp = getattr(item, "timestamp", None) + if timestamp is not None: + return timestamp.isoformat() # type: ignore[attr-defined,no-any-return] return "" @staticmethod def _get_timestamp_field_for_scope(scope: str) -> str: - """Get the timestamp field name for filtering based on scope. - - Args: - scope: The type of items. - - Returns: - The field name to use in filters. + """Get the v2 observations filter column for resume by last-processed + timestamp. The v2 endpoint filters on observation ``startTime``; + ``scope='traces'`` uses this as a proxy for the trace's timestamp. """ - if scope == "traces": - return "timestamp" - elif scope == "observations": - return "start_time" - return "timestamp" # Default + return "startTime" @staticmethod def _dedupe_tags(tags: Optional[List[str]]) -> List[str]: diff --git a/tests/e2e/test_batch_evaluation_root_selection.py b/tests/e2e/test_batch_evaluation_root_selection.py new file mode 100644 index 000000000..6438f28e4 --- /dev/null +++ b/tests/e2e/test_batch_evaluation_root_selection.py @@ -0,0 +1,399 @@ +"""End-to-end coverage for trace-root selection in batch evaluation. + +Unit tests cover the request the runner sends, but only a real server can show +that the ``isRootObservation`` filter is honoured and that the cross-page bug +this replaced actually occurred. Both are asserted here against the deployment +the e2e suite runs against. + +The bug: the v2 endpoint pages by cursor over observations, not traces, so a +trace's root and its children can land on different pages. Choosing a +representative per page let the first page fix the representative for the whole +run, and the cross-page ``seen`` set then suppressed the root entirely. These +tests seed a wide trace so the root cannot share a page with all its children, +which is the shape that triggers it. +""" + +import base64 +import json +import os +import time + +from langfuse import get_client +from langfuse.batch_evaluation import _collapse_observations_to_traces +from tests.support.utils import create_uuid + +ROOT_OUTPUT = "root-output-marker" +CHILD_OUTPUT = "child-output-marker" + + +def _seed_wide_trace(*, child_count: int = 12) -> str: + """Seed one trace with a root observation and many children. + + A wide trace is what makes the cross-page case reachable: with a small + `page_size` the root cannot share a page with every child. + """ + langfuse_client = get_client() + # Trace IDs must be 32 lowercase hex chars; `create_uuid()` is dashed. + trace_id = create_uuid().replace("-", "") + name = f"root-selection-{create_uuid()}" + + with langfuse_client.start_as_current_observation( + name=f"{name}-root", + trace_context={"trace_id": trace_id}, + input=f"{name}-root-input", + output=ROOT_OUTPUT, + ): + for index in range(child_count): + with langfuse_client.start_as_current_observation( + name=f"{name}-child-{index}", + input=f"{name}-child-{index}-input", + output=CHILD_OUTPUT, + ): + pass + + langfuse_client.flush() + return trace_id + + +def _fetch_observations(trace_id: str, *, page_size: int = 2, root_only: bool = False): + """Read one trace back page by page, optionally narrowed to root observations.""" + client = get_client() + conditions = None + if root_only: + conditions = json.dumps( + [ + { + "type": "boolean", + "column": "isRootObservation", + "operator": "=", + "value": True, + } + ] + ) + + pages = [] + cursor = None + while True: + response = client.api.observations.get_many( + trace_id=trace_id, limit=page_size, cursor=cursor, filter=conditions + ) + rows = list(response.data) + if rows: + pages.append(rows) + cursor = response.meta.cursor if response.meta else None + if cursor is None or not rows: + break + return pages + + +def _wait_for_trace_observations( + trace_id: str, *, expected_min: int, timeout: float = 90.0 +): + """Ingestion is async; wait until the trace has its observations. + + Returns the rows so callers do not immediately re-fetch. + """ + deadline = time.time() + timeout + rows: list = [] + while time.time() < deadline: + rows = list( + get_client().api.observations.get_many(trace_id=trace_id, limit=100).data + ) + if len(rows) >= expected_min: + return rows + time.sleep(2.0) + raise AssertionError( + f"expected at least {expected_min} observations for {trace_id}, got {len(rows)}" + ) + + +def test_is_root_observation_filter_is_honoured_by_the_server(): + """The filter the runner sends must actually narrow the response. + + If the server ignored the condition, a trace's children would come back and + the root could still lose the page-local comparison. + """ + child_count = 12 + trace_id = _seed_wide_trace(child_count=child_count) + all_rows = _wait_for_trace_observations(trace_id, expected_min=child_count + 1) + + # The server flags exactly one observation on the trace as the root. + flagged_roots = [row for row in all_rows if row.is_root_observation] + assert len(flagged_roots) == 1, ( + f"expected exactly one server-side root, got {[r.name for r in flagged_roots]}" + ) + + # And the root filter returns exactly that row, not its siblings. + root_pages = _fetch_observations(trace_id, root_only=True) + filtered = [row for page in root_pages for row in page] + assert len(filtered) == 1 + assert filtered[0].id == flagged_roots[0].id + + # Sanity check that the negation excludes it, so the filter is not a no-op. + negated = json.dumps( + [ + { + "type": "boolean", + "column": "isRootObservation", + "operator": "<>", + "value": True, + } + ] + ) + non_root = list( + get_client() + .api.observations.get_many(trace_id=trace_id, limit=100, filter=negated) + .data + ) + assert len(non_root) >= child_count + assert all(row.id != flagged_roots[0].id for row in non_root) + + +def test_root_wins_even_when_it_is_not_on_the_first_page(): + """The bug this change fixes, reproduced against real server ordering. + + Without the root filter the representative is chosen per page, so a page of + children arriving before the root's page decides the representative for the + whole run. With the filter, every page carries roots only. + """ + child_count = 12 + page_size = 2 + trace_id = _seed_wide_trace(child_count=child_count) + all_rows = _wait_for_trace_observations(trace_id, expected_min=child_count + 1) + server_root = next(row for row in all_rows if row.is_root_observation) + + # Unfiltered pages: the root may not share a page with the earlier children. + # Whether it does depends on the server's ordering, so this is recorded and + # not asserted; what must hold either way is that the root-filtered path + # yields only the root, once. + pages = _fetch_observations(trace_id, page_size=page_size, root_only=False) + root_on_first_page = any(row.id == server_root.id for row in pages[0]) + + # Pre-fix behaviour: page-local collapse with no root filter upstream. + seen: set = set() + pre_fix_ids = [ + row.id for page in pages for row in _collapse_observations_to_traces(page, seen) + ] + + # Post-fix behaviour: the same helper, fed by root-only pages. + root_pages = _fetch_observations(trace_id, page_size=page_size, root_only=True) + seen_post: set = set() + post_fix_ids = [ + row.id + for page in root_pages + for row in _collapse_observations_to_traces(page, seen_post) + ] + + assert post_fix_ids == [server_root.id], ( + f"expected only the root {[server_root.name]}, got {post_fix_ids}" + ) + # The pre-fix path evaluated exactly one observation for this trace too -- + # that was the original defect, not a duplicate-evaluation bug. Whether it + # picked the root depended on page ordering. + assert len(pre_fix_ids) == 1 + print( + f"\n[info] root_on_first_page={root_on_first_page} " + f"pre_fix_picked_root={pre_fix_ids[0] == server_root.id} " + f"post_fix_picked_root={post_fix_ids[0] == server_root.id}" + ) + + +def test_run_batched_evaluation_on_traces_evaluates_the_root(): + """End to end through the public API: the mapper must receive the root. + + The seeded root and its children carry different outputs, so a child + representative is detectable from what the mapper saw. + """ + trace_id = _seed_wide_trace(child_count=12) + rows = _wait_for_trace_observations(trace_id, expected_min=13) + root = next(row for row in rows if row.is_root_observation) + + seen_ids: list = [] + seen_outputs: list = [] + + def mapper(*, item): + seen_ids.append(item.id) + seen_outputs.append(getattr(item, "output", None)) + return None + + get_client().run_batched_evaluation( + scope="traces", + mapper=mapper, + evaluators=[], + fetch_batch_size=2, + ) + get_client().flush() + + assert root.id in seen_ids, "the root observation was never evaluated" + assert CHILD_OUTPUT not in [o for o in seen_outputs if o is not None], ( + "a child observation was evaluated instead of the root" + ) + + +def test_multiple_flagged_roots_on_one_trace_still_collapse_to_one(): + """The server can flag more than one root on a single trace. + + Two sibling spans on the same trace have both been observed carrying + ``isRootObservation=True``, so "one root per trace" is not a guarantee the + filter provides. What has to hold is that the runner still evaluates such a + trace exactly once, which is the collapse's job rather than the filter's. + """ + langfuse_client = get_client() + trace_id = create_uuid().replace("-", "") + name = f"multi-root-{create_uuid()}" + + # Siblings: neither is nested in the other, so both are eligible app roots. + for index in range(2): + with langfuse_client.start_as_current_observation( + name=f"{name}-sibling-{index}", + trace_context={"trace_id": trace_id}, + input=f"{name}-in-{index}", + output=f"sibling-out-{index}", + ): + pass + langfuse_client.flush() + + rows = _wait_for_trace_observations(trace_id, expected_min=2) + flagged = [row for row in rows if row.is_root_observation] + root_pages = _fetch_observations(trace_id, root_only=True) + filtered = [row for page in root_pages for row in page] + + print( + f"\n[info] observations={len(rows)} flagged_roots={len(flagged)} " + f"root_filter_returned={len(filtered)}" + ) + + # Whether the server flags one root or several here is server behaviour, not + # something this test should pin. Either way the runner must collapse the + # page to one item for the trace. + seen: set = set() + collapsed = _collapse_observations_to_traces(filtered, seen) + assert len(collapsed) == 1, ( + f"expected one representative for the trace, got {[c.name for c in collapsed]}" + ) + assert collapsed[0].is_root_observation + + # And across pages: two flagged roots on separate pages must still yield one + # evaluation, which is what `seen_trace_ids` is for. + if len(flagged) > 1: + seen_across: set = set() + first = _collapse_observations_to_traces([flagged[0]], seen_across) + second = _collapse_observations_to_traces(flagged[1:], seen_across) + assert len(first) + len(second) == 1, ( + f"a multi-root trace was evaluated {len(first) + len(second)} times" + ) + + +def test_trace_with_no_flagged_root_is_not_evaluated(): + """A trace whose observations carry no root is silently skipped. + + Traces ingested by other clients -- an OTel collector, another language SDK, + direct OTLP -- never pass through this SDK's app-root marking, so nothing on + them is flagged. The root filter then returns nothing for them and + `scope='traces'` never evaluates them. + + The span is written through the OTLP endpoint with a parent id that does not + exist, so there is no parent-less observation either. The observation write + endpoints refuse observation creates on an events_only deployment + (`/api/public/observations` answers 405, `/api/public/ingestion` answers + "Event type not accepted ... only accepts score events"), and the error names + the OTLP path as the events_only-compatible route, so that is what is used. + """ + import httpx + + client = get_client() + trace_id = create_uuid().replace("-", "") + span_id = create_uuid().replace("-", "")[:16] + name = f"no-root-{create_uuid()}" + + public_key = os.environ["LANGFUSE_PUBLIC_KEY"] + secret_key = os.environ["LANGFUSE_SECRET_KEY"] + base_url = os.environ.get("LANGFUSE_BASE_URL", "http://localhost:3000") + + payload = { + "resourceSpans": [ + { + "resource": { + "attributes": [ + {"key": "service.name", "value": {"stringValue": name}}, + ] + }, + "scopeSpans": [ + { + "scope": {"name": name}, + "spans": [ + { + "traceId": trace_id, + "spanId": span_id, + # A parent that does not exist: no observation + # is parent-less, and nothing is flagged. + "parentSpanId": create_uuid().replace("-", "")[:16], + "name": name, + "kind": 1, + "startTimeUnixNano": "1767225600000000000", + "endTimeUnixNano": "1767225601000000000", + "attributes": [ + { + "key": "langfuse.observation.type", + "value": {"stringValue": "SPAN"}, + }, + { + "key": "langfuse.trace.name", + "value": {"stringValue": name}, + }, + ], + } + ], + } + ], + } + ] + } + + response = httpx.post( + f"{base_url}/api/public/otel/v1/traces", + json=payload, + headers={ + "Authorization": "Basic " + + base64.b64encode(f"{public_key}:{secret_key}".encode()).decode("ascii"), + "x-langfuse-sdk-name": "python", + "x-langfuse-sdk-version": "e2e", + "x-langfuse-public-key": public_key, + "Content-Type": "application/json", + }, + timeout=30.0, + ) + assert response.status_code == 200, ( + f"OTLP ingest failed: {response.status_code} {response.text[:200]}" + ) + client.flush() + + rows = _wait_for_trace_observations(trace_id, expected_min=1) + assert rows, ( + "the OTLP span never became readable; the no-root case was not exercised" + ) + assert all(not row.is_root_observation for row in rows), ( + "the span was flagged as a root, so this deployment does not reproduce the no-root case" + ) + + root_pages = _fetch_observations(trace_id, root_only=True) + assert root_pages == [], "the root filter returned rows for a trace with no root" + + # The consequence: scope='traces' never sees this trace at all. + seen_ids: list = [] + + def mapper(*, item): + seen_ids.append(item.id) + return None + + get_client().run_batched_evaluation( + scope="traces", + mapper=mapper, + evaluators=[], + fetch_batch_size=2, + ) + get_client().flush() + + assert not any(row.id in seen_ids for row in rows), ( + "a trace with no flagged root was evaluated; the documented consequence changed" + ) diff --git a/tests/unit/test_batch_evaluation_fetch.py b/tests/unit/test_batch_evaluation_fetch.py new file mode 100644 index 000000000..e818fca34 --- /dev/null +++ b/tests/unit/test_batch_evaluation_fetch.py @@ -0,0 +1,842 @@ +"""Unit tests for BatchEvaluationRunner._fetch_batch_with_retry. + +These tests cover the v2 observations API path used by `batch_evaluation`, +so the SDK works on Langfuse platform v4 events_only deployments where the +legacy `/api/public/traces` and `/api/public/observations` endpoints are +unavailable. See langfuse/langfuse#1861. +""" + +from __future__ import annotations + +from datetime import datetime +from typing import Any +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from langfuse.batch_evaluation import ( + BatchEvaluationRunner, + _collapse_observations_to_traces, + _v2_observations_fields, +) + + +def _v2_response(*, items: list[Any], cursor: str | None = None) -> MagicMock: + response = MagicMock() + response.data = items + response.meta.cursor = cursor + return response + + +def _obs( + *, + id: str, + trace_id: str, + input: Any = None, + output: Any = None, + is_root: bool = False, +) -> MagicMock: + obs = MagicMock() + obs.id = id + obs.trace_id = trace_id + obs.input = input + obs.output = output + obs.is_root_observation = is_root + return obs + + +class _StubRunner(BatchEvaluationRunner): + """Runner whose v3 endpoints raise and whose _process_batch_evaluation_item + is replaced so the unit test focuses on the fetch path.""" + + def __init__(self) -> None: + super().__init__(client=MagicMock()) + self.client = MagicMock() + self.client.api.trace.list.side_effect = AssertionError( + "v3 GET /api/public/traces must not be called on v4 events_only" + ) + self.client.api.legacy.observations_v1.get_many.side_effect = AssertionError( + "v1 GET /api/public/observations must not be called on v4 events_only" + ) + self.client.flush = MagicMock() + + +@pytest.mark.asyncio +async def test_fetch_batch_uses_v2_observations_api_for_observations_scope() -> None: + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response(items=[_obs(id="obs-1", trace_id="t-1")], cursor="c-2"), + _v2_response(items=[_obs(id="obs-2", trace_id="t-2")], cursor=None), + ] + + items: list = [] + cursor: Any = None + for _ in range(3): + batch, cursor = await runner._fetch_batch_with_retry( + scope="observations", + filter=None, + cursor=cursor, + limit=50, + max_retries=1, + fields=None, + ) + items.extend(batch) + if cursor is None: + break + + assert [item.id for item in items] == ["obs-1", "obs-2"] + calls = runner.client.api.observations.get_many.call_args_list + assert len(calls) == 2 + assert calls[0].kwargs["cursor"] is None + assert calls[0].kwargs["limit"] == 50 + assert calls[1].kwargs["cursor"] == "c-2" + + +@pytest.mark.asyncio +async def test_fetch_batch_groups_observations_per_trace_for_traces_scope() -> None: + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[ + _obs(id="root-a", trace_id="ta", is_root=True), + _obs(id="child-a", trace_id="ta", is_root=False), + _obs(id="root-b", trace_id="tb", is_root=True), + ], + ) + ] + + items, _cursor = await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + # Each trace appears once, with its root observation preferred. + assert [item.id for item in items] == ["root-a", "root-b"] + assert runner.client.api.observations.get_many.call_count == 1 + kwargs = runner.client.api.observations.get_many.call_args.kwargs + assert kwargs["cursor"] is None + assert kwargs["limit"] == 50 + + +@pytest.mark.asyncio +async def test_fetch_batch_rejects_unknown_scope() -> None: + runner = _StubRunner() + runner.client.api.observations.get_many.side_effect = AssertionError( + "v2 observations must not be called for an unknown scope" + ) + + with pytest.raises(ValueError, match="bogus"): + await runner._fetch_batch_with_retry( + scope="bogus", + filter=None, + cursor=None, + limit=10, + max_retries=1, + fields=None, + ) + + +@pytest.mark.asyncio +async def test_fetch_batch_falls_back_when_trace_has_no_root_observation() -> None: + """When scope='traces' but no observation on a page is marked as the + root, the helper must still collapse to one item per trace using the + first observation seen.""" + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[ + _obs(id="first-t", trace_id="t1", is_root=False), + _obs(id="second-t", trace_id="t1", is_root=False), + ], + ) + ] + + items, _cursor = await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + assert [item.id for item in items] == ["first-t"] + + +@pytest.mark.asyncio +async def test_fetch_batch_filters_out_observations_without_trace_id() -> None: + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + no_trace = MagicMock() + no_trace.id = "orphan" + no_trace.trace_id = None + no_trace.is_root_observation = False + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[ + no_trace, + _obs(id="kept", trace_id="t1", is_root=True), + ], + ) + ] + + items, _cursor = await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + assert [item.id for item in items] == ["kept"] + + +@pytest.mark.asyncio +async def test_fetch_batch_prefers_root_observation_regardless_of_page_order() -> None: + """When the root is not the first observation on the page, the helper + must still pick it as the trace's representative.""" + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[ + _obs(id="child-ta", trace_id="ta", is_root=False), + _obs(id="root-ta", trace_id="ta", is_root=True), + ], + ) + ] + + items, _cursor = await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + assert [item.id for item in items] == ["root-ta"] + + +def test_v2_observations_fields_defaults() -> None: + assert ( + _v2_observations_fields(None) + == "core,basic,io,metadata,model,usage,trace_context" + ) + + +def test_v2_observations_fields_merges_user_supplied_group_with_defaults() -> None: + # ``io`` is the legacy default for ``fetch_trace_fields``; the user supply + # is preserved and the v2 default groups are added. The merged set is + # sorted alphabetically to produce a stable comma-separated string. + result = _v2_observations_fields("io") + assert result == "basic,core,io,metadata,model,trace_context,usage" + + +def test_v2_observations_fields_drops_legacy_only_groups() -> None: + result = _v2_observations_fields("observations,scores,io") + assert "observations" not in result.split(",") + assert "scores" not in result.split(",") + assert "io" in result.split(",") + assert "metadata" in result.split(",") + + +def test_v2_observations_fields_falls_back_when_user_supply_is_all_legacy() -> None: + # If the caller passes only legacy-only groups, we cannot satisfy them on + # v2 and we fall back to the full default set. The merged set is then + # sorted alphabetically to produce a stable comma-separated string. + result = _v2_observations_fields("observations,scores") + assert result == "basic,core,io,metadata,model,trace_context,usage" + + +@pytest.mark.asyncio +async def test_fetch_batch_does_not_re_evaluate_traces_already_seen() -> None: + """When a trace's observations span cursor pages, the trace-collapse state + must remember it on the first page so the same trace is not evaluated + twice on a later page.""" + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + # Page 1: trace ``ta`` appears with its root. Page 2 also returns + # observations on ``ta`` (e.g. via ``order_by``), but the trace has + # already been processed. + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[_obs(id="root-ta", trace_id="ta", is_root=True)], + cursor="c-2", + ), + _v2_response( + items=[_obs(id="child-ta", trace_id="ta", is_root=False)], + cursor=None, + ), + ] + + seen: list = [] + cursor: Any = None + for _ in range(3): + batch, cursor = await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=cursor, + limit=50, + max_retries=1, + fields=None, + ) + seen.extend(item.id for item in batch) + if cursor is None: + break + + assert seen == ["root-ta"] + + +@pytest.mark.asyncio +async def test_fetch_batch_translates_trace_filter_to_v2_columns() -> None: + """For ``scope='traces'`` the v3 trace-level filter columns ``name`` and + ``timestamp`` are rewritten to ``traceName`` / ``startTime`` so the v2 + endpoint returns the same traces.""" + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response(items=[_obs(id="o1", trace_id="t1")], cursor=None), + ] + + filter_json = ( + '[{"type":"string","column":"name","operator":"=","value":"checkout"},' + '{"type":"datetime","column":"timestamp","operator":">","value":"2026-01-01T00:00:00Z"},' + '{"type":"string","column":"id","operator":"=","value":"trace-abc"}]' + ) + + await runner._fetch_batch_with_retry( + scope="traces", + filter=filter_json, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + kwargs = runner.client.api.observations.get_many.call_args.kwargs + import json as _json + + sent_filter = _json.loads(kwargs["filter"]) + columns = {c["column"] for c in sent_filter} + assert "traceName" in columns + assert "startTime" in columns + assert "traceId" in columns + assert "name" not in columns + assert "timestamp" not in columns + assert "id" not in columns + + +@pytest.mark.asyncio +async def test_fetch_batch_does_not_translate_filter_for_observations_scope() -> None: + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response(items=[_obs(id="o1", trace_id="t1")], cursor=None), + ] + + filter_json = ( + '[{"type":"string","column":"name","operator":"=","value":"checkout"}]' + ) + + await runner._fetch_batch_with_retry( + scope="observations", + filter=filter_json, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + kwargs = runner.client.api.observations.get_many.call_args.kwargs + import json as _json + + sent_filter = _json.loads(kwargs["filter"]) + assert sent_filter[0]["column"] == "name" + + +@pytest.mark.asyncio +async def test_fetch_batch_requests_root_observations_for_traces_scope() -> None: + """For ``scope='traces'`` the v2 request is narrowed to root observations. + + Root selection has to be global to the run. Choosing a representative per + page cannot work: a trace's root and its children can straddle a cursor + page boundary, and the page that happens to arrive first would fix the + representative for the whole run. Asking the server for + ``isRootObservation = true`` makes every returned row a root, so each page + yields at most one observation per trace. + """ + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[_obs(id="root-t1", trace_id="t1", is_root=True)], cursor=None + ), + ] + + await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + kwargs = runner.client.api.observations.get_many.call_args.kwargs + import json as _json + + conditions = _json.loads(kwargs["filter"]) + root_conditions = [c for c in conditions if c["column"] == "isRootObservation"] + assert root_conditions == [ + { + "type": "boolean", + "column": "isRootObservation", + "operator": "=", + "value": True, + } + ] + + +@pytest.mark.asyncio +async def test_fetch_batch_keeps_caller_root_condition_when_filter_supplied() -> None: + """A caller filter that already constrains ``isRootObservation`` is not + duplicated, and the trace-scope translation still applies alongside it.""" + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[_obs(id="root-t1", trace_id="t1", is_root=True)], cursor=None + ), + ] + + filter_json = ( + '[{"type":"string","column":"name","operator":"=","value":"checkout"},' + '{"type":"boolean","column":"isRootObservation","operator":"=","value":false}]' + ) + + await runner._fetch_batch_with_retry( + scope="traces", + filter=filter_json, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + kwargs = runner.client.api.observations.get_many.call_args.kwargs + import json as _json + + conditions = _json.loads(kwargs["filter"]) + assert [c["column"] for c in conditions].count("isRootObservation") == 1 + # The conflicting caller condition is replaced rather than merely deduped, + # otherwise the request could return no rows at all. + root_conditions = [c for c in conditions if c["column"] == "isRootObservation"] + assert root_conditions[0]["value"] is True + assert "traceName" in {c["column"] for c in conditions} + + +@pytest.mark.asyncio +async def test_fetch_batch_does_not_request_roots_for_observations_scope() -> None: + """``scope='observations'`` evaluates every observation, so the root + narrowing must not be applied there.""" + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response(items=[_obs(id="o1", trace_id="t1")], cursor=None), + ] + + await runner._fetch_batch_with_retry( + scope="observations", + filter=None, + cursor=None, + limit=50, + max_retries=1, + fields=None, + ) + + kwargs = runner.client.api.observations.get_many.call_args.kwargs + assert kwargs["filter"] is None + + +@pytest.mark.asyncio +async def test_fetch_batch_sends_root_filter_on_every_cursor_page() -> None: + """The root narrowing is re-derived per page, not carried over from page 1. + + ``_translate_trace_filter`` re-parses the filter string on every request, so + a resumed or multi-page run could lose the condition if the merge were + stateful. Each page's request must carry it independently. + """ + + runner = _StubRunner() + runner._process_batch_evaluation_item = MagicMock( # type: ignore[method-assign] + return_value=(0, 0, 0, []) + ) + + runner.client.api.observations.get_many.side_effect = [ + _v2_response( + items=[_obs(id="root-tb", trace_id="tb", is_root=True)], cursor="c-2" + ), + _v2_response( + items=[_obs(id="root-ta", trace_id="ta", is_root=True)], cursor=None + ), + ] + + seen: list = [] + cursor: Any = None + for _ in range(3): + batch, cursor = await runner._fetch_batch_with_retry( + scope="traces", + filter=None, + cursor=cursor, + limit=50, + max_retries=1, + fields=None, + ) + seen.extend(item.id for item in batch) + if cursor is None: + break + + assert seen == ["root-tb", "root-ta"] + assert runner.client.api.observations.get_many.call_count == 2 + import json as _json + + for call in runner.client.api.observations.get_many.call_args_list: + conditions = _json.loads(call.kwargs["filter"]) + assert any(c["column"] == "isRootObservation" for c in conditions) + + +def test_page_local_collapse_would_lose_a_root_on_a_later_page() -> None: + """Why the narrowing has to be server-side, shown directly on the helper. + + A v2 page can carry a trace's child while its root sits on a later page. + Given only that page, ``_collapse_observations_to_traces`` picks the child -- + it cannot see the root -- and once the trace is marked seen the root is + never evaluated at all. This is the defect the request-level root filter + removes; it is pinned here so the fallback cannot silently become the + primary path again. + """ + + page_with_only_a_child = [ + _obs(id="child-ta", trace_id="ta", is_root=False), + ] + later_page_with_the_root = [ + _obs(id="root-ta", trace_id="ta", is_root=True), + ] + + seen: set = set() + first = _collapse_observations_to_traces(page_with_only_a_child, seen) + second = _collapse_observations_to_traces(later_page_with_the_root, seen) + + assert [o.id for o in first] == ["child-ta"] + # The root is suppressed, so the trace is evaluated on its child alone. + assert second == [] + assert "root-ta" not in {o.id for o in first + second} + + +def test_collapse_prefers_root_when_a_page_carries_both() -> None: + """Within a single page the root still wins, which is why the helper keeps + its preference logic even though the request now asks for roots only.""" + + collapsed = _collapse_observations_to_traces( + [ + _obs(id="child-a", trace_id="ta", is_root=False), + _obs(id="root-a", trace_id="ta", is_root=True), + ] + ) + + assert [o.id for o in collapsed] == ["root-a"] + + +@pytest.mark.asyncio +async def test_run_async_raises_on_malformed_filter_instead_of_resuming() -> None: + """A filter that is not a JSON array must fail at the call site. + + ``_translate_trace_filter`` merges the root condition into the caller's + filter array, so a malformed filter cannot be honoured. Validating inside + the fetch loop would be swallowed by its ``except Exception`` into a + ``completed=False`` result whose resume token carries an empty timestamp + bound -- a guaranteed 400 on the next run, with no indication that the + filter was the cause. + """ + + runner = _StubRunner() + + with pytest.raises(ValueError, match="JSON array"): + await runner.run_async( + scope="traces", + filter='{"column":"name","operator":"=","value":"checkout"}', + mapper=lambda **kw: None, + evaluators=[], + ) + + # The run must not have started fetching, and no resume token is produced. + assert runner.client.api.observations.get_many.call_count == 0 + + +@pytest.mark.asyncio +async def test_run_async_raises_on_unparseable_filter() -> None: + runner = _StubRunner() + + with pytest.raises(ValueError, match="JSON array"): + await runner.run_async( + scope="traces", + filter="not json at all", + mapper=lambda **kw: None, + evaluators=[], + ) + + assert runner.client.api.observations.get_many.call_count == 0 + + +@pytest.mark.asyncio +async def test_run_async_accepts_a_valid_filter_for_traces_scope() -> None: + """The validation must not reject the ordinary case.""" + + runner = _StubRunner() + runner.client.api.observations.get_many.side_effect = [ + _v2_response(items=[], cursor=None), + ] + + result = await runner.run_async( + scope="traces", + filter='[{"type":"string","column":"name","operator":"=","value":"checkout"}]', + mapper=lambda **kw: None, + evaluators=[], + ) + + assert result.completed is True + kwargs = runner.client.api.observations.get_many.call_args.kwargs + import json as _json + + conditions = _json.loads(kwargs["filter"]) + assert {c["column"] for c in conditions} == {"traceName", "isRootObservation"} + + +@pytest.mark.asyncio +async def test_run_async_does_not_validate_filter_for_observations_scope() -> None: + """``scope='observations'`` forwards the caller's filter unchanged, so the + traces-scope array requirement does not apply to it.""" + + runner = _StubRunner() + runner.client.api.observations.get_many.side_effect = [ + _v2_response(items=[], cursor=None), + ] + + result = await runner.run_async( + scope="observations", + filter="not json at all", + mapper=lambda **kw: None, + evaluators=[], + ) + + assert result.completed is True + + +def test_get_item_id_returns_trace_id_for_scope_traces() -> None: + """When ``scope='traces'`` the item is now an ``ObservationV2`` and its + ``id`` is the observation ID; downstream score-create calls need the + trace ID, which is ``trace_id``.""" + obs = _obs(id="obs-id", trace_id="trace-id") + assert BatchEvaluationRunner._get_item_id(obs, "traces") == "trace-id" + assert BatchEvaluationRunner._get_item_id(obs, "observations") == "obs-id" + + +def test_get_item_timestamp_uses_observation_start_time() -> None: + obs = _obs(id="o", trace_id="t") + obs.start_time = datetime(2026, 9, 24, 7, 0, 0) + assert BatchEvaluationRunner._get_item_timestamp(obs, "traces") == ( + "2026-09-24T07:00:00" + ) + assert BatchEvaluationRunner._get_item_timestamp(obs, "observations") == ( + "2026-09-24T07:00:00" + ) + + +def test_get_timestamp_field_for_scope_uses_v2_start_time() -> None: + """The v2 observations filter column for resume-by-timestamp is + ``startTime``; the trace's timestamp is approximated by the root + observation's start time.""" + assert BatchEvaluationRunner._get_timestamp_field_for_scope("traces") == "startTime" + assert ( + BatchEvaluationRunner._get_timestamp_field_for_scope("observations") + == "startTime" + ) + + +def _loop_stub_runner(pages: list[Any]) -> _StubRunner: + """Runner whose fetch and process layers are stubbed so the test drives + the ``run_async`` pagination loop itself.""" + runner = _StubRunner() + + async def fake_fetch(**kwargs: Any) -> Any: + return pages.pop(0) + + async def fake_process(*args: Any, **kwargs: Any) -> Any: + return (0, 0, 0, []) + + runner._fetch_batch_with_retry = fake_fetch # type: ignore[method-assign] + runner._process_batch_evaluation_item = fake_process # type: ignore[method-assign] + return runner + + +@pytest.mark.asyncio +async def test_run_async_reports_done_when_max_items_reached_on_last_page() -> None: + """Scenario: the server returns ``cursor=None`` on the page where + ``max_items`` is reached. The run must report ``completed=True`` and + ``has_more_items=False`` — not claim more work exists.""" + + obs_page = [_obs(id="o1", trace_id="t1")] + runner = _loop_stub_runner([(obs_page, None)]) + + result = await runner.run_async( + scope="observations", + mapper=lambda **kw: None, + evaluators=[], + max_items=1, + ) + + assert result.total_items_fetched == 1 + assert result.completed is True + assert result.has_more_items is False + + +@pytest.mark.asyncio +async def test_run_async_reports_more_when_max_items_reached_mid_stream() -> None: + """Scenario: ``max_items`` is reached while the server still has pages + (``cursor`` is set). The run reports ``has_more_items=True`` so callers + know there is remaining work.""" + + page1 = [_obs(id="o1", trace_id="t1")] + runner = _loop_stub_runner([(page1, "next-cursor")]) + + result = await runner.run_async( + scope="observations", + mapper=lambda **kw: None, + evaluators=[], + max_items=1, + ) + + assert result.completed is True + assert result.has_more_items is True + + +@pytest.mark.asyncio +async def test_run_async_completed_on_empty_first_page() -> None: + """Scenario: no items at all. The run completes with zero processed items + and ``completed=True``.""" + runner = _loop_stub_runner([([], None)]) + + result = await runner.run_async( + scope="observations", + mapper=lambda **kw: None, + evaluators=[], + ) + + assert result.total_items_processed == 0 + assert result.completed is True + assert result.has_more_items is False + + +@pytest.mark.asyncio +async def test_run_async_clears_seen_trace_ids_between_runs() -> None: + """``_seen_trace_ids`` must be reset at the start of every ``run_async`` + call so consecutive runs on the same runner instance do not skip traces.""" + + runner = _loop_stub_runner( + [ + ([_obs(id="root-a", trace_id="ta", is_root=True)], None), + ([_obs(id="root-a2", trace_id="ta", is_root=True)], None), + ] + ) + + first = await runner.run_async( + scope="traces", + mapper=lambda **kw: None, + evaluators=[], + ) + assert first.total_items_fetched == 1 + + second = await runner.run_async( + scope="traces", + mapper=lambda **kw: None, + evaluators=[], + ) + # The second run sees the same trace again — it must not be filtered out + # by state left over from the first run. + assert second.total_items_fetched == 1 + + +@pytest.mark.asyncio +async def test_run_async_flushes_before_returning_resume_token() -> None: + """When a fetch fails after retries, scores already created for earlier + pages must be flushed before the early return, not left in the buffer.""" + + async def failing_fetch(**kwargs: Any) -> Any: + raise RuntimeError("fetch exploded") + + runner = _StubRunner() + runner._process_batch_evaluation_item = ( # type: ignore[method-assign] + AsyncMock(return_value=(0, 0, 0, [])) + ) + runner._fetch_batch_with_retry = failing_fetch # type: ignore[method-assign] + + result = await runner.run_async( + scope="observations", + mapper=lambda **kw: None, + evaluators=[], + max_retries=1, + ) + + assert result.completed is False + assert result.resume_token is not None + runner.client.flush.assert_called()