diff --git a/.github/langfuse-v4-dual-write.override.yml b/.github/langfuse-v4-dual-write.override.yml deleted file mode 100644 index 48ab6ff48..000000000 --- a/.github/langfuse-v4-dual-write.override.yml +++ /dev/null @@ -1,7 +0,0 @@ -services: - langfuse-worker: - environment: - LANGFUSE_MIGRATION_V4_WRITE_MODE: dual - langfuse-web: - environment: - LANGFUSE_MIGRATION_V4_WRITE_MODE: dual diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 104cda66f..ee355e49b 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -159,8 +159,6 @@ jobs: LANGFUSE_SERVER_SHA="$(git ls-remote https://github.com/langfuse/langfuse.git HEAD | cut -f1)" curl -fsSL "https://raw.githubusercontent.com/langfuse/langfuse/${LANGFUSE_SERVER_SHA}/docker-compose.yml" \ -o ./langfuse-server/docker-compose.yml - cp ./.github/langfuse-v4-dual-write.override.yml \ - ./langfuse-server/docker-compose.override.yml echo "${LANGFUSE_SERVER_SHA}" - name: Run langfuse server diff --git a/langfuse/_client/attributes.py b/langfuse/_client/attributes.py index 43a85c2fd..e7027a192 100644 --- a/langfuse/_client/attributes.py +++ b/langfuse/_client/attributes.py @@ -32,8 +32,6 @@ class LangfuseOtelSpanAttributes: TRACE_TAGS = "langfuse.trace.tags" TRACE_PUBLIC = "langfuse.trace.public" TRACE_METADATA = "langfuse.trace.metadata" - TRACE_INPUT = "langfuse.trace.input" - TRACE_OUTPUT = "langfuse.trace.output" # Langfuse-observation attributes OBSERVATION_TYPE = "langfuse.observation.type" @@ -76,13 +74,9 @@ class LangfuseOtelSpanAttributes: def create_trace_attributes( *, - input: Optional[Any] = None, - output: Optional[Any] = None, public: Optional[bool] = None, ) -> dict: attributes = { - LangfuseOtelSpanAttributes.TRACE_INPUT: _serialize(input), - LangfuseOtelSpanAttributes.TRACE_OUTPUT: _serialize(output), LangfuseOtelSpanAttributes.TRACE_PUBLIC: public, } diff --git a/langfuse/_client/client.py b/langfuse/_client/client.py index dbcb8c9ee..fe3b09426 100644 --- a/langfuse/_client/client.py +++ b/langfuse/_client/client.py @@ -38,7 +38,6 @@ _agnosticcontextmanager, ) from packaging.version import Version -from typing_extensions import deprecated from langfuse._client.attributes import ( LangfuseOtelSpanAttributes, @@ -213,7 +212,7 @@ class Langfuse: release (Optional[str]): Release version/hash of your application. Used for grouping analytics by release. media_upload_thread_count (Optional[int]): Number of background threads for handling media uploads. Defaults to 1. Can also be set via LANGFUSE_MEDIA_UPLOAD_THREAD_COUNT environment variable. sample_rate (Optional[float]): Sampling rate for traces (0.0 to 1.0). Defaults to 1.0 (100% of traces are sampled). Can also be set via LANGFUSE_SAMPLE_RATE environment variable. - mask (Optional[MaskFunction]): Function to mask sensitive data synchronously when Langfuse SDK attributes are created. This applies only to data set through Langfuse SDK APIs such as `start_observation()`, `update()`, and `set_trace_io()`. + mask (Optional[MaskFunction]): Function to mask sensitive data synchronously when Langfuse SDK attributes are created. This applies only to data set through Langfuse SDK APIs such as `start_observation()` and `update()`. mask_otel_spans (Optional[MaskOtelSpansFunction]): Synchronous export-stage hook for masking raw OpenTelemetry span attributes before this Langfuse client sends them to Langfuse. Use this for spans created by third-party OpenTelemetry instrumentations, or when you need to inspect final span attributes after export filtering and Langfuse media handling. It does not modify spans already exported through other OpenTelemetry exporters. The hook receives one OpenTelemetry export batch. A batch is not guaranteed to contain a complete trace, request, or Langfuse observation tree. The hook usually runs on the OpenTelemetry batch span processor worker thread; during `flush()` and shutdown it may run on the caller thread. Keep it synchronous, deterministic, and fast. @@ -1534,55 +1533,6 @@ def update_current_span( status_message=status_message, ) - @deprecated( - "Trace-level input/output is deprecated. " - "For trace attributes (user_id, session_id, tags, etc.), use propagate_attributes() instead. " - "This method will be removed in a future major version." - ) - def set_current_trace_io( - self, - *, - input: Optional[Any] = None, - output: Optional[Any] = None, - ) -> None: - """Set trace-level input and output for the current span's trace. - - .. deprecated:: - This is a legacy method for backward compatibility with Langfuse platform - features that still rely on trace-level input/output (e.g., legacy LLM-as-a-judge - evaluators). It will be removed in a future major version. - - For setting other trace attributes (user_id, session_id, metadata, tags, version), - use :func:`langfuse.propagate_attributes` (top-level import) instead. - - Args: - input: Input data to associate with the trace. - output: Output data to associate with the trace. - """ - if not self._tracing_enabled: - langfuse_logger.debug( - "Operation skipped: set_current_trace_io - Tracing is disabled or client is in no-op mode." - ) - return - - current_otel_span = self._get_current_otel_span() - - if current_otel_span is not None and current_otel_span.is_recording(): - span_class = self._get_span_class( - self._get_observation_type_from_otel_span(current_otel_span) - ) - span = span_class( - otel_span=current_otel_span, - langfuse_client=self, - environment=self._environment, - release=self._release, - ) - - span.set_trace_io( - input=input, - output=output, - ) - def set_current_trace_as_public(self) -> None: """Make the current trace publicly accessible via its URL. @@ -3203,11 +3153,10 @@ def _create_experiment_run_name( def run_batched_evaluation( self, *, - scope: Literal["traces", "observations"], mapper: MapperFunction, filter: Optional[str] = None, fetch_batch_size: int = 50, - fetch_trace_fields: Optional[str] = None, + fields: Optional[str] = "core,basic,io,metadata", max_items: Optional[int] = None, max_retries: int = 3, evaluators: List[EvaluatorFunction], @@ -3215,16 +3164,15 @@ def run_batched_evaluation( max_concurrency: int = 5, metadata: Optional[Dict[str, Any]] = None, _add_observation_scores_to_trace: bool = False, - _additional_trace_tags: Optional[List[str]] = None, resume_from: Optional[BatchEvaluationResumeToken] = None, verbose: bool = False, ) -> BatchEvaluationResult: - """Fetch traces or observations using legacy read APIs and evaluate each item. + """Fetch observations from Langfuse and evaluate each of them. 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 - evaluators on each item, and creates scores that are linked back to the original - entities. This is ideal for: + It fetches observations based on filters, transforms them using a mapper function, + runs evaluators on each item, and creates scores that are linked back to the + original entities. This is ideal for: - Running evaluations on production traces after deployment - Backtesting new evaluation metrics on historical data @@ -3235,46 +3183,51 @@ 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. + Items are read from `GET /api/public/v2/observations` with cursor pagination, + newest first (by start time). Args: - scope: The type of items to evaluate. Must be one of: - - "traces": Evaluate complete traces with all their observations - - "observations": Evaluate individual observations (spans, generations, events) - mapper: Function that transforms API response objects into evaluator inputs. - Receives a trace/observation object and returns an EvaluatorInputs + mapper: Function that transforms an `ObservationV2` into evaluator inputs. + Called as `mapper(item=observation)` and must return an EvaluatorInputs instance with input, output, expected_output, and metadata fields. - Can be sync or async. + `input`/`output` are raw strings (not JSON-parsed), and only the + requested `fields` groups are populated. Can be sync or async. evaluators: List of evaluation functions to run on each item. Each evaluator receives the mapped inputs and returns Evaluation object(s). Evaluator failures are logged but don't stop the batch evaluation. - filter: Optional JSON filter string for querying items (same format as Langfuse API). Examples: - - '{"tags": ["production"]}' - - '{"user_id": "user123", "timestamp": {"operator": ">", "value": "2024-01-01"}}' - Default: None (fetches all items). + filter: Optional JSON array of filter conditions in the v2 observations + filter format, for example: + - '[{"type": "arrayOptions", "column": "tags", "operator": "any of", "value": ["production"]}]' + - '[{"type": "string", "column": "traceName", "operator": "=", "value": "chat"}]' + - '[{"type": "datetime", "column": "startTime", "operator": ">=", "value": "2026-01-01T00:00:00Z"}]' + Default: None (fetches all observations). 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'. + Larger values may be faster but use more memory. Maximum 1000. Default: 50. + fields: Comma-separated list of observation field groups to fetch. Available + groups: 'core' (always included), 'basic', 'time', 'io', 'metadata', + 'model', 'usage', 'prompt', 'metrics', 'trace_context'. Fields of groups + that are not requested are None on the item passed to the mapper. + Metadata values longer than 200 characters are truncated. + Default: "core,basic,io,metadata". 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). max_concurrency: Maximum number of items to evaluate concurrently. Controls parallelism and resource usage. Default: 5. composite_evaluator: Optional function that creates a composite score from - item-level evaluations. Receives the original item and its evaluations, - returns a single Evaluation. Useful for weighted averages or combined metrics. - Default: None. + item-level evaluations. Receives the mapped inputs and the item's + evaluations, returns Evaluation(s). Useful for weighted averages or + combined metrics. Default: None. metadata: Optional metadata dict to add to all created scores. Useful for tracking evaluation runs, versions, or other context. Default: None. max_retries: Maximum number of retry attempts for failed batch fetches. - Uses exponential backoff (1s, 2s, 4s). Default: 3. + Default: 3. verbose: If True, logs progress information to console. Useful for monitoring long-running evaluations. Default: False. - resume_from: Optional resume token from a previous incomplete run. Allows - continuing evaluation after interruption or failure. Default: None. + resume_from: Optional resume token from a previous run that stopped early + (fetch failure or `max_items`). Continues exactly after the last + processed page. Pass the same `filter`, or omit it to reuse the + token's filter. Default: None. Returns: @@ -3286,134 +3239,112 @@ def run_batched_evaluation( - total_composite_scores_created: Scores created by composite evaluator - total_evaluations_failed: Individual evaluator failures - evaluator_stats: Per-evaluator statistics (success rate, scores created) - - resume_token: Token for resuming if incomplete (None if completed) - - completed: True if all items processed + - resume_token: Token for continuing the run (set after a fetch failure + or when max_items was reached while more items exist) + - completed: False if the run stopped because a fetch failed - duration_seconds: Total execution time - - failed_item_ids: IDs of items that failed + - failed_item_ids: Observation IDs of items that failed - error_summary: Error types and counts - has_more_items: True if max_items reached but more exist + - item_evaluations: Evaluations per observation ID Raises: - ValueError: If invalid scope is provided. + ValueError: If a non-array filter is provided, or the resume token was + created for a different filter. Examples: - Basic trace evaluation: + Evaluate production observations: ```python from langfuse import Langfuse, EvaluatorInputs, Evaluation client = Langfuse() - # Define mapper to extract fields from traces - def trace_mapper(trace): + def simple_mapper(*, item): return EvaluatorInputs( - input=trace.input, - output=trace.output, + input=item.input, + output=item.output, expected_output=None, - metadata={"trace_id": trace.id} + metadata={"trace_id": item.trace_id}, ) - # Define evaluator def length_evaluator(*, input, output, expected_output, metadata): return Evaluation( name="output_length", value=len(output) if output else 0 ) - # Run batch evaluation result = client.run_batched_evaluation( - scope="traces", - mapper=trace_mapper, + mapper=simple_mapper, evaluators=[length_evaluator], - filter='{"tags": ["production"]}', + filter='[{"type": "arrayOptions", "column": "tags", "operator": "any of", "value": ["production"]}]', max_items=1000, verbose=True ) - print(f"Processed {result.total_items_processed} traces") + print(f"Processed {result.total_items_processed} observations") print(f"Created {result.total_scores_created} scores") ``` - Evaluation with composite scorer: + Evaluate generations with a composite scorer: ```python + import json + + def generation_mapper(*, item): + return EvaluatorInputs( + input=json.loads(item.input) if item.input else None, + output=item.output, + expected_output=None, + metadata={"model": item.model}, + ) + def accuracy_evaluator(*, input, output, expected_output, metadata): - # ... evaluation logic return Evaluation(name="accuracy", value=0.85) def relevance_evaluator(*, input, output, expected_output, metadata): - # ... evaluation logic return Evaluation(name="relevance", value=0.92) - def composite_evaluator(*, item, evaluations): - # Weighted average of evaluations + def composite_evaluator(*, input, output, expected_output, metadata, evaluations): weights = {"accuracy": 0.6, "relevance": 0.4} total = sum( e.value * weights.get(e.name, 0) for e in evaluations if isinstance(e.value, (int, float)) ) - return Evaluation( - name="composite_score", - value=total, - comment=f"Weighted average of {len(evaluations)} metrics" - ) + return Evaluation(name="composite_score", value=total) result = client.run_batched_evaluation( - scope="traces", - mapper=trace_mapper, + mapper=generation_mapper, evaluators=[accuracy_evaluator, relevance_evaluator], composite_evaluator=composite_evaluator, - filter='{"user_id": "important_user"}', - verbose=True + filter='[{"type": "string", "column": "type", "operator": "=", "value": "GENERATION"}]', + fields="core,basic,io,model", ) ``` - Handling incomplete runs with resume: + Continuing a run that stopped early: ```python - # Initial run that may fail or timeout result = client.run_batched_evaluation( - scope="observations", - mapper=obs_mapper, - evaluators=[my_evaluator], + mapper=generation_mapper, + evaluators=[accuracy_evaluator], max_items=10000, - verbose=True ) - # Check if incomplete - if not result.completed and result.resume_token: - print(f"Processed {result.resume_token.items_processed} items before interruption") - - # Resume from where it left off + while result.resume_token: result = client.run_batched_evaluation( - scope="observations", - mapper=obs_mapper, - evaluators=[my_evaluator], + mapper=generation_mapper, + evaluators=[accuracy_evaluator], + max_items=10000, resume_from=result.resume_token, - verbose=True ) - - print(f"Total items processed: {result.total_items_processed}") - ``` - - Monitoring evaluator performance: - ```python - result = client.run_batched_evaluation(...) - - for stats in result.evaluator_stats: - success_rate = stats.successful_runs / stats.total_runs - print(f"{stats.name}:") - print(f" Success rate: {success_rate:.1%}") - print(f" Scores created: {stats.total_scores_created}") - - if stats.failed_runs > 0: - print(f" ⚠️ Failed {stats.failed_runs} times") ``` Note: - Evaluator failures are logged but don't stop the batch evaluation - Individual item failures are tracked but don't stop processing - - Fetch failures are retried with exponential backoff + - Fetch failures are retried up to `max_retries` times - All scores are automatically flushed to Langfuse at the end - - The resume mechanism uses timestamp-based filtering to avoid duplicates + - Resuming uses the pagination cursor, so items are neither skipped nor + evaluated twice """ runner = BatchEvaluationRunner(self) @@ -3421,18 +3352,16 @@ def composite_evaluator(*, item, evaluations): BatchEvaluationResult, run_async_safely( runner.run_async( - scope=scope, mapper=mapper, evaluators=evaluators, filter=filter, fetch_batch_size=fetch_batch_size, - fetch_trace_fields=fetch_trace_fields, + fields=fields, max_items=max_items, max_concurrency=max_concurrency, composite_evaluator=composite_evaluator, metadata=metadata, _add_observation_scores_to_trace=_add_observation_scores_to_trace, - _additional_trace_tags=_additional_trace_tags, max_retries=max_retries, verbose=verbose, resume_from=resume_from, diff --git a/langfuse/_client/span.py b/langfuse/_client/span.py index 96879499c..9e2b5b718 100644 --- a/langfuse/_client/span.py +++ b/langfuse/_client/span.py @@ -36,7 +36,6 @@ if TYPE_CHECKING: from langfuse._client.client import Langfuse -from typing_extensions import deprecated from langfuse._client.attributes import ( LangfuseOtelSpanAttributes, @@ -226,53 +225,6 @@ def end(self, *, end_time: Optional[int] = None) -> "LangfuseObservationWrapper" return self - @deprecated( - "Trace-level input/output is deprecated. " - "For trace attributes (user_id, session_id, tags, etc.), use propagate_attributes() instead. " - "This method will be removed in a future major version." - ) - def set_trace_io( - self, - *, - input: Optional[Any] = None, - output: Optional[Any] = None, - ) -> "LangfuseObservationWrapper": - """Set trace-level input and output for the trace this span belongs to. - - .. deprecated:: - This is a legacy method for backward compatibility with Langfuse platform - features that still rely on trace-level input/output (e.g., legacy LLM-as-a-judge - evaluators). It will be removed in a future major version. - - For setting other trace attributes (user_id, session_id, metadata, tags, version), - use :func:`langfuse.propagate_attributes` (top-level import) instead. - - Args: - input: Input data to associate with the trace. - output: Output data to associate with the trace. - - Returns: - The span instance for method chaining. - """ - if not self._otel_span.is_recording(): - return self - - media_processed_input = self._process_media_and_apply_mask( - data=input, field="input", span=self._otel_span - ) - media_processed_output = self._process_media_and_apply_mask( - data=output, field="output", span=self._otel_span - ) - - attributes = create_trace_attributes( - input=media_processed_input, - output=media_processed_output, - ) - - self._otel_span.set_attributes(attributes) - - return self - def set_trace_as_public(self) -> "LangfuseObservationWrapper": """Make this trace publicly accessible via its URL. diff --git a/langfuse/_client/span_exporter.py b/langfuse/_client/span_exporter.py index 661a1a26c..efbb04261 100644 --- a/langfuse/_client/span_exporter.py +++ b/langfuse/_client/span_exporter.py @@ -27,7 +27,6 @@ _INPUT_MEDIA_ATTRIBUTE_KEYS = frozenset( { - LangfuseOtelSpanAttributes.TRACE_INPUT, LangfuseOtelSpanAttributes.OBSERVATION_INPUT, "ai.prompt.messages", "ai.prompt", @@ -54,7 +53,6 @@ _OUTPUT_MEDIA_ATTRIBUTE_KEYS = frozenset( { - LangfuseOtelSpanAttributes.TRACE_OUTPUT, LangfuseOtelSpanAttributes.OBSERVATION_OUTPUT, "ai.response.text", "ai.result.text", diff --git a/langfuse/batch_evaluation.py b/langfuse/batch_evaluation.py index 723b45757..a761f67ce 100644 --- a/langfuse/batch_evaluation.py +++ b/langfuse/batch_evaluation.py @@ -1,9 +1,10 @@ """Batch evaluation functionality for Langfuse. This module provides comprehensive batch evaluation capabilities for running evaluations -on traces and observations fetched from Langfuse. It includes type definitions, -protocols, result classes, and the implementation for large-scale evaluation workflows -with error handling, retry logic, and resume capability. +on observations fetched from Langfuse via the v2 observations API +(`GET /api/public/v2/observations`). It includes type definitions, protocols, result +classes, and the implementation for large-scale evaluation workflows with error +handling, retry logic, and resume capability. """ import asyncio @@ -20,35 +21,36 @@ Set, Tuple, Union, - cast, ) -from langfuse.api import ( - ObservationsView, - TraceWithFullDetails, -) +from langfuse.api import ObservationV2 from langfuse.experiment import Evaluation, EvaluatorFunction from langfuse.logger import langfuse_logger as logger if TYPE_CHECKING: from langfuse._client.client import Langfuse +DEFAULT_BATCH_EVALUATION_FIELDS = "core,basic,io,metadata" +"""Default v2 observation field groups fetched for batch evaluation. + +Includes ``io`` (raw ``input``/``output`` strings) and ``metadata`` so that mappers +work without extra configuration. Available groups: core, basic, time, io, +metadata, model, usage, prompt, metrics, trace_context. +""" + class EvaluatorInputs: """Input data structure for evaluators, returned by mapper functions. - This class provides a strongly-typed container for transforming API response - objects (traces, observations) into the standardized format expected - by evaluator functions. It ensures consistent access to input, output, expected - output, and metadata regardless of the source entity type. + This class provides a strongly-typed container for transforming `ObservationV2` + objects returned by the v2 observations API into the standardized format + expected by evaluator functions. Attributes: - input: The input data that was provided to generate the output being evaluated. - For traces, this might be the initial prompt or request. For observations, - this could be the span's input. The exact meaning depends on your use case. - output: The actual output that was produced and needs to be evaluated. - For traces, this is typically the final response. For observations, - this might be the generation output or span result. + input: The input data that was provided to generate the output being evaluated, + for example the observation's input. + output: The actual output that was produced and needs to be evaluated, + for example the generation output or span result. expected_output: Optional ground truth or expected result for comparison. Used by evaluators to assess correctness. May be None if no ground truth is available for the entity being evaluated. @@ -57,38 +59,31 @@ class EvaluatorInputs: or any other relevant data that evaluators might use. Examples: - Simple mapper for traces: + Simple observation mapper: ```python from langfuse import EvaluatorInputs - def trace_mapper(trace): + def simple_mapper(*, item): return EvaluatorInputs( - input=trace.input, - output=trace.output, + input=item.input, # raw string as returned by the API + output=item.output, expected_output=None, # No ground truth available - metadata={"user_id": trace.user_id, "tags": trace.tags} + metadata={"trace_id": item.trace_id, "user_id": item.user_id}, ) ``` - Mapper for observations extracting specific fields: + Mapper that decodes JSON input/output: ```python - def observation_mapper(observation): - # Extract input/output from observation's data - input_data = observation.input if hasattr(observation, 'input') else None - output_data = observation.output if hasattr(observation, 'output') else None + import json + def observation_mapper(*, item): return EvaluatorInputs( - input=input_data, - output=output_data, + input=json.loads(item.input) if item.input else None, + output=json.loads(item.output) if item.output else None, expected_output=None, - metadata={ - "observation_type": observation.type, - "model": observation.model, - "latency_ms": observation.end_time - observation.start_time - } + metadata={"observation_type": item.type, "name": item.name}, ) ``` - ``` Note: All arguments must be passed as keywords when instantiating this class. @@ -122,34 +117,39 @@ def __init__( class MapperFunction(Protocol): """Protocol defining the interface for mapper functions in batch evaluation. - Mapper functions transform API response objects (traces or observations) - into the standardized EvaluatorInputs format that evaluators expect. This abstraction - allows you to define how to extract and structure evaluation data from different - entity types. + Mapper functions transform `ObservationV2` objects from the v2 observations API + into the standardized EvaluatorInputs format that evaluators expect. Mapper functions must: - - Accept a single item parameter (trace, observation) + - Accept a single keyword argument `item` (an `ObservationV2`) - Return an EvaluatorInputs instance with input, output, expected_output, metadata - Can be either synchronous or asynchronous - Should handle missing or malformed data gracefully + + Notes on `ObservationV2`: + - `input` and `output` are raw strings exactly as stored; the SDK does not + parse them. Use `json.loads` in the mapper if you need structured data. + - Only the requested field groups (see the `fields` argument of + `Langfuse.run_batched_evaluation`) are populated; fields of other groups + are None. + - `metadata` values longer than 200 characters are truncated by the API. + - Price fields (`input_price`, `output_price`, `total_price`) are strings. """ def __call__( self, *, - item: Union["TraceWithFullDetails", "ObservationsView"], + item: ObservationV2, **kwargs: Dict[str, Any], ) -> Union[EvaluatorInputs, Awaitable[EvaluatorInputs]]: - """Transform an API response object into evaluator inputs. + """Transform an observation into evaluator inputs. - This method defines how to extract evaluation-relevant data from the raw - API response object. The implementation should map entity-specific fields - to the standardized input/output/expected_output/metadata structure. + This method defines how to extract evaluation-relevant data from the + observation. The implementation should map its fields to the standardized + input/output/expected_output/metadata structure. Args: - item: The API response object to transform. The type depends on the scope: - - TraceWithFullDetails: When evaluating traces - - ObservationsView: When evaluating observations + item: The `ObservationV2` to transform. Returns: EvaluatorInputs: A structured container with: @@ -162,48 +162,45 @@ def __call__( (for async mappers that need to fetch additional data). Examples: - Basic trace mapper: + Basic observation mapper: ```python - def map_trace(trace): + def map_basic(*, item): return EvaluatorInputs( - input=trace.input, - output=trace.output, + input=item.input, + output=item.output, expected_output=None, - metadata={"trace_id": trace.id, "user": trace.user_id} + metadata={"trace_id": item.trace_id, "user": item.user_id} ) ``` Observation mapper with conditional logic: ```python - def map_observation(observation): - # Extract fields based on observation type - if observation.type == "GENERATION": - input_data = observation.input - output_data = observation.output + import json + + def map_observation(*, item): + if item.type == "GENERATION": + input_data = json.loads(item.input) if item.input else None else: - # For other types, use different fields - input_data = observation.metadata.get("input") - output_data = observation.metadata.get("output") + input_data = item.input return EvaluatorInputs( input=input_data, - output=output_data, + output=item.output, expected_output=None, - metadata={"obs_id": observation.id, "type": observation.type} + metadata={"obs_id": item.id, "type": item.type} ) ``` Async mapper (if additional processing needed): ```python - async def map_trace_async(trace): - # Could do async processing here if needed - processed_output = await some_async_transformation(trace.output) + async def map_async(*, item): + processed_output = await some_async_transformation(item.output) return EvaluatorInputs( - input=trace.input, + input=item.input, output=processed_output, expected_output=None, - metadata={"trace_id": trace.id} + metadata={"trace_id": item.trace_id} ) ``` """ @@ -452,122 +449,87 @@ def __init__( class BatchEvaluationResumeToken: - """Token for resuming a failed batch evaluation run. + """Token for resuming an interrupted or limited batch evaluation run. + + The v2 observations API returns observations ordered by start time, newest + first, and paginates with an opaque cursor. The token stores the cursor of + the next page that has not been processed yet, so a resumed run continues + exactly where the previous one stopped, without re-evaluating or skipping + items, even if new observations were ingested in the meantime. - This class encapsulates all the information needed to resume a batch evaluation - that was interrupted or failed partway through. It uses timestamp-based filtering - to avoid re-processing items that were already evaluated, even if the underlying - dataset changed between runs. + A token is returned when a run stops because a batch fetch failed after all + retries (`completed=False`) or because `max_items` was reached while more + items exist (`has_more_items=True`). Attributes: - scope: The type of items being evaluated ("traces", "observations"). - filter: The original JSON filter string used to query items. - last_processed_timestamp: ISO 8601 timestamp of the last successfully processed item. - Used to construct a filter that only fetches items after this timestamp. - last_processed_id: The ID of the last successfully processed item, for reference. - items_processed: Count of items successfully processed before interruption. + filter: The original JSON filter string used to query items. Pass the + same filter when resuming. + cursor: Cursor of the next page to fetch. None if no page was fetched yet. + last_processed_timestamp: ISO 8601 start time of the oldest processed + observation. Only used to resume when `cursor` is None, by fetching + observations that started strictly before this timestamp. + last_processed_id: The ID of the last processed observation, for reference. + items_processed: Number of items successfully processed so far, including + items processed by the runs this token was resumed from. Examples: - Resuming a failed batch evaluation: + Resuming a run that stopped early: ```python - # Initial run that fails partway through - try: - result = client.run_batched_evaluation( - scope="traces", - mapper=my_mapper, - evaluators=[evaluator1, evaluator2], - filter='{"tags": ["production"]}', - max_items=10000 - ) - except Exception as e: - print(f"Evaluation failed: {e}") - - # Save the resume token - if result.resume_token: - # Store resume token for later (e.g., in a file or database) - import json - with open("resume_token.json", "w") as f: - json.dump({ - "scope": result.resume_token.scope, - "filter": result.resume_token.filter, - "last_timestamp": result.resume_token.last_processed_timestamp, - "last_id": result.resume_token.last_processed_id, - "items_done": result.resume_token.items_processed - }, f) - - # Later, resume from where it left off - with open("resume_token.json") as f: - token_data = json.load(f) - - resume_token = BatchEvaluationResumeToken( - scope=token_data["scope"], - filter=token_data["filter"], - last_processed_timestamp=token_data["last_timestamp"], - last_processed_id=token_data["last_id"], - items_processed=token_data["items_done"] - ) - - # Resume the evaluation result = client.run_batched_evaluation( - scope="traces", mapper=my_mapper, evaluators=[evaluator1, evaluator2], - resume_from=resume_token + filter=my_filter, + max_items=10000, ) - print(f"Processed {result.total_items_processed} additional items") + if result.resume_token: + result = client.run_batched_evaluation( + mapper=my_mapper, + evaluators=[evaluator1, evaluator2], + filter=my_filter, + resume_from=result.resume_token, + ) ``` - Handling partial completion: + Persisting a token between processes: ```python - result = client.run_batched_evaluation(...) + import json - if not result.completed: - print(f"Evaluation incomplete. Processed {result.resume_token.items_processed} items") - print(f"Last item: {result.resume_token.last_processed_id}") - print(f"Resume from: {result.resume_token.last_processed_timestamp}") + token = result.resume_token + with open("resume_token.json", "w") as f: + json.dump(vars(token), f) - # Optionally retry automatically - if result.resume_token: - print("Retrying...") - result = client.run_batched_evaluation( - scope=result.resume_token.scope, - mapper=my_mapper, - evaluators=my_evaluators, - resume_from=result.resume_token - ) + with open("resume_token.json") as f: + token = BatchEvaluationResumeToken(**json.load(f)) ``` Note: All arguments must be passed as keywords when instantiating this class. - The timestamp-based approach means that items created after the initial run - but before the timestamp will be skipped. This is intentional to avoid - duplicates and ensure consistent evaluation. """ def __init__( self, *, - scope: str, filter: Optional[str], last_processed_timestamp: str, last_processed_id: str, items_processed: int, + cursor: Optional[str] = None, ): """Initialize BatchEvaluationResumeToken with the provided state. Args: - scope: The scope type ("traces", "observations"). filter: The original JSON filter string. - last_processed_timestamp: ISO 8601 timestamp of last processed item. + last_processed_timestamp: ISO 8601 start time of the oldest processed item. last_processed_id: ID of last processed item. items_processed: Count of items processed before interruption. + cursor: Cursor of the next page to fetch. Note: All arguments must be provided as keywords. """ - self.scope = scope self.filter = filter + self.cursor = cursor self.last_processed_timestamp = last_processed_timestamp self.last_processed_id = last_processed_id self.items_processed = items_processed @@ -588,13 +550,16 @@ class BatchEvaluationResult: total_composite_scores_created: Scores created by the composite evaluator. total_evaluations_failed: Number of individual evaluator failures across all items. evaluator_stats: List of per-evaluator statistics (success/failure rates, scores created). - resume_token: Token for resuming if evaluation was interrupted (None if completed). - completed: True if all items were processed, False if stopped early or failed. + resume_token: Token for continuing the run. Set when a batch fetch failed + (`completed=False`) or when `max_items` was reached while more items + exist (`has_more_items=True`); None otherwise. + completed: False if the run stopped because a batch fetch failed, True otherwise + (including when it stopped at `max_items`). duration_seconds: Total time taken to execute the batch evaluation. - failed_item_ids: List of IDs for items that failed evaluation. + failed_item_ids: List of observation IDs for items that failed evaluation. error_summary: Dictionary mapping error types to occurrence counts. has_more_items: True if max_items limit was reached but more items exist. - item_evaluations: Dictionary mapping item IDs to their evaluation results (both regular and composite). + item_evaluations: Dictionary mapping observation IDs to their evaluation results (both regular and composite). Examples: Basic result inspection: @@ -644,9 +609,7 @@ class BatchEvaluationResult: if result.resume_token: print(f"Processed {result.resume_token.items_processed} items before failure") - print(f"Use resume_from parameter to continue from:") - print(f" Timestamp: {result.resume_token.last_processed_timestamp}") - print(f" Last ID: {result.resume_token.last_processed_id}") + print("Pass it as resume_from to continue") if result.has_more_items: print(f"ℹ️ More items available beyond max_items limit") @@ -839,57 +802,65 @@ def __init__(self, client: "Langfuse"): async def run_async( self, *, - scope: str, mapper: MapperFunction, evaluators: List[EvaluatorFunction], filter: Optional[str] = None, fetch_batch_size: int = 50, - fetch_trace_fields: Optional[str] = "io", + fields: Optional[str] = DEFAULT_BATCH_EVALUATION_FIELDS, max_items: Optional[int] = None, max_concurrency: int = 5, composite_evaluator: Optional[CompositeEvaluatorFunction] = None, metadata: Optional[Dict[str, Any]] = None, _add_observation_scores_to_trace: bool = False, - _additional_trace_tags: Optional[List[str]] = None, max_retries: int = 3, verbose: bool = False, resume_from: Optional[BatchEvaluationResumeToken] = None, ) -> BatchEvaluationResult: - """Run batch evaluation asynchronously using legacy read APIs. + """Run batch evaluation asynchronously. 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. + evaluation process: fetching items from `GET /api/public/v2/observations` + with cursor pagination, mapping, evaluating, creating scores, and tracking + statistics. Args: - scope: The type of items to evaluate ("traces", "observations"). - mapper: Function to transform API response items to evaluator inputs. + mapper: Function to transform `ObservationV2` items to evaluator inputs. 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' + filter: JSON filter string (v2 observations filter schema). + fetch_batch_size: Number of items to fetch per API call (max 1000). + fields: Comma-separated v2 observation field groups to fetch. 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. metadata: Metadata to add to all created scores. _add_observation_scores_to_trace: Private option to duplicate observation-level scores onto the parent trace. - _additional_trace_tags: Private option to add tags on traces via - ingestion trace-create events. max_retries: Maximum retries for failed batch fetches. verbose: If True, log progress to console. - resume_from: Resume token from a previous failed run. + resume_from: Resume token from a previous run. If `filter` is omitted, + the token's filter is reused. Returns: BatchEvaluationResult with comprehensive statistics. + + Raises: + ValueError: If the filter is not a JSON array, or the resume token was + created for a different filter. """ start_time = time.time() - # Initialize tracking variables + # The token's cursor only continues the query that produced it. + if resume_from is not None: + if filter is None: + filter = resume_from.filter + elif filter != resume_from.filter: + raise ValueError( + "Resume token was created for a different filter. Pass the " + "same filter, or omit it to reuse the token's filter." + ) + + effective_filter = self._build_filter(filter=filter, resume_from=resume_from) + total_items_fetched = 0 total_items_processed = 0 total_items_failed = 0 @@ -899,8 +870,8 @@ async def run_async( failed_item_ids: List[str] = [] error_summary: Dict[str, int] = {} item_evaluations: Dict[str, List[Evaluation]] = {} + previously_processed = resume_from.items_processed if resume_from else 0 - # Initialize evaluator stats evaluator_stats_dict = { getattr(evaluator, "__name__", "unknown_evaluator"): EvaluatorStats( name=getattr(evaluator, "__name__", "unknown_evaluator") @@ -908,65 +879,69 @@ async def run_async( for evaluator in evaluators } - # Handle resume token by modifying filter - effective_filter = self._build_timestamp_filter(filter, resume_from) - normalized_additional_trace_tags = ( - self._dedupe_tags(_additional_trace_tags) - if _additional_trace_tags is not None - else [] - ) - updated_trace_ids: Set[str] = set() - - # Create semaphore for concurrency control semaphore = asyncio.Semaphore(max_concurrency) - # Pagination state - page = 1 + cursor: Optional[str] = resume_from.cursor if resume_from else None has_more = True - last_item_timestamp: Optional[str] = None - last_item_id: Optional[str] = None + last_item_timestamp = ( + resume_from.last_processed_timestamp if resume_from else "" + ) + last_item_id = resume_from.last_processed_id if resume_from else "" + # The start-time fallback is inclusive so tied items are not lost; skip + # the one item the token says was already processed. + resumed_item_id = ( + resume_from.last_processed_id + if resume_from is not None + and resume_from.cursor is None + and resume_from.last_processed_timestamp + else None + ) + previous_page_ids: Set[Optional[str]] = set() + batch_number = 0 if verbose: - logger.info("Starting batch evaluation on %s", scope) - if scope == "traces" and fetch_trace_fields: - logger.info("Fetching trace fields: %s", fetch_trace_fields) + logger.info("Starting batch evaluation on observations") + if fields: + logger.info("Fetching observation fields: %s", fields) if resume_from: logger.info( - "Resuming from %s (%s items already processed)", - resume_from.last_processed_timestamp, + "Resuming after %s (%s items already processed)", + resume_from.last_processed_timestamp or "start", resume_from.items_processed, ) - # Main pagination loop - while has_more: - # Check if we've reached max_items + def build_resume_token() -> BatchEvaluationResumeToken: + return BatchEvaluationResumeToken( + filter=filter, + cursor=cursor, + last_processed_timestamp=last_item_timestamp, + last_processed_id=last_item_id, + items_processed=previously_processed + total_items_processed, + ) + + while True: if max_items is not None and total_items_fetched >= max_items: if verbose: logger.info("Reached max_items limit (%s)", max_items) - has_more = True # More items may exist break - # Fetch next batch with retry logic + # Clamping the page size keeps the cursor aligned with the last + # processed item, so resuming after max_items skips nothing. + limit = fetch_batch_size + if max_items is not None: + limit = min(limit, max_items - total_items_fetched) + try: - items = await self._fetch_batch_with_retry( - scope=scope, + items, next_cursor = await self._fetch_batch_with_retry( filter=effective_filter, - page=page, - limit=fetch_batch_size, + cursor=cursor, + limit=limit, max_retries=max_retries, - fields=fetch_trace_fields, + fields=fields, ) except Exception as e: - # Failed after max_retries - create resume token and return - error_msg = f"Failed to fetch batch after {max_retries} retries" - logger.error("%s: %s", error_msg, e) - - resume_token = BatchEvaluationResumeToken( - scope=scope, - filter=filter, # Original filter, not modified - last_processed_timestamp=last_item_timestamp or "", - last_processed_id=last_item_id or "", - items_processed=total_items_processed, + logger.error( + "Failed to fetch batch after %s retries: %s", max_retries, e ) return self._build_result( @@ -977,51 +952,28 @@ async def run_async( total_composite_scores_created=total_composite_scores_created, total_evaluations_failed=total_evaluations_failed, evaluator_stats_dict=evaluator_stats_dict, - resume_token=resume_token, + resume_token=build_resume_token(), completed=False, start_time=start_time, failed_item_ids=failed_item_ids, error_summary=error_summary, - has_more_items=has_more, + has_more_items=False, item_evaluations=item_evaluations, ) - # Check if we got any items - if not items: - has_more = False - if verbose: - logger.info("No more items to fetch") - break - + batch_number += 1 total_items_fetched += len(items) if verbose: - logger.info("Fetched batch %s (%s items)", page, len(items)) - - # Limit items if max_items would be exceeded - items_to_process = items - if max_items is not None: - remaining_capacity = max_items - total_items_processed - if len(items) > remaining_capacity: - items_to_process = items[:remaining_capacity] - if verbose: - logger.info( - "Limiting batch to %s items to respect max_items=%s", - len(items_to_process), - max_items, - ) + logger.info("Fetched batch %s (%s items)", batch_number, len(items)) - # Process items concurrently async def process_item( - item: Union[TraceWithFullDetails, ObservationsView], + item: ObservationV2, ) -> Tuple[str, Union[Tuple[int, int, int, List[Evaluation]], Exception]]: - """Process a single item and return (item_id, result).""" async with semaphore: - item_id = self._get_item_id(item, scope) try: result = await self._process_batch_evaluation_item( item=item, - scope=scope, mapper=mapper, evaluators=evaluators, composite_evaluator=composite_evaluator, @@ -1029,25 +981,27 @@ async def process_item( _add_observation_scores_to_trace=_add_observation_scores_to_trace, evaluator_stats_dict=evaluator_stats_dict, ) - return (item_id, result) + return (item.id, result) except Exception as e: - return (item_id, e) + return (item.id, e) - # Run all items in batch concurrently - tasks = [process_item(item) for item in items_to_process] - results = await asyncio.gather(*tasks) + items_to_process = self._deduplicate_page( + items, skip_ids=previous_page_ids | {resumed_item_id} + ) + previous_page_ids = {item.id for item in items} + + results = await asyncio.gather( + *[process_item(item) for item in items_to_process] + ) - # Process results and update statistics for item, (item_id, result) in zip(items_to_process, results): if isinstance(result, Exception): - # Item processing failed total_items_failed += 1 failed_item_ids.append(item_id) error_type = type(result).__name__ error_summary[error_type] = error_summary.get(error_type, 0) + 1 logger.warning("Item %s failed: %s", item_id, result) else: - # Item processed successfully total_items_processed += 1 scores_created, composite_created, evals_failed, evaluations = ( result @@ -1055,36 +1009,19 @@ async def process_item( total_scores_created += scores_created total_composite_scores_created += composite_created total_evaluations_failed += evals_failed - - # Store evaluations for this item item_evaluations[item_id] = evaluations - if normalized_additional_trace_tags: - trace_id = ( - item_id - if scope == "traces" - else cast(ObservationsView, item).trace_id - ) - - if trace_id and trace_id not in updated_trace_ids: - self.client._create_trace_tags_via_ingestion( - trace_id=trace_id, - tags=normalized_additional_trace_tags, - ) - updated_trace_ids.add(trace_id) - - # Update last processed tracking - last_item_timestamp = self._get_item_timestamp(item, scope) - last_item_id = item_id + if items: + last_item_timestamp = items[-1].start_time.isoformat() + last_item_id = items[-1].id if verbose: if max_items is not None and max_items > 0: - progress_pct = total_items_processed / max_items * 100 logger.info( "Progress: %s/%s items (%.1f%%), %s scores created", total_items_processed, max_items, - progress_pct, + total_items_processed / max_items * 100, total_scores_created, ) else: @@ -1094,40 +1031,22 @@ 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 + cursor = next_cursor + if cursor is None or not items: 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: - has_more = True # More items exist but we're stopping - break + break - # Flush all scores to Langfuse if verbose: logger.info("Flushing scores to Langfuse...") self.client.flush() - # Build final result - duration = time.time() - start_time - if verbose: logger.info( "Batch evaluation complete: %s items processed in %.2fs", total_items_processed, - duration, + time.time() - start_time, ) - # Completed successfully if we either: - # 1. Ran out of items (has_more is False), OR - # 2. Hit max_items limit (intentionally stopped) - completed_successfully = not has_more or ( - max_items is not None and total_items_fetched >= max_items - ) - return self._build_result( total_items_fetched=total_items_fetched, total_items_processed=total_items_processed, @@ -1136,69 +1055,54 @@ async def process_item( total_composite_scores_created=total_composite_scores_created, total_evaluations_failed=total_evaluations_failed, evaluator_stats_dict=evaluator_stats_dict, - resume_token=None, # No resume needed on successful completion - completed=completed_successfully, + resume_token=build_resume_token() if has_more else None, + completed=True, start_time=start_time, failed_item_ids=failed_item_ids, error_summary=error_summary, - has_more_items=( - has_more and max_items is not None and total_items_fetched >= max_items - ), + has_more_items=has_more, item_evaluations=item_evaluations, ) async def _fetch_batch_with_retry( self, *, - 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[ObservationV2], Optional[str]]: + """Fetch one page of observations from the v2 observations API. 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: Cursor returned with the previous page; None for the first page. + limit: Number of items to request. max_retries: Maximum number of retry attempts. - verbose: Whether to log retry attempts. - fields: Trace fields to fetch + fields: Comma-separated v2 field groups to include. Returns: - List of items from the API. + Tuple of the page's observations and the cursor for the next page + (None when there are no more pages). Raises: 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: - error_message = f"Invalid scope: {scope}" - raise ValueError(error_message) + response = await asyncio.to_thread( + self.client.api.observations.get_many, + fields=fields, + limit=limit, + cursor=cursor, + filter=filter, + request_options={"max_retries": max_retries}, + ) + + return list(response.data), response.meta.cursor async def _process_batch_evaluation_item( self, - item: Union[TraceWithFullDetails, ObservationsView], - scope: str, + item: ObservationV2, mapper: MapperFunction, evaluators: List[EvaluatorFunction], composite_evaluator: Optional[CompositeEvaluatorFunction], @@ -1209,8 +1113,7 @@ async def _process_batch_evaluation_item( """Process a single item: map, evaluate, create scores. Args: - item: The API response object to evaluate. - scope: The type of item ("traces", "observations"). + item: The observation to evaluate. mapper: Function to transform item to evaluator inputs. evaluators: List of evaluator functions. composite_evaluator: Optional composite evaluator function. @@ -1225,14 +1128,16 @@ async def _process_batch_evaluation_item( Raises: Exception: If mapping fails or item processing encounters fatal error. """ + if not item.trace_id: + message = f"Observation {item.id} has no trace_id" + raise ValueError(message) + scores_created = 0 composite_scores_created = 0 evaluations_failed = 0 - # Run mapper to transform item evaluator_inputs = await self._run_mapper(mapper, item) - # Run all evaluators evaluations: List[Evaluation] = [] for evaluator in evaluators: evaluator_name = getattr(evaluator, "__name__", "unknown_evaluator") @@ -1253,31 +1158,23 @@ async def _process_batch_evaluation_item( evaluations.extend(eval_results) except Exception as e: - # Evaluator failed - log warning and continue with other evaluators stats.failed_runs += 1 evaluations_failed += 1 logger.warning( "Evaluator %s failed on item %s: %s", evaluator_name, - self._get_item_id(item, scope), + item.id, e, ) - # Create scores for item-level evaluations - item_id = self._get_item_id(item, scope) for evaluation in evaluations: - scores_created += self._create_score_for_scope( - scope=scope, - item_id=item_id, - trace_id=cast(ObservationsView, item).trace_id - if scope == "observations" - else None, + scores_created += self._create_score( + item=item, evaluation=evaluation, additional_metadata=metadata, add_observation_score_to_trace=_add_observation_scores_to_trace, ) - # Run composite evaluator if provided and we have evaluations if composite_evaluator and evaluations: try: composite_evals = await self._run_composite_evaluator( @@ -1289,24 +1186,18 @@ async def _process_batch_evaluation_item( evaluations=evaluations, ) - # Create scores for all composite evaluations for composite_eval in composite_evals: - composite_scores_created += self._create_score_for_scope( - scope=scope, - item_id=item_id, - trace_id=cast(ObservationsView, item).trace_id - if scope == "observations" - else None, + composite_scores_created += self._create_score( + item=item, evaluation=composite_eval, additional_metadata=metadata, add_observation_score_to_trace=_add_observation_scores_to_trace, ) - # Add composite evaluations to the list evaluations.extend(composite_evals) except Exception as e: - logger.warning("Composite evaluator failed on item %s: %s", item_id, e) + logger.warning("Composite evaluator failed on item %s: %s", item.id, e) return ( scores_created, @@ -1337,11 +1228,9 @@ async def _run_evaluator_internal( """ result = evaluator(**kwargs) - # Handle async evaluators if asyncio.iscoroutine(result): result = await result - # Normalize to list if isinstance(result, (dict, Evaluation)): return [result] # type: ignore elif isinstance(result, list): @@ -1352,13 +1241,13 @@ async def _run_evaluator_internal( async def _run_mapper( self, mapper: MapperFunction, - item: Union[TraceWithFullDetails, ObservationsView], + item: ObservationV2, ) -> EvaluatorInputs: """Run mapper function (handles both sync and async mappers). Args: mapper: The mapper function to run. - item: The API response object to map. + item: The observation to map. Returns: EvaluatorInputs instance. @@ -1406,7 +1295,6 @@ async def _run_composite_evaluator( if asyncio.iscoroutine(result): result = await result - # Normalize to list (same as regular evaluator) if isinstance(result, (dict, Evaluation)): return [result] # type: ignore elif isinstance(result, list): @@ -1414,22 +1302,18 @@ async def _run_composite_evaluator( else: return [] - def _create_score_for_scope( + def _create_score( self, *, - scope: str, - item_id: str, - trace_id: Optional[str] = None, + item: ObservationV2, evaluation: Evaluation, additional_metadata: Optional[Dict[str, Any]], add_observation_score_to_trace: bool = False, ) -> int: - """Create a score linked to the appropriate entity based on scope. + """Create a score on the evaluated observation, optionally duplicated onto its trace. Args: - scope: The type of entity ("traces", "observations"). - item_id: The ID of the entity. - trace_id: The trace ID of the entity; required if scope=observations + item: The evaluated observation. evaluation: The evaluation result to create a score from. additional_metadata: Additional metadata to merge with evaluation metadata. add_observation_score_to_trace: Whether to duplicate observation @@ -1438,165 +1322,104 @@ def _create_score_for_scope( Returns: Number of score events created. """ - # Merge metadata score_metadata = { **(evaluation.metadata or {}), **(additional_metadata or {}), } + score_kwargs: Dict[str, Any] = { + "name": evaluation.name, + "value": evaluation.value, + "comment": evaluation.comment, + "metadata": score_metadata, + "data_type": evaluation.data_type, + "config_id": evaluation.config_id, + } - if scope == "traces": - self.client.create_score( - trace_id=item_id, - name=evaluation.name, - value=evaluation.value, # type: ignore - comment=evaluation.comment, - metadata=score_metadata, - data_type=evaluation.data_type, # type: ignore[arg-type] - config_id=evaluation.config_id, - ) + self.client.create_score( + observation_id=item.id, trace_id=item.trace_id, **score_kwargs + ) + if not add_observation_score_to_trace: return 1 - elif scope == "observations": - self.client.create_score( - observation_id=item_id, - trace_id=trace_id, - name=evaluation.name, - value=evaluation.value, # type: ignore - comment=evaluation.comment, - metadata=score_metadata, - data_type=evaluation.data_type, # type: ignore[arg-type] - config_id=evaluation.config_id, - ) - score_count = 1 - - if add_observation_score_to_trace and trace_id: - self.client.create_score( - trace_id=trace_id, - name=evaluation.name, - value=evaluation.value, # type: ignore - comment=evaluation.comment, - metadata=score_metadata, - data_type=evaluation.data_type, # type: ignore[arg-type] - config_id=evaluation.config_id, - ) - score_count += 1 - return score_count + self.client.create_score(trace_id=item.trace_id, **score_kwargs) + return 2 - return 0 - - def _build_timestamp_filter( - self, - original_filter: Optional[str], + @staticmethod + def _build_filter( + *, + filter: Optional[str], resume_from: Optional[BatchEvaluationResumeToken], ) -> Optional[str]: - """Build filter with timestamp constraint for resume capability. + """Combine the user filter with the resume constraint. + + Constraints are added as JSON filter conditions rather than query + parameters because the API drops a query-parameter filter whenever the + JSON filter has a condition on the same column. Args: - original_filter: The original JSON filter string. - resume_from: Optional resume token with timestamp information. + filter: The user-provided JSON filter string (a JSON array). + resume_from: Optional resume token. Returns: - Modified filter string with timestamp constraint, or original filter. - """ - if not resume_from: - return original_filter + The JSON filter string to send, or None if there are no conditions. - # Parse original filter (should be array) or create empty array - try: - filter_list = json.loads(original_filter) if original_filter else [] - if not isinstance(filter_list, list): - logger.warning( - "Filter should be a JSON array, got: %s", type(filter_list).__name__ + Raises: + ValueError: If the filter is not a JSON array. + """ + conditions: List[Any] = [] + if filter: + try: + parsed = json.loads(filter) + except json.JSONDecodeError as e: + message = f"filter must be a JSON array: {e}" + raise ValueError(message) from e + if not isinstance(parsed, list): + message = ( + "filter must be a JSON array of conditions, " + f"got {type(parsed).__name__}" ) - filter_list = [] - except json.JSONDecodeError: - logger.warning( - "Invalid JSON in original filter, ignoring: %s", original_filter + raise ValueError(message) + conditions.extend(parsed) + + # Results are ordered by start time descending, so items that remain + # after the last processed one started before it. The cursor resumes + # exactly; the timestamp is the fallback for tokens without a cursor. + if ( + resume_from is not None + and resume_from.cursor is None + and resume_from.last_processed_timestamp + ): + conditions.append( + { + "type": "datetime", + "column": "startTime", + "operator": "<=", + "value": resume_from.last_processed_timestamp, + } ) - filter_list = [] - - # Add timestamp constraint to filter array - timestamp_field = self._get_timestamp_field_for_scope(resume_from.scope) - timestamp_filter = { - "type": "datetime", - "column": timestamp_field, - "operator": ">", - "value": resume_from.last_processed_timestamp, - } - filter_list.append(timestamp_filter) - - return json.dumps(filter_list) - - @staticmethod - def _get_item_id( - item: Union[TraceWithFullDetails, ObservationsView], - scope: str, - ) -> str: - """Extract ID from item based on scope. - Args: - item: The API response object. - scope: The type of item. - - Returns: - The item's ID. - """ - return item.id + return json.dumps(conditions) if conditions else None @staticmethod - def _get_item_timestamp( - item: Union[TraceWithFullDetails, ObservationsView], - scope: str, - ) -> str: - """Extract timestamp from item based on scope. + def _deduplicate_page( + items: List[ObservationV2], *, skip_ids: Set[Optional[str]] + ) -> List[ObservationV2]: + """Keep one row per observation id, preferring the most recently updated. - Args: - item: The API response object. - scope: The type of item. - - Returns: - ISO 8601 timestamp string. + Until ClickHouse merges them, the events table can return several rows + for one observation, within a page or across a page boundary. """ - 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] - 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. - """ - if scope == "traces": - return "timestamp" - elif scope == "observations": - return "start_time" - return "timestamp" # Default - - @staticmethod - def _dedupe_tags(tags: Optional[List[str]]) -> List[str]: - """Deduplicate tags while preserving order.""" - if tags is None: - return [] - - deduped: List[str] = [] - seen = set() - for tag in tags: - if tag not in seen: - deduped.append(tag) - seen.add(tag) - - return deduped + latest_by_id: Dict[str, ObservationV2] = {} + for item in items: + if item.id in skip_ids: + continue + current = latest_by_id.get(item.id) + if current is None or (item.updated_at or item.start_time) > ( + current.updated_at or current.start_time + ): + latest_by_id[item.id] = item + + return list(latest_by_id.values()) def _build_result( self, @@ -1625,12 +1448,12 @@ def _build_result( total_composite_scores_created: Scores from composite evaluator. total_evaluations_failed: Individual evaluator failures. evaluator_stats_dict: Per-evaluator statistics. - resume_token: Resume token if incomplete. - completed: Whether evaluation completed fully. + resume_token: Resume token if the run can be continued. + completed: Whether evaluation completed without a fetch failure. start_time: Start time (unix timestamp). failed_item_ids: IDs of failed items. error_summary: Error type counts. - has_more_items: Whether more items exist. + has_more_items: Whether more items exist beyond max_items. item_evaluations: Dictionary mapping item IDs to their evaluation results. Returns: diff --git a/tests/e2e/test_batch_evaluation.py b/tests/e2e/test_batch_evaluation.py index 0632b21b8..edbea8ebc 100644 --- a/tests/e2e/test_batch_evaluation.py +++ b/tests/e2e/test_batch_evaluation.py @@ -1,1139 +1,268 @@ -"""Comprehensive tests for batch evaluation functionality. +"""End-to-end tests for run_batched_evaluation against a Langfuse server. -This test suite covers the run_batched_evaluation method which allows evaluating -traces, observations, and sessions fetched from Langfuse with mappers, evaluators, -and composite evaluators. +Every run is restricted to a corpus seeded by this module (filtered by a unique +tag), so assertions do not depend on other data in the project. Runner logic +that does not need a server is covered in tests/unit/test_batch_evaluation.py. """ -import asyncio -import time +import json +from dataclasses import dataclass +from typing import Any, List import pytest from langfuse import get_client, propagate_attributes -from langfuse.batch_evaluation import ( - BatchEvaluationResult, - BatchEvaluationResumeToken, - EvaluatorInputs, - EvaluatorStats, -) +from langfuse.api import ObservationV2 +from langfuse.batch_evaluation import EvaluatorInputs from langfuse.experiment import Evaluation from tests.support.utils import create_uuid, get_api, wait_for_result -# ============================================================================ -# FIXTURES & SETUP -# ============================================================================ +TRACE_COUNT = 4 -# pytestmark = pytest.mark.skip(reason="Github CI runner overwhelmed by score volume") +@dataclass +class Corpus: + tag: str + filter: str + trace_ids: List[str] + root_ids: List[str] + child_ids: List[str] -@pytest.fixture -def langfuse_client(): - """Get a Langfuse client for testing.""" - return get_client() - - -@pytest.fixture -def sample_trace_name(): - """Generate a unique trace name for filtering.""" - return f"batch-eval-test-{create_uuid()}" - - -def _seed_trace_corpus( - *, trace_count: int = 6, tag: str | None = None -) -> tuple[str, list[str]]: - langfuse_client = get_client() - corpus_tag = tag or f"batch-eval-seed-{create_uuid()}" - trace_names: list[str] = [] - - for index in range(trace_count): - trace_name = f"{corpus_tag}-trace-{index}" - trace_names.append(trace_name) - with langfuse_client.start_as_current_observation(name=trace_name) as span: - with propagate_attributes(tags=[corpus_tag]): - span.set_trace_io( - input=f"Seed input {index}", - output=f"Seed output {index}", - ) +def _tag_filter(tag: str) -> str: + return json.dumps( + [ + { + "type": "arrayOptions", + "column": "tags", + "operator": "any of", + "value": [tag], + } + ] + ) - langfuse_client.flush() - filter_json = f'[{{"type": "arrayOptions", "column": "tags", "operator": "any of", "value": ["{corpus_tag}"]}}]' +def _wait_for_observations(filter_json: str, expected_count: int) -> List[Any]: api = get_api(retry=False) - wait_for_result( - lambda: api.trace.list(filter=filter_json, limit=trace_count), - is_result_ready=lambda response: len(response.data) >= trace_count, + response = wait_for_result( + lambda: api.observations.get_many(filter=filter_json, limit=100), + is_result_ready=lambda r: len(r.data) >= expected_count, ) - - return corpus_tag, trace_names - - -@pytest.fixture(scope="module", autouse=True) -def seeded_batch_evaluation_traces(): - _seed_trace_corpus() + return list(response.data) -def simple_trace_mapper(*, item): - """Simple mapper for traces.""" +def _wait_for_scores(**kwargs: Any) -> List[Any]: + api = get_api(retry=False) + response = wait_for_result( + lambda: api.scores_v3.get_many_v3(fields="details,subject", **kwargs), + is_result_ready=lambda r: len(r.data) > 0, + ) + return list(response.data) + + +@pytest.fixture(scope="module") +def corpus() -> Corpus: + langfuse = get_client() + tag = f"batch-eval-{create_uuid()}" + trace_ids, root_ids, child_ids = [], [], [] + + for index in range(TRACE_COUNT): + with langfuse.start_as_current_observation(name=f"{tag}-root") as root: + with propagate_attributes(tags=[tag]): + with langfuse.start_as_current_observation( + as_type="generation", + name=f"{tag}-child", + input={"question": index}, + output=f"child answer {index}", + ) as child: + child_ids.append(child.id) + root.update(input={"question": index}, output=f"answer {index}") + trace_ids.append(root.trace_id) + root_ids.append(root.id) + + langfuse.flush() + _wait_for_observations(_tag_filter(tag), 2 * TRACE_COUNT) + + return Corpus( + tag=tag, + filter=_tag_filter(tag), + trace_ids=trace_ids, + root_ids=root_ids, + child_ids=child_ids, + ) + + +def io_mapper(*, item: ObservationV2) -> EvaluatorInputs: return EvaluatorInputs( - input=item.input if hasattr(item, "input") else None, - output=item.output if hasattr(item, "output") else None, - expected_output=None, - metadata={"trace_id": item.id}, + input=item.input, + output=item.output, + metadata={"trace_id": item.trace_id}, ) -def simple_evaluator(*, input, output, expected_output=None, metadata=None, **kwargs): - """Simple evaluator that returns a score based on output length.""" - if output is None: - return Evaluation(name="length_score", value=0.0, comment="No output") - - return Evaluation( - name="length_score", - value=float(len(str(output))) / 10.0, - comment=f"Length: {len(str(output))}", - ) +def length_evaluator(*, output, **kwargs): + return Evaluation(name="length", value=float(len(output or ""))) -# ============================================================================ -# BASIC FUNCTIONALITY TESTS -# ============================================================================ +def test_evaluates_every_observation(corpus): + seen: List[ObservationV2] = [] + def recording_mapper(*, item): + seen.append(item) + return io_mapper(item=item) -def test_run_batched_evaluation_on_observations_basic(langfuse_client): - """Test basic batch evaluation on traces.""" - result = langfuse_client.run_batched_evaluation( - scope="observations", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=1, - verbose=True, + result = get_client().run_batched_evaluation( + mapper=recording_mapper, + evaluators=[length_evaluator], + filter=corpus.filter, ) - # Validate result structure - assert isinstance(result, BatchEvaluationResult) - assert result.total_items_fetched >= 0 - assert result.total_items_processed >= 0 - assert result.total_scores_created >= 0 assert result.completed is True - assert isinstance(result.duration_seconds, float) - assert result.duration_seconds > 0 - - # Verify evaluator stats - assert len(result.evaluator_stats) == 1 - stats = result.evaluator_stats[0] - assert isinstance(stats, EvaluatorStats) - assert stats.name == "simple_evaluator" - - -def test_run_batched_evaluation_on_traces_basic(langfuse_client): - """Test basic batch evaluation on traces.""" - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=5, - verbose=True, - ) - - # Validate result structure - assert isinstance(result, BatchEvaluationResult) - assert result.total_items_fetched >= 0 - assert result.total_items_processed >= 0 - assert result.total_scores_created >= 0 - assert result.completed is True - assert isinstance(result.duration_seconds, float) - assert result.duration_seconds > 0 - - # Verify evaluator stats - assert len(result.evaluator_stats) == 1 - stats = result.evaluator_stats[0] - assert isinstance(stats, EvaluatorStats) - assert stats.name == "simple_evaluator" - - -def test_batch_evaluation_with_filter(langfuse_client): - """Test batch evaluation with JSON filter.""" - # Create a trace with specific tag - unique_tag = f"test-filter-{create_uuid()}" - with langfuse_client.start_as_current_observation( - name=f"filtered-trace-{create_uuid()}" - ) as span: - with propagate_attributes(tags=[unique_tag]): - span.set_trace_io( - input="Filtered test", - output="Filtered output", - ) - - langfuse_client.flush() - time.sleep(3) - - # Filter format: array of filter conditions - filter_json = f'[{{"type": "arrayOptions", "column": "tags", "operator": "any of", "value": ["{unique_tag}"]}}]' - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - filter=filter_json, - verbose=True, - ) - - # Should only process the filtered trace - assert result.total_items_fetched >= 1 - assert result.completed is True - - -def test_batch_evaluation_with_metadata(langfuse_client): - """Test that additional metadata is added to all scores.""" - - def metadata_checking_evaluator(*, input, output, metadata=None, **kwargs): - return Evaluation( - name="test_score", - value=1.0, - metadata={"evaluator_data": "test"}, - ) - - additional_metadata = { - "batch_run_id": "test-batch-123", - "evaluation_version": "v2.0", + assert result.has_more_items is False + assert result.resume_token is None + assert result.total_items_fetched == 2 * TRACE_COUNT + assert result.total_items_processed == 2 * TRACE_COUNT + assert result.total_scores_created == 2 * TRACE_COUNT + assert set(result.item_evaluations) == set(corpus.root_ids + corpus.child_ids) + + by_id = {item.id: item for item in seen} + child = by_id[corpus.child_ids[0]] + assert isinstance(child.input, str) + assert json.loads(child.input) == {"question": 0} + assert child.output == "child answer 0" + + +def test_observation_scores_are_attached_to_observations(corpus): + score_name = f"obs-score-{create_uuid()}" + child_filter = json.loads(corpus.filter) + [ + { + "type": "string", + "column": "id", + "operator": "=", + "value": corpus.child_ids[1], + } + ] + + def observation_evaluator(**kwargs): + return Evaluation(name=score_name, value=0.5, comment="ok") + + result = get_client().run_batched_evaluation( + mapper=io_mapper, + evaluators=[observation_evaluator], + filter=json.dumps(child_filter), + metadata={"run": "e2e"}, + ) + + assert result.total_items_processed == 1 + + scores = _wait_for_scores( + trace_id=corpus.trace_ids[1], + observation_id=corpus.child_ids[1], + name=score_name, + ) + assert len(scores) == 1 + assert scores[0].subject.kind == "observation" + assert scores[0].subject.id == corpus.child_ids[1] + assert scores[0].subject.trace_id == corpus.trace_ids[1] + assert scores[0].comment == "ok" + assert scores[0].metadata == {"run": "e2e"} + + +def test_max_items_then_resume_covers_corpus_exactly_once(corpus): + langfuse = get_client() + run_kwargs: Any = { + "mapper": io_mapper, + "evaluators": [length_evaluator], + "filter": corpus.filter, + "fetch_batch_size": 2, } - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[metadata_checking_evaluator], - metadata=additional_metadata, - max_items=2, - ) - - assert result.total_scores_created > 0 - - # Verify scores were created with merged metadata - langfuse_client.flush() - time.sleep(3) - - # Note: In a real test, you'd verify via API that metadata was merged - # For now, just verify the operation completed - assert result.completed is True + first = langfuse.run_batched_evaluation(max_items=3, **run_kwargs) + assert first.total_items_fetched == 3 + assert first.has_more_items is True + assert first.resume_token is not None + assert first.resume_token.cursor is not None -def test_result_structure_fields(langfuse_client): - """Test that BatchEvaluationResult has all expected fields.""" - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=3, + second = langfuse.run_batched_evaluation( + resume_from=first.resume_token, **run_kwargs ) - # Check all result fields exist - assert hasattr(result, "total_items_fetched") - assert hasattr(result, "total_items_processed") - assert hasattr(result, "total_items_failed") - assert hasattr(result, "total_scores_created") - assert hasattr(result, "total_composite_scores_created") - assert hasattr(result, "total_evaluations_failed") - assert hasattr(result, "evaluator_stats") - assert hasattr(result, "resume_token") - assert hasattr(result, "completed") - assert hasattr(result, "duration_seconds") - assert hasattr(result, "failed_item_ids") - assert hasattr(result, "error_summary") - assert hasattr(result, "has_more_items") - assert hasattr(result, "item_evaluations") - - # Check types - assert isinstance(result.evaluator_stats, list) - assert isinstance(result.failed_item_ids, list) - assert isinstance(result.error_summary, dict) - assert isinstance(result.completed, bool) - assert isinstance(result.has_more_items, bool) - assert isinstance(result.item_evaluations, dict) - - -# ============================================================================ -# MAPPER FUNCTION TESTS -# ============================================================================ - - -def test_simple_mapper(langfuse_client): - """Test basic mapper functionality.""" - - def custom_mapper(*, item): - return EvaluatorInputs( - input=item.input if hasattr(item, "input") else "no input", - output=item.output if hasattr(item, "output") else "no output", - expected_output=None, - metadata={"custom_field": "test_value"}, - ) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=custom_mapper, - evaluators=[simple_evaluator], - max_items=2, + assert second.completed is True + assert second.resume_token is None + assert set(first.item_evaluations).isdisjoint(second.item_evaluations) + assert set(first.item_evaluations) | set(second.item_evaluations) == set( + corpus.root_ids + corpus.child_ids ) - assert result.total_items_processed > 0 - - -@pytest.mark.asyncio -async def test_async_mapper(langfuse_client): - """Test that async mappers work correctly.""" - - async def async_mapper(*, item): - await asyncio.sleep(0.01) # Simulate async work - return EvaluatorInputs( - input=item.input if hasattr(item, "input") else None, - output=item.output if hasattr(item, "output") else None, - expected_output=None, - metadata={"async": True}, - ) - - # Note: run_batched_evaluation is synchronous but handles async mappers - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=async_mapper, - evaluators=[simple_evaluator], - max_items=2, - ) - - assert result.total_items_processed > 0 - - -def test_mapper_failure_handling(langfuse_client): - """Test that mapper failures cause items to be skipped.""" - - def failing_mapper(*, item): - raise ValueError("Intentional mapper failure") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=failing_mapper, - evaluators=[simple_evaluator], - max_items=3, - ) - - # All items should fail due to mapper failures - assert result.total_items_failed > 0 - assert len(result.failed_item_ids) > 0 - assert "ValueError" in result.error_summary or "Exception" in result.error_summary - - -def test_mapper_with_missing_fields(langfuse_client): - """Test mapper handles traces with missing fields gracefully.""" - - def robust_mapper(*, item): - # Handle missing fields with defaults - input_val = getattr(item, "input", None) or "default_input" - output_val = getattr(item, "output", None) or "default_output" - - return EvaluatorInputs( - input=input_val, - output=output_val, - expected_output=None, - metadata={}, - ) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=robust_mapper, - evaluators=[simple_evaluator], - max_items=2, - ) - - assert result.total_items_processed > 0 - - -# ============================================================================ -# EVALUATOR TESTS -# ============================================================================ - - -def test_single_evaluator(langfuse_client): - """Test with a single evaluator.""" - - def quality_evaluator(*, input, output, **kwargs): - return Evaluation(name="quality", value=0.85, comment="High quality") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[quality_evaluator], - max_items=2, - ) - assert result.total_scores_created > 0 - assert len(result.evaluator_stats) == 1 - assert result.evaluator_stats[0].name == "quality_evaluator" +def test_fields_control_populated_field_groups(corpus): + seen: List[ObservationV2] = [] + def recording_mapper(*, item): + seen.append(item) + return EvaluatorInputs(input=None, output=None) -def test_multiple_evaluators(langfuse_client): - """Test with multiple evaluators running in parallel.""" - - def accuracy_evaluator(*, input, output, **kwargs): - return Evaluation(name="accuracy", value=0.9) - - def relevance_evaluator(*, input, output, **kwargs): - return Evaluation(name="relevance", value=0.8) - - def safety_evaluator(*, input, output, **kwargs): - return Evaluation(name="safety", value=1.0) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[accuracy_evaluator, relevance_evaluator, safety_evaluator], - max_items=2, - ) - - # Should have 3 evaluators - assert len(result.evaluator_stats) == 3 - assert result.total_scores_created >= result.total_items_processed * 3 - - -@pytest.mark.asyncio -async def test_async_evaluator(langfuse_client): - """Test that async evaluators work correctly.""" - - async def async_evaluator(*, input, output, **kwargs): - await asyncio.sleep(0.01) # Simulate async work - return Evaluation(name="async_score", value=0.75) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[async_evaluator], - max_items=2, - ) - - assert result.total_scores_created > 0 - - -def test_evaluator_returning_list(langfuse_client): - """Test evaluator that returns multiple Evaluations.""" - - def multi_score_evaluator(*, input, output, **kwargs): - return [ - Evaluation(name="score_1", value=0.8), - Evaluation(name="score_2", value=0.9), - Evaluation(name="score_3", value=0.7), - ] - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[multi_score_evaluator], - max_items=2, - ) - - # Should create 3 scores per item - assert result.total_scores_created >= result.total_items_processed * 3 - - -def test_evaluator_failure_statistics(langfuse_client): - """Test that evaluator failures are tracked in statistics.""" - - def working_evaluator(*, input, output, **kwargs): - return Evaluation(name="working", value=1.0) - - def failing_evaluator(*, input, output, **kwargs): - raise RuntimeError("Intentional evaluator failure") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[working_evaluator, failing_evaluator], - max_items=3, - ) - - # Verify evaluator stats - assert len(result.evaluator_stats) == 2 - - working_stats = next( - s for s in result.evaluator_stats if s.name == "working_evaluator" - ) - assert working_stats.successful_runs > 0 - assert working_stats.failed_runs == 0 - - failing_stats = next( - s for s in result.evaluator_stats if s.name == "failing_evaluator" - ) - assert failing_stats.failed_runs > 0 - assert failing_stats.successful_runs == 0 - - # Total evaluations failed should be tracked - assert result.total_evaluations_failed > 0 - - -def test_mixed_sync_async_evaluators(langfuse_client): - """Test mixing synchronous and asynchronous evaluators.""" - - def sync_evaluator(*, input, output, **kwargs): - return Evaluation(name="sync_score", value=0.8) - - async def async_evaluator(*, input, output, **kwargs): - await asyncio.sleep(0.01) - return Evaluation(name="async_score", value=0.9) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[sync_evaluator, async_evaluator], - max_items=2, - ) - - assert len(result.evaluator_stats) == 2 - assert result.total_scores_created >= result.total_items_processed * 2 - - -# ============================================================================ -# COMPOSITE EVALUATOR TESTS -# ============================================================================ - - -def test_composite_evaluator_weighted_average(langfuse_client): - """Test composite evaluator that computes weighted average.""" - - def accuracy_evaluator(*, input, output, **kwargs): - return Evaluation(name="accuracy", value=0.8) - - def relevance_evaluator(*, input, output, **kwargs): - return Evaluation(name="relevance", value=0.9) - - def composite_evaluator(*, input, output, expected_output, metadata, evaluations): - weights = {"accuracy": 0.6, "relevance": 0.4} - total = sum( - e.value * weights.get(e.name, 0) - for e in evaluations - if isinstance(e.value, (int, float)) - ) - - return Evaluation( - name="composite_score", - value=total, - comment=f"Weighted average of {len(evaluations)} metrics", - ) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[accuracy_evaluator, relevance_evaluator], - composite_evaluator=composite_evaluator, - max_items=2, - ) - - # Should have both regular and composite scores - assert result.total_scores_created > 0 - assert result.total_composite_scores_created > 0 - assert result.total_scores_created > result.total_composite_scores_created - - -def test_composite_evaluator_pass_fail(langfuse_client): - """Test composite evaluator that implements pass/fail logic.""" - - def metric1_evaluator(*, input, output, **kwargs): - return Evaluation(name="metric1", value=0.9) - - def metric2_evaluator(*, input, output, **kwargs): - return Evaluation(name="metric2", value=0.7) - - def pass_fail_composite(*, input, output, expected_output, metadata, evaluations): - thresholds = {"metric1": 0.8, "metric2": 0.6} - - passes = all( - e.value >= thresholds.get(e.name, 0) - for e in evaluations - if isinstance(e.value, (int, float)) - ) - - return Evaluation( - name="passes_all_checks", - value=1.0 if passes else 0.0, - comment="All checks passed" if passes else "Some checks failed", - ) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[metric1_evaluator, metric2_evaluator], - composite_evaluator=pass_fail_composite, - max_items=2, - ) - - assert result.total_composite_scores_created > 0 - - -@pytest.mark.asyncio -async def test_async_composite_evaluator(langfuse_client): - """Test async composite evaluator.""" - - def evaluator1(*, input, output, **kwargs): - return Evaluation(name="eval1", value=0.8) - - async def async_composite(*, input, output, expected_output, metadata, evaluations): - await asyncio.sleep(0.01) # Simulate async processing - avg = sum( - e.value for e in evaluations if isinstance(e.value, (int, float)) - ) / len(evaluations) - return Evaluation(name="async_composite", value=avg) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[evaluator1], - composite_evaluator=async_composite, - max_items=2, - ) - - assert result.total_composite_scores_created > 0 - - -def test_composite_evaluator_with_no_evaluations(langfuse_client): - """Test composite evaluator when no evaluations are present.""" - - def always_failing_evaluator(*, input, output, **kwargs): - raise Exception("Always fails") - - def composite_evaluator(*, input, output, expected_output, metadata, evaluations): - # Should not be called if no evaluations succeed - return Evaluation(name="composite", value=0.0) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[always_failing_evaluator], - composite_evaluator=composite_evaluator, - max_items=2, - ) - - # Composite evaluator should not create scores if no evaluations - assert result.total_composite_scores_created == 0 - - -def test_composite_evaluator_failure_handling(langfuse_client): - """Test that composite evaluator failures are handled gracefully.""" - - def evaluator1(*, input, output, **kwargs): - return Evaluation(name="eval1", value=0.8) - - def failing_composite(*, input, output, expected_output, metadata, evaluations): - raise ValueError("Composite evaluator failed") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[evaluator1], - composite_evaluator=failing_composite, - max_items=2, - ) - - # Regular scores should still be created - assert result.total_scores_created > 0 - # But no composite scores - assert result.total_composite_scores_created == 0 - - -# ============================================================================ -# ERROR HANDLING TESTS -# ============================================================================ - - -def test_mapper_failure_skips_item(langfuse_client): - """Test that mapper failure causes item to be skipped.""" - - call_count = {"count": 0} - - def sometimes_failing_mapper(*, item): - call_count["count"] += 1 - if call_count["count"] % 2 == 0: - raise Exception("Mapper failed") - return simple_trace_mapper(item=item) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=sometimes_failing_mapper, - evaluators=[simple_evaluator], - max_items=4, - ) - - # Some items should fail, some should succeed - assert result.total_items_failed > 0 - assert result.total_items_processed > 0 - - -def test_evaluator_failure_continues(langfuse_client): - """Test that one evaluator failing doesn't stop others.""" - - def working_evaluator1(*, input, output, **kwargs): - return Evaluation(name="working1", value=0.8) - - def failing_evaluator(*, input, output, **kwargs): - raise Exception("Evaluator failed") - - def working_evaluator2(*, input, output, **kwargs): - return Evaluation(name="working2", value=0.9) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[working_evaluator1, failing_evaluator, working_evaluator2], - max_items=2, - ) - - # Working evaluators should still create scores - assert result.total_scores_created >= result.total_items_processed * 2 - - # Failing evaluator should be tracked - failing_stats = next( - s for s in result.evaluator_stats if s.name == "failing_evaluator" - ) - assert failing_stats.failed_runs > 0 - - -def test_all_evaluators_fail(langfuse_client): - """Test when all evaluators fail but item is still processed.""" - - def failing_evaluator1(*, input, output, **kwargs): - raise Exception("Failed 1") - - def failing_evaluator2(*, input, output, **kwargs): - raise Exception("Failed 2") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[failing_evaluator1, failing_evaluator2], - max_items=2, - ) - - # Items should be processed even if all evaluators fail - assert result.total_items_processed > 0 - # But no scores created - assert result.total_scores_created == 0 - # All evaluations failed - assert result.total_evaluations_failed > 0 - - -# ============================================================================ -# EDGE CASES TESTS -# ============================================================================ - - -def test_empty_results_handling(langfuse_client): - """Test batch evaluation when filter returns no items.""" - nonexistent_name = f"nonexistent-trace-{create_uuid()}" - nonexistent_filter = f'[{{"type": "string", "column": "name", "operator": "=", "value": "{nonexistent_name}"}}]' - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - filter=nonexistent_filter, - ) - - assert result.total_items_fetched == 0 - assert result.total_items_processed == 0 - assert result.total_scores_created == 0 - assert result.completed is True - assert result.has_more_items is False - - -def test_max_items_zero(langfuse_client): - """Test with max_items=0 (should process no items).""" - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=0, - ) - - assert result.total_items_fetched == 0 - assert result.total_items_processed == 0 - - -def test_evaluation_value_type_conversions(langfuse_client): - """Test that different evaluation value types are handled correctly.""" - - def multi_type_evaluator(*, input, output, **kwargs): - return [ - Evaluation(name="int_score", value=5), # int - Evaluation(name="float_score", value=0.85), # float - Evaluation(name="bool_score", value=True), # bool - Evaluation(name="none_score", value=None), # None - ] - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[multi_type_evaluator], + get_client().run_batched_evaluation( + mapper=recording_mapper, + evaluators=[length_evaluator], + filter=corpus.filter, + fields="core,basic", max_items=1, ) - # All value types should be converted and scores created - assert result.total_scores_created >= 4 - - -# ============================================================================ -# PAGINATION TESTS -# ============================================================================ - - -def test_pagination_with_max_items(langfuse_client): - """Test that max_items limit is respected.""" - # Create more traces to ensure we have enough data - for i in range(10): - with langfuse_client.start_as_current_observation( - name=f"pagination-test-{create_uuid()}" - ) as span: - with propagate_attributes(tags=["pagination_test"]): - span.set_trace_io( - input=f"Input {i}", - output=f"Output {i}", - ) - - langfuse_client.flush() - time.sleep(3) - - filter_json = '[{"type": "arrayOptions", "column": "tags", "operator": "any of", "value": ["pagination_test"]}]' - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - filter=filter_json, - max_items=5, - fetch_batch_size=2, - ) - - # Should not exceed max_items - assert result.total_items_processed <= 5 - - -def test_has_more_items_flag(langfuse_client): - """Test that has_more_items flag is set correctly when max_items is reached.""" - # Create enough traces to exceed max_items - batch_tag = f"batch-test-{create_uuid()}" - for i in range(15): - with langfuse_client.start_as_current_observation( - name=f"more-items-test-{i}" - ) as span: - with propagate_attributes(tags=[batch_tag]): - span.set_trace_io( - input=f"Input {i}", - output=f"Output {i}", - ) - - langfuse_client.flush() - time.sleep(3) - - filter_json = f'[{{"type": "arrayOptions", "column": "tags", "operator": "any of", "value": ["{batch_tag}"]}}]' - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - filter=filter_json, - max_items=5, - fetch_batch_size=2, - ) - - # has_more_items should be True if we hit the limit - if result.total_items_fetched >= 5: - assert result.has_more_items is True - - -def test_fetch_batch_size_parameter(langfuse_client): - """Test that different fetch_batch_size values work correctly.""" - for batch_size in [1, 5, 10]: - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=3, - fetch_batch_size=batch_size, - ) - - # Should complete regardless of batch size - assert result.completed is True or result.total_items_processed > 0 - - -# ============================================================================ -# RESUME FUNCTIONALITY TESTS -# ============================================================================ - - -def test_resume_token_structure(langfuse_client): - """Test that BatchEvaluationResumeToken has correct structure.""" - resume_token = BatchEvaluationResumeToken( - scope="traces", - filter='{"test": "filter"}', - last_processed_timestamp="2024-01-01T00:00:00Z", - last_processed_id="trace-123", - items_processed=10, - ) - - assert resume_token.scope == "traces" - assert resume_token.filter == '{"test": "filter"}' - assert resume_token.last_processed_timestamp == "2024-01-01T00:00:00Z" - assert resume_token.last_processed_id == "trace-123" - assert resume_token.items_processed == 10 - - -# ============================================================================ -# CONCURRENCY TESTS -# ============================================================================ - - -def test_max_concurrency_parameter(langfuse_client): - """Test that max_concurrency parameter works correctly.""" - for concurrency in [1, 5, 10]: - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=3, - max_concurrency=concurrency, - ) - - # Should complete regardless of concurrency - assert result.completed is True or result.total_items_processed > 0 + assert len(seen) == 1 + assert seen[0].input is None + assert seen[0].output is None -# ============================================================================ -# STATISTICS TESTS -# ============================================================================ +def test_composite_evaluator_and_failures(corpus): + def failing_evaluator(**kwargs): + raise RuntimeError("intentional") + def composite(*, evaluations, **kwargs): + return Evaluation(name="composite", value=float(len(evaluations))) -def test_evaluator_stats_structure(langfuse_client): - """Test that EvaluatorStats has correct structure.""" - - def test_evaluator(*, input, output, **kwargs): - return Evaluation(name="test", value=1.0) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[test_evaluator], - max_items=2, + result = get_client().run_batched_evaluation( + mapper=io_mapper, + evaluators=[length_evaluator, failing_evaluator], + composite_evaluator=composite, + filter=corpus.filter, ) - assert len(result.evaluator_stats) == 1 - stats = result.evaluator_stats[0] - - # Check all fields exist - assert hasattr(stats, "name") - assert hasattr(stats, "total_runs") - assert hasattr(stats, "successful_runs") - assert hasattr(stats, "failed_runs") - assert hasattr(stats, "total_scores_created") - - # Check values - assert stats.name == "test_evaluator" - assert stats.total_runs == result.total_items_processed - assert stats.successful_runs == result.total_items_processed - assert stats.failed_runs == 0 - + item_count = 2 * TRACE_COUNT + assert result.total_items_processed == item_count + assert result.total_scores_created == item_count + assert result.total_composite_scores_created == item_count + assert result.total_evaluations_failed == item_count + stats = {s.name: s for s in result.evaluator_stats} + assert stats["failing_evaluator"].failed_runs == item_count -def test_evaluator_stats_tracking(langfuse_client): - """Test that evaluator statistics are tracked correctly.""" - - call_count = {"count": 0} - - def sometimes_failing_evaluator(*, input, output, **kwargs): - call_count["count"] += 1 - if call_count["count"] % 2 == 0: - raise Exception("Failed") - return Evaluation(name="test", value=1.0) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[sometimes_failing_evaluator], - max_items=4, - ) - - stats = result.evaluator_stats[0] - assert stats.total_runs == result.total_items_processed - assert stats.successful_runs > 0 - assert stats.failed_runs > 0 - assert stats.successful_runs + stats.failed_runs == stats.total_runs - - -def test_error_summary_aggregation(langfuse_client): - """Test that error types are aggregated correctly in error_summary.""" - - def failing_mapper(*, item): - raise ValueError("Mapper error") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=failing_mapper, - evaluators=[simple_evaluator], - max_items=3, - ) - - # Error summary should contain the error type - assert len(result.error_summary) > 0 - assert any("Error" in key for key in result.error_summary.keys()) - - -def test_failed_item_ids_collected(langfuse_client): - """Test that failed item IDs are collected.""" +def test_mapper_failures_are_reported_per_item(corpus): def failing_mapper(*, item): - raise Exception("Failed") + raise ValueError("intentional") - result = langfuse_client.run_batched_evaluation( - scope="traces", + result = get_client().run_batched_evaluation( mapper=failing_mapper, - evaluators=[simple_evaluator], - max_items=3, - ) - - assert len(result.failed_item_ids) > 0 - # Each failed ID should be a string - assert all(isinstance(item_id, str) for item_id in result.failed_item_ids) - - -# ============================================================================ -# PERFORMANCE TESTS -# ============================================================================ - - -def test_duration_tracking(langfuse_client): - """Test that duration is tracked correctly.""" - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=2, - ) - - assert result.duration_seconds > 0 - assert result.duration_seconds < 60 # Should complete quickly for small batch - - -def test_verbose_logging(langfuse_client): - """Test that verbose=True doesn't cause errors.""" - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[simple_evaluator], - max_items=2, - verbose=True, # Should log progress + evaluators=[length_evaluator], + filter=corpus.filter, ) assert result.completed is True + assert result.total_items_failed == 2 * TRACE_COUNT + assert set(result.failed_item_ids) == set(corpus.root_ids + corpus.child_ids) + assert result.error_summary == {"ValueError": 2 * TRACE_COUNT} -# ============================================================================ -# ITEM EVALUATIONS TESTS -# ============================================================================ - - -def test_item_evaluations_basic(langfuse_client): - """Test that item_evaluations dict contains correct structure.""" - - def test_evaluator(*, input, output, **kwargs): - return Evaluation(name="test_metric", value=0.5) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[test_evaluator], - max_items=3, - ) - - # Check that item_evaluations is a dict - assert isinstance(result.item_evaluations, dict) - - # Should have evaluations for each processed item - assert len(result.item_evaluations) == result.total_items_processed - - # Each entry should be a list of Evaluation objects - for item_id, evaluations in result.item_evaluations.items(): - assert isinstance(item_id, str) - assert isinstance(evaluations, list) - assert all(isinstance(e, Evaluation) for e in evaluations) - # Should have one evaluation per evaluator - assert len(evaluations) == 1 - assert evaluations[0].name == "test_metric" - - -def test_item_evaluations_multiple_evaluators(langfuse_client): - """Test item_evaluations with multiple evaluators.""" - - def accuracy_evaluator(*, input, output, **kwargs): - return Evaluation(name="accuracy", value=0.8) - - def relevance_evaluator(*, input, output, **kwargs): - return Evaluation(name="relevance", value=0.9) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[accuracy_evaluator, relevance_evaluator], - max_items=2, - ) - - # Check structure - assert len(result.item_evaluations) == result.total_items_processed - - # Each item should have evaluations from both evaluators - for item_id, evaluations in result.item_evaluations.items(): - assert len(evaluations) == 2 - eval_names = {e.name for e in evaluations} - assert eval_names == {"accuracy", "relevance"} - - -def test_item_evaluations_with_composite(langfuse_client): - """Test that item_evaluations includes composite evaluations.""" - - def base_evaluator(*, input, output, **kwargs): - return Evaluation(name="base_score", value=0.7) - - def composite_evaluator(*, input, output, expected_output, metadata, evaluations): - return Evaluation( - name="composite_score", - value=sum( - e.value for e in evaluations if isinstance(e.value, (int, float)) - ), - ) - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=simple_trace_mapper, - evaluators=[base_evaluator], - composite_evaluator=composite_evaluator, - max_items=2, - ) - - # Each item should have both base and composite evaluations - for item_id, evaluations in result.item_evaluations.items(): - assert len(evaluations) == 2 - eval_names = {e.name for e in evaluations} - assert eval_names == {"base_score", "composite_score"} - - # Verify composite scores were created - assert result.total_composite_scores_created > 0 - - -def test_item_evaluations_empty_on_failure(langfuse_client): - """Test that failed items don't appear in item_evaluations.""" - - def failing_mapper(*, item): - raise Exception("Mapper failed") - - result = langfuse_client.run_batched_evaluation( - scope="traces", - mapper=failing_mapper, - evaluators=[simple_evaluator], - max_items=3, +def test_filter_without_matches_completes_empty(): + result = get_client().run_batched_evaluation( + mapper=io_mapper, + evaluators=[length_evaluator], + filter=_tag_filter(f"nonexistent-{create_uuid()}"), ) - # All items failed, so item_evaluations should be empty - assert len(result.item_evaluations) == 0 - assert result.total_items_failed > 0 + assert result.completed is True + assert result.total_items_fetched == 0 + assert result.has_more_items is False + assert result.resume_token is None diff --git a/tests/e2e/test_core_sdk.py b/tests/e2e/test_core_sdk.py index 614d6da41..5330932c3 100644 --- a/tests/e2e/test_core_sdk.py +++ b/tests/e2e/test_core_sdk.py @@ -1,3 +1,4 @@ +import json import os import time from asyncio import gather @@ -5,17 +6,20 @@ from time import sleep import pytest -from tenacity import Retrying, stop_after_delay, wait_fixed from langfuse import Langfuse, propagate_attributes from langfuse._client.resource_manager import LangfuseResourceManager from langfuse._utils import _get_timestamp -from tests.support.api_wrapper import LangfuseAPI from tests.support.utils import ( create_uuid, get_api, - wait_for_result, - wait_for_trace, + get_observations, + get_root_observation, + get_scores, + user_metadata, + wait_for_observations, + wait_for_root_observation, + wait_for_scores, ) @@ -38,31 +42,28 @@ async def update_generation(i, langfuse: Langfuse): # End the generation generation.end() + return generation.trace_id + # Create Langfuse client langfuse = Langfuse() # Run concurrent operations - await gather(*(update_generation(i, langfuse) for i in range(100))) + trace_ids = await gather(*(update_generation(i, langfuse) for i in range(100))) langfuse.flush() - # Allow time for all operations to be processed - sleep(10) - # Verify that all spans were created properly - api = get_api() - for i in range(100): - # Find the observations with the expected name - observations = api.legacy.observations_v1.get_many(name=str(i)).data + for i, trace_id in enumerate(trace_ids): + observations = wait_for_observations(trace_id, min_count=2) - # Find generation observations (there should be at least one) generation_obs = [obs for obs in observations if obs.type == "GENERATION"] - assert len(generation_obs) > 0 + assert len(generation_obs) == 1 # Verify metadata observation = generation_obs[0] assert observation.name == str(i) - assert observation.metadata["count"] == i + assert user_metadata(observation)["count"] == i + assert get_root_observation(observations).trace_name == str(i) def test_flush(): @@ -80,14 +81,10 @@ def test_flush(): # Flush all pending spans to the Langfuse API langfuse.flush() - # Allow time for API to process - sleep(2) - # Verify traces were sent by checking they exist in the API - api = get_api() for i, trace_id in enumerate(trace_ids): - trace = api.trace.get(trace_id) - assert trace.name == str(i) + root = wait_for_root_observation(trace_id) + assert root.trace_name == str(i) def test_invalid_score_data_does_not_raise_exception(): @@ -152,21 +149,20 @@ def test_create_session_score(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - score = get_api().scores.get_by_id(score_id) + scores = wait_for_scores(session_id=session_id, id=score_id) - # find the score by name (server may transform the id format) - assert score is not None + assert len(scores) == 1 + score = scores[0] assert score.value == 1 assert score.data_type == "NUMERIC" - assert score.session_id == session_id + assert score.subject.kind == "session" + assert score.subject.id == session_id def test_create_numeric_score(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span and set trace properties with langfuse.start_as_current_observation(name="test-span") as span: @@ -202,22 +198,21 @@ def test_create_numeric_score(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = api_wrapper.get_trace(trace_id) + wait_for_observations(trace_id, min_count=2) + scores = wait_for_scores(trace_id=trace_id, name="this-is-a-score") - # Find the score by name (server may transform the ID format) - score = next((s for s in trace["scores"] if s["name"] == "this-is-a-score"), None) - assert score is not None - assert score["value"] == 1 - assert score["dataType"] == "NUMERIC" - assert score["stringValue"] is None + assert len(scores) == 1 + score = scores[0] + assert score.id == score_id + assert score.value == 1 + assert score.data_type == "NUMERIC" + assert score.subject.kind == "trace" def test_create_boolean_score(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span and set trace properties with langfuse.start_as_current_observation(name="test-span") as span: @@ -231,7 +226,7 @@ def test_create_boolean_score(): # Ensure data is sent langfuse.flush() - api_wrapper.get_trace(trace_id) + wait_for_root_observation(trace_id) # Create a boolean score score_id = create_uuid() @@ -256,27 +251,17 @@ def test_create_boolean_score(): langfuse.flush() # Retrieve and verify - trace = api_wrapper.get_trace( - trace_id, - is_result_ready=lambda trace: any( - score["name"] == "this-is-a-score" for score in trace.get("scores", []) - ), - ) + scores = wait_for_scores(trace_id=trace_id, name="this-is-a-score") - # Find the score we created by name - created_score = next( - (s for s in trace["scores"] if s["name"] == "this-is-a-score"), None - ) - assert created_score is not None, "Score not found in trace" - assert created_score["id"] == score_id - assert created_score["dataType"] == "BOOLEAN" - assert created_score["value"] == 1 - assert created_score["stringValue"] == "True" + assert len(scores) == 1, "Score not found in trace" + created_score = scores[0] + assert created_score.id == score_id + assert created_score.data_type == "BOOLEAN" + assert created_score.value is True def test_create_categorical_score(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span and set trace properties with langfuse.start_as_current_observation(name="test-span") as span: @@ -290,7 +275,7 @@ def test_create_categorical_score(): # Ensure data is sent langfuse.flush() - api_wrapper.get_trace(trace_id) + wait_for_root_observation(trace_id) # Create a categorical score score_id = create_uuid() @@ -314,27 +299,17 @@ def test_create_categorical_score(): langfuse.flush() # Retrieve and verify - trace = api_wrapper.get_trace( - trace_id, - is_result_ready=lambda trace: any( - score["name"] == "this-is-a-score" for score in trace.get("scores", []) - ), - ) + scores = wait_for_scores(trace_id=trace_id, name="this-is-a-score") - # Find the score we created by name - created_score = next( - (s for s in trace["scores"] if s["name"] == "this-is-a-score"), None - ) - assert created_score is not None, "Score not found in trace" - assert created_score["id"] == score_id - assert created_score["dataType"] == "CATEGORICAL" - assert created_score["value"] == 0 - assert created_score["stringValue"] == "high score" + assert len(scores) == 1, "Score not found in trace" + created_score = scores[0] + assert created_score.id == score_id + assert created_score.data_type == "CATEGORICAL" + assert created_score.value == "high score" def test_create_text_score(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span and set trace properties with langfuse.start_as_current_observation(name="test-span") as span: @@ -372,30 +347,21 @@ def test_create_text_score(): # Ensure data is sent langfuse.flush() - # Retrieve and verify with retry - for attempt in Retrying( - stop=stop_after_delay(10), wait=wait_fixed(0.1), reraise=True - ): - with attempt: - trace = api_wrapper.get_trace(trace_id) - - # Find the score we created by name - created_score = next( - (s for s in trace["scores"] if s["name"] == "this-is-a-score"), None - ) - assert created_score is not None, "Score not found in trace" - assert created_score["id"] == score_id - assert created_score["dataType"] == "TEXT" + # Retrieve and verify + scores = wait_for_scores(trace_id=trace_id, name="this-is-a-score") - assert ( - created_score["stringValue"] - == "This is a detailed text evaluation of the output quality." - ) + assert len(scores) == 1, "Score not found in trace" + created_score = scores[0] + assert created_score.id == score_id + assert created_score.data_type == "TEXT" + assert ( + created_score.value + == "This is a detailed text evaluation of the output quality." + ) def test_create_score_with_custom_timestamp(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span and set trace properties with langfuse.start_as_current_observation(name="test-span") as span: @@ -409,7 +375,7 @@ def test_create_score_with_custom_timestamp(): # Ensure data is sent langfuse.flush() - api_wrapper.get_trace(trace_id) + wait_for_root_observation(trace_id) custom_timestamp = datetime.now(timezone.utc) - timedelta(hours=1) score_id = create_uuid() @@ -426,28 +392,16 @@ def test_create_score_with_custom_timestamp(): langfuse.flush() # Retrieve and verify - trace = api_wrapper.get_trace( - trace_id, - is_result_ready=lambda trace: any( - score["name"] == "custom-timestamp-score" - for score in trace.get("scores", []) - ), - ) + scores = wait_for_scores(trace_id=trace_id, name="custom-timestamp-score") - # Find the score we created by name - created_score = next( - (s for s in trace["scores"] if s["name"] == "custom-timestamp-score"), None - ) - assert created_score is not None, "Score not found in trace" - assert created_score["id"] == score_id - assert created_score["dataType"] == "NUMERIC" - assert created_score["value"] == 0.85 + assert len(scores) == 1, "Score not found in trace" + created_score = scores[0] + assert created_score.id == score_id + assert created_score.data_type == "NUMERIC" + assert created_score.value == 0.85 # Verify timestamp is close to our custom timestamp - # Parse the timestamp from the API response - response_timestamp = datetime.fromisoformat( - created_score["timestamp"].replace("Z", "+00:00") - ) + response_timestamp = created_score.timestamp # Check that the timestamps are within 1 second of each other # (allowing for some processing time and rounding) @@ -477,24 +431,14 @@ def test_create_trace(): langfuse.flush() # Retrieve the trace from the API - trace = LangfuseAPI().get_trace( - trace_id, - is_result_ready=lambda trace: ( - trace.get("name") == trace_name - and trace.get("userId") == "test" - and trace.get("metadata", {}).get("key") == "value" - and trace.get("tags") == ["tag1", "tag2"] - and trace.get("public") is True - ), - ) + root = wait_for_root_observation(trace_id) # Verify all trace properties - assert trace["name"] == trace_name - assert trace["userId"] == "test" - assert trace["metadata"]["key"] == "value" - assert trace["tags"] == ["tag1", "tag2"] - assert trace["public"] is True - assert True if not trace["externalId"] else False + assert root.trace_name == trace_name + assert root.user_id == "test" + assert user_metadata(root) == {"key": "value"} + assert root.tags == ["tag1", "tag2"] + assert root.public is True def test_create_update_trace(): @@ -525,23 +469,12 @@ def test_create_update_trace(): assert isinstance(trace_id, str) # Retrieve and verify trace - trace = wait_for_trace( - trace_id, - is_result_ready=lambda trace: ( - trace.name == trace_name - and trace.user_id == "test" - and trace.metadata is not None - and trace.metadata.get("key") == "value" - and trace.metadata.get("key2") == "value2" - and trace.public is True - ), - ) + root = wait_for_root_observation(trace_id) - assert trace.name == trace_name - assert trace.user_id == "test" - assert trace.metadata["key"] == "value" - assert trace.metadata["key2"] == "value2" - assert trace.public is True + assert root.trace_name == trace_name + assert root.user_id == "test" + assert user_metadata(root) == {"key": "value", "key2": "value2"} + assert root.public is True def test_create_update_current_trace(): @@ -549,14 +482,14 @@ def test_create_update_current_trace(): trace_name = create_uuid() - # Create initial span with trace properties using propagate_attributes and set_current_trace_io + # Create initial span with trace properties using propagate_attributes with langfuse.start_as_current_observation(name="test-span-current") as span: with propagate_attributes( trace_name=trace_name, user_id="test", metadata={"key": "value"}, ): - langfuse.set_current_trace_io(input="test_input") + langfuse.update_current_span(input="test_input") langfuse.set_current_trace_as_public() # Get trace ID for later reference trace_id = span.trace_id @@ -570,20 +503,18 @@ def test_create_update_current_trace(): # Ensure data is sent to the API langfuse.flush() - sleep(2) assert isinstance(trace_id, str) # Retrieve and verify trace - trace = get_api().trace.get(trace_id) + root = wait_for_root_observation(trace_id) # The 2nd update to the trace must not erase previously set attributes - assert trace.name == trace_name - assert trace.user_id == "test" - assert trace.metadata["key"] == "value" - assert trace.metadata["key2"] == "value2" - assert trace.public is True - assert trace.version == "1.0" - assert trace.input == "test_input" + assert root.trace_name == trace_name + assert root.user_id == "test" + assert user_metadata(root) == {"key": "value", "key2": "value2"} + assert root.public is True + assert root.version == "1.0" + assert root.input == "test_input" def test_create_generation(): @@ -620,19 +551,19 @@ def test_create_generation(): # Flush to ensure all data is sent langfuse.flush() - sleep(2) # Retrieve the trace from the API - trace = get_api().trace.get(trace_id) - - # Verify trace details - assert trace.name == "query-generation" - assert trace.user_id is None + observations = wait_for_observations(trace_id) - assert len(trace.observations) == 1 + assert len(observations) == 1 # Verify generation details - generation_api = trace.observations[0] + generation_api = observations[0] + + # Verify trace details + assert generation_api.is_root_observation is True + assert generation_api.trace_name == "query-generation" + assert generation_api.user_id is None assert generation_api.name == "query-generation" assert generation_api.start_time is not None @@ -707,15 +638,14 @@ def test_create_generation_complex( langfuse.flush() trace_id = generation.trace_id - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id) - assert trace.name == "query-generation" - assert trace.user_id is None + assert len(observations) == 1 - assert len(trace.observations) == 1 - - generation_api = trace.observations[0] + generation_api = observations[0] + assert generation_api.trace_name == "query-generation" + assert generation_api.user_id is None assert generation_api.id == generation.id assert generation_api.name == "query-generation" assert generation_api.input == [ @@ -727,13 +657,7 @@ def test_create_generation_complex( ] assert generation_api.output == [{"foo": "bar"}] - # Check if metadata exists and has tags before asserting - if ( - hasattr(generation_api, "metadata") - and generation_api.metadata is not None - and "tags" in generation_api.metadata - ): - assert generation_api.metadata["tags"] == ["yo"] + assert user_metadata(generation_api) == {"tags": ["yo"]} assert generation_api.start_time is not None assert generation_api.usage_details == {"input": 51, "output": 0, "total": 100} @@ -759,19 +683,18 @@ def test_create_span(): # Ensure all data is sent langfuse.flush() - sleep(2) # Retrieve from API - trace = get_api().trace.get(trace_id) - - # Verify trace details - assert trace.name == "span" - assert trace.user_id is None + observations = wait_for_observations(trace_id) - assert len(trace.observations) == 1 + assert len(observations) == 1 # Verify span details - span_api = trace.observations[0] + span_api = observations[0] + + # Verify trace details + assert span_api.trace_name == "span" + assert span_api.user_id is None assert span_api.id == span_id assert span_api.name == "span" @@ -784,7 +707,6 @@ def test_create_span(): def test_score_trace(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() trace_name = create_uuid() @@ -803,20 +725,17 @@ def test_score_trace(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = api_wrapper.get_trace(trace_id) - - assert trace["name"] == trace_name + assert wait_for_root_observation(trace_id).trace_name == trace_name - # Find the score we created by name (server may create additional auto-scores) - score = next((s for s in trace["scores"] if s["name"] == "valuation"), None) - assert score is not None - assert score["value"] == 0.5 - assert score["comment"] == "This is a comment" - assert score["observationId"] is None - assert score["dataType"] == "NUMERIC" + scores = wait_for_scores(trace_id=trace_id, name="valuation") + assert len(scores) == 1 + score = scores[0] + assert score.value == 0.5 + assert score.comment == "This is a comment" + assert score.subject.kind == "trace" + assert score.data_type == "NUMERIC" def test_score_trace_nested_trace(): @@ -839,19 +758,16 @@ def test_score_trace_nested_trace(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + assert wait_for_root_observation(trace_id).trace_name == trace_name - assert trace.name == trace_name - - # Find the score we created by name (server may create additional auto-scores) - score = next((s for s in trace.scores if s.name == "valuation"), None) - assert score is not None + scores = wait_for_scores(trace_id=trace_id, name="valuation") + assert len(scores) == 1 + score = scores[0] assert score.value == 0.5 assert score.comment == "This is a comment" - assert score.observation_id is None # API returns this field name + assert score.subject.kind == "trace" assert score.data_type == "NUMERIC" @@ -882,25 +798,23 @@ def test_score_trace_nested_observation(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) - - assert trace.name == trace_name + assert wait_for_root_observation(trace_id).trace_name == trace_name - # Find the score we created by name (server may create additional auto-scores) - score = next((s for s in trace.scores if s.name == "valuation"), None) - assert score is not None + scores = wait_for_scores(trace_id=trace_id, name="valuation") + assert len(scores) == 1 + score = scores[0] assert score.value == 0.5 assert score.comment == "This is a comment" - assert score.observation_id == child_span_id # API returns this field name + assert score.subject.kind == "observation" + assert score.subject.id == child_span_id + assert score.subject.trace_id == trace_id assert score.data_type == "NUMERIC" def test_score_span(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span span = langfuse.start_observation( @@ -928,20 +842,19 @@ def test_score_span(): # Ensure data is sent langfuse.flush() - sleep(3) # Retrieve and verify - trace = api_wrapper.get_trace(trace_id) + assert len(wait_for_observations(trace_id)) == 1 - assert len(trace["observations"]) == 1 - - # Find the score we created by name (server may create additional auto-scores) - score = next((s for s in trace["scores"] if s["name"] == "valuation"), None) - assert score is not None - assert score["value"] == 1 - assert score["comment"] == "This is a comment" - assert score["observationId"] == span_id - assert score["dataType"] == "NUMERIC" + scores = wait_for_scores(trace_id=trace_id, observation_id=span_id) + assert len(scores) == 1 + score = scores[0] + assert score.name == "valuation" + assert score.value == 1 + assert score.comment == "This is a comment" + assert score.subject.kind == "observation" + assert score.subject.id == span_id + assert score.data_type == "NUMERIC" def test_create_trace_and_span(): @@ -963,16 +876,15 @@ def test_create_trace_and_span(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id, min_count=2) - assert trace.name == trace_name - assert len(trace.observations) == 2 # Parent span and child span + assert get_root_observation(observations).trace_name == trace_name + assert len(observations) == 2 # Parent span and child span # Find the child span - child_spans = [obs for obs in trace.observations if obs.name == "span"] + child_spans = [obs for obs in observations if obs.name == "span"] assert len(child_spans) == 1 span = child_spans[0] @@ -989,7 +901,7 @@ def test_create_trace_and_generation(): # Create parent span and set trace properties with langfuse.start_as_current_observation(name=trace_name) as parent_span: with propagate_attributes(trace_name=trace_name, session_id="test-session-id"): - parent_span.set_trace_io(input={"key": "value"}) + parent_span.update(input={"key": "value"}) # Create a generation as child generation = parent_span.start_observation( @@ -1004,35 +916,30 @@ def test_create_trace_and_generation(): # Ensure data is sent langfuse.flush() - sleep(2) - # Retrieve traces in two ways - dbTrace = get_api().trace.get(trace_id) - getTrace = get_api().trace.get( - trace_id - ) # Using API as direct getTrace not available + # Retrieve and verify + observations = wait_for_observations(trace_id, min_count=2) + root = get_root_observation(observations) # Verify trace details - assert dbTrace.name == trace_name - assert len(dbTrace.observations) == 2 # Parent span and generation - assert getTrace.name == trace_name - assert len(getTrace.observations) == 2 - assert getTrace.session_id == "test-session-id" + assert root.trace_name == trace_name + assert len(observations) == 2 # Parent span and generation + assert root.session_id == "test-session-id" # Find the generation - generations = [obs for obs in getTrace.observations if obs.name == "generation"] + generations = [obs for obs in observations if obs.name == "generation"] assert len(generations) == 1 generation = generations[0] assert generation.name == "generation" assert generation.trace_id == trace_id + assert generation.session_id == "test-session-id" assert generation.start_time is not None - assert getTrace.input == {"key": "value"} + assert root.input == {"key": "value"} def test_create_generation_and_trace(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() trace_name = create_uuid() @@ -1063,23 +970,24 @@ def test_create_generation_and_trace(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = api_wrapper.get_trace(trace_id) - - assert trace["name"] == trace_name + observations = wait_for_observations(trace_id, min_count=2) # We should have 2 observations (the generation and the span for updating trace) - assert len(trace["observations"]) == 2 + assert len(observations) == 2 + + trace_update_spans = [obs for obs in observations if obs.name == "trace-update"] + assert len(trace_update_spans) == 1 + assert trace_update_spans[0].trace_name == trace_name # Find the generation - generations = [obs for obs in trace["observations"] if obs["name"] == "generation"] + generations = [obs for obs in observations if obs.name == "generation"] assert len(generations) == 1 generation_obs = generations[0] - assert generation_obs["name"] == "generation" - assert generation_obs["traceId"] == trace["id"] + assert generation_obs.name == "generation" + assert generation_obs.trace_id == trace_id def test_create_span_and_get_observation(): @@ -1096,12 +1004,12 @@ def test_create_span_and_get_observation(): # Flush and wait langfuse.flush() - sleep(2) - # Use API to fetch the observation by ID - observation = get_api().legacy.observations_v1.get(span_id) + observations = wait_for_observations(span.trace_id) # Verify observation properties + assert len(observations) == 1 + observation = observations[0] assert observation.name == "span" assert observation.id == span_id @@ -1123,20 +1031,19 @@ def test_update_generation(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id) # Verify trace properties - assert trace.name == "generation" - assert len(trace.observations) == 1 + assert len(observations) == 1 # Verify generation updates - retrieved_generation = trace.observations[0] + retrieved_generation = observations[0] + assert retrieved_generation.trace_name == "generation" assert retrieved_generation.name == "generation" assert retrieved_generation.trace_id == trace_id - assert retrieved_generation.metadata["dict"] == "value" + assert user_metadata(retrieved_generation) == {"dict": "value"} # Note: With OTEL, we can't verify exact start times from manually set timestamps, # as they are managed internally by the OTEL SDK @@ -1159,20 +1066,19 @@ def test_update_span(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id) # Verify trace properties - assert trace.name == "span" - assert len(trace.observations) == 1 + assert len(observations) == 1 # Verify span updates - retrieved_span = trace.observations[0] + retrieved_span = observations[0] + assert retrieved_span.trace_name == "span" assert retrieved_span.name == "span" assert retrieved_span.trace_id == trace_id - assert retrieved_span.metadata["dict"] == "value" + assert user_metadata(retrieved_span) == {"dict": "value"} def test_create_span_and_generation(): @@ -1197,17 +1103,16 @@ def test_create_span_and_generation(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id, min_count=2) # Verify trace details - assert len(trace.observations) == 2 + assert len(observations) == 2 # Find span and generation - spans = [obs for obs in trace.observations if obs.name == "span"] - generations = [obs for obs in trace.observations if obs.name == "generation"] + spans = [obs for obs in observations if obs.name == "span"] + generations = [obs for obs in observations if obs.name == "generation"] assert len(spans) == 1 assert len(generations) == 1 @@ -1222,7 +1127,6 @@ def test_create_span_and_generation(): def test_create_trace_with_id_and_generation(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() trace_name = create_uuid() @@ -1244,28 +1148,27 @@ def test_create_trace_with_id_and_generation(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = api_wrapper.get_trace(trace_id) + observations = wait_for_observations(trace_id, min_count=2) # Verify trace properties - assert trace["name"] == trace_name - assert trace["id"] == trace_id - assert len(trace["observations"]) == 2 # Parent span and generation + root = get_root_observation(observations) + assert root.trace_name == trace_name + assert root.trace_id == trace_id + assert len(observations) == 2 # Parent span and generation # Find the generation - generations = [obs for obs in trace["observations"] if obs["name"] == "generation"] + generations = [obs for obs in observations if obs.name == "generation"] assert len(generations) == 1 gen = generations[0] - assert gen["name"] == "generation" - assert gen["traceId"] == trace["id"] + assert gen.name == "generation" + assert gen.trace_id == trace_id def test_end_generation(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a generation generation = langfuse.start_observation( @@ -1294,21 +1197,14 @@ def test_end_generation(): langfuse.flush() # Retrieve and verify - trace = api_wrapper.get_trace( - trace_id, - is_result_ready=lambda trace: any( - obs["name"] == "query-generation" for obs in trace.get("observations", []) - ), - ) + observations = wait_for_observations(trace_id) # Find generation by name - generations = [ - obs for obs in trace["observations"] if obs["name"] == "query-generation" - ] + generations = [obs for obs in observations if obs.name == "query-generation"] assert len(generations) == 1 gen = generations[0] - assert gen["endTime"] is not None + assert gen.end_time is not None def test_end_generation_with_data(): @@ -1353,15 +1249,12 @@ def test_end_generation_with_data(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - fetched_trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id, min_count=2) # Find generation by name - generations = [ - obs for obs in fetched_trace.observations if obs.name == "query-generation" - ] + generations = [obs for obs in observations if obs.name == "query-generation"] assert len(generations) == 1 generation = generations[0] @@ -1371,7 +1264,7 @@ def test_end_generation_with_data(): 2023, 1, 1, 12, 3, tzinfo=timezone.utc ) assert generation.name == "query-generation" - assert generation.metadata["dict"] == "value" + assert user_metadata(generation) == {"dict": "value"} assert generation.level == "ERROR" assert generation.status_message == "Generation ended" assert generation.version == "1.0" @@ -1379,12 +1272,9 @@ def test_end_generation_with_data(): assert generation.model_parameters == {"param1": "value1", "param2": "value2"} assert generation.input == [{"test_input_key": "test_input_value"}] assert generation.output == {"test_output_key": "test_output_value"} - assert generation.usage.input == 100 - assert generation.usage.output == 200 - assert generation.usage.total == 500 - assert generation.calculated_input_cost == 111 - assert generation.calculated_output_cost == 222 - assert generation.calculated_total_cost == 444 + assert generation.usage_details == {"input": 100, "output": 200, "total": 500} + assert generation.cost_details == {"input": 111, "output": 222, "total": 444} + assert generation.total_cost == 444 def test_end_generation_with_openai_token_format(): @@ -1418,31 +1308,26 @@ def test_end_generation_with_openai_token_format(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id) # Find generation - generations = [obs for obs in trace.observations if obs.name == "query-generation"] + generations = [obs for obs in observations if obs.name == "query-generation"] assert len(generations) == 1 generation_api = generations[0] # Verify properties were converted correctly assert generation_api.end_time is not None - assert generation_api.usage.input == 100 # prompt_tokens mapped to input - assert generation_api.usage.output == 200 # completion_tokens mapped to output - assert generation_api.usage.total == 500 - assert generation_api.usage.unit == "TOKENS" # Default unit for OpenAI format - assert generation_api.calculated_input_cost == 111 - assert generation_api.calculated_output_cost == 222 - assert generation_api.calculated_total_cost == 444 + # OpenAI-style keys are mapped to input/output/total + assert generation_api.usage_details == {"input": 100, "output": 200, "total": 500} + assert generation_api.cost_details == {"input": 111, "output": 222, "total": 444} + assert generation_api.total_cost == 444 def test_end_span(): langfuse = Langfuse() - api_wrapper = LangfuseAPI() # Create a span span = langfuse.start_observation( @@ -1460,19 +1345,18 @@ def test_end_span(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = api_wrapper.get_trace(trace_id) + observations = wait_for_observations(trace_id) # Find span - spans = [obs for obs in trace["observations"] if obs["name"] == "span"] + spans = [obs for obs in observations if obs.name == "span"] assert len(spans) == 1 span_api = spans[0] # Verify end time was set - assert span_api["endTime"] is not None + assert span_api.end_time is not None def test_end_span_with_data(): @@ -1495,21 +1379,19 @@ def test_end_span_with_data(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id) # Find span - spans = [obs for obs in trace.observations if obs.name == "span"] + spans = [obs for obs in observations if obs.name == "span"] assert len(spans) == 1 span_api = spans[0] # Verify end time and metadata were updated assert span_api.end_time is not None - assert span_api.metadata["dict"] == "value" - assert span_api.metadata["interface"] == "whatsapp" + assert user_metadata(span_api) == {"dict": "value", "interface": "whatsapp"} def test_get_generations(): @@ -1535,16 +1417,15 @@ def test_get_generations(): # Ensure data is sent langfuse.flush() - sleep(3) # Fetch generations using API - generations = get_api().legacy.observations_v1.get_many(name=generation_name) + generations = wait_for_observations(name=generation_name) # Verify fetched generation matches what we created - assert len(generations.data) == 1 - assert generations.data[0].name == generation_name - assert generations.data[0].input == "great-prompt" - assert generations.data[0].output == "great-completion" + assert len(generations) == 1 + assert generations[0].name == generation_name + assert generations[0].input == "great-prompt" + assert generations[0].output == "great-completion" def test_get_generations_by_user(): @@ -1574,18 +1455,15 @@ def test_get_generations_by_user(): # Ensure data is sent langfuse.flush() - sleep(3) # Fetch generations by user ID using the API - generations = get_api().legacy.observations_v1.get_many( - user_id=user_id, type="GENERATION" - ) + generations = wait_for_observations(user_id=user_id, type="GENERATION") # Verify fetched generation matches what we created - assert len(generations.data) == 1 - assert generations.data[0].name == generation_name - assert generations.data[0].input == "great-prompt" - assert generations.data[0].output == "great-completion" + assert len(generations) == 1 + assert generations[0].name == generation_name + assert generations[0].input == "great-prompt" + assert generations[0].output == "great-completion" def test_kwargs(): @@ -1614,16 +1492,17 @@ def test_kwargs(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - observation = get_api().legacy.observations_v1.get(span_id) + observations = wait_for_observations(span.trace_id) + assert [observation.id for observation in observations] == [span_id] + observation = observations[0] # Verify kwargs were properly set as attributes assert observation.start_time is not None assert observation.input == {"key": "value"} assert observation.output == {"key": "value"} - assert observation.metadata["interface"] == "whatsapp" + assert user_metadata(observation) == {"interface": "whatsapp"} @pytest.mark.skip("Flaky") @@ -1660,16 +1539,13 @@ def test_timezone_awareness(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id, min_count=4) # Verify timestamps are in UTC regardless of local timezone - assert ( - len(trace.observations) == 4 - ) # Parent span, child span, generation, and event - for observation in trace.observations: + assert len(observations) == 4 # Parent span, child span, generation, and event + for observation in observations: # Check that start_time is within 5 seconds of current time delta = observation.start_time - utc_now assert delta.seconds < 5 @@ -1720,16 +1596,13 @@ def test_timezone_awareness_setting_timestamps(): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - trace = get_api().trace.get(trace_id) + observations = wait_for_observations(trace_id, min_count=4) # Verify timestamps are in UTC regardless of local timezone - assert ( - len(trace.observations) == 4 - ) # Parent span, child span, generation, and event - for observation in trace.observations: + assert len(observations) == 4 # Parent span, child span, generation, and event + for observation in observations: # Check that start_time is within 5 seconds of current time delta = abs((utc_now - observation.start_time).total_seconds()) assert delta < 5 @@ -1762,17 +1635,17 @@ def test_get_trace_by_session_id(): # Ensure data is sent langfuse.flush() - sleep(2) - # Retrieve the trace using the session_id - traces = get_api().trace.list(session_id=session_id) + # Retrieve the trace's observations using the session_id + observations = wait_for_observations(session_id=session_id) # Verify that the trace was retrieved correctly - assert len(traces.data) == 1 - retrieved_trace = traces.data[0] - assert retrieved_trace.name == trace_name - assert retrieved_trace.session_id == session_id - assert retrieved_trace.id == trace_id + assert len(observations) == 1 + retrieved_root = observations[0] + assert retrieved_root.is_root_observation is True + assert retrieved_root.trace_name == trace_name + assert retrieved_root.session_id == session_id + assert retrieved_root.trace_id == trace_id def test_fetch_trace(): @@ -1789,15 +1662,12 @@ def test_fetch_trace(): # Ensure data is sent langfuse.flush() - sleep(2) - # Fetch the trace using the get_api client - # Note: In the OTEL-based client, we use the API client directly - trace = get_api().trace.get(trace_id) + root = wait_for_root_observation(trace_id) # Verify trace properties - assert trace.id == trace_id - assert trace.name == name + assert root.trace_id == trace_id + assert root.trace_name == name def test_fetch_traces(): @@ -1812,7 +1682,7 @@ def test_fetch_traces(): # First trace with langfuse.start_as_current_observation(name="test1") as span: with propagate_attributes(trace_name=name, session_id="session-1"): - span.set_trace_io(input={"key": "value"}, output="output-value") + span.update(input={"key": "value"}, output="output-value") trace_ids.append(span.trace_id) sleep(1) # Ensure traces have different timestamps @@ -1820,7 +1690,7 @@ def test_fetch_traces(): # Second trace with langfuse.start_as_current_observation(name="test2") as span: with propagate_attributes(trace_name=name, session_id="session-1"): - span.set_trace_io(input={"key": "value"}, output="output-value") + span.update(input={"key": "value"}, output="output-value") trace_ids.append(span.trace_id) sleep(1) # Ensure traces have different timestamps @@ -1828,7 +1698,7 @@ def test_fetch_traces(): # Third trace with langfuse.start_as_current_observation(name="test3") as span: with propagate_attributes(trace_name=name, session_id="session-1"): - span.set_trace_io(input={"key": "value"}, output="output-value") + span.update(input={"key": "value"}, output="output-value") trace_ids.append(span.trace_id) # Ensure data is sent @@ -1836,38 +1706,50 @@ def test_fetch_traces(): expected_trace_ids = set(trace_ids) api = get_api(retry=False) + trace_name_filter = json.dumps( + [{"type": "string", "column": "traceName", "operator": "=", "value": name}] + ) - # Fetch all traces with the same name. - all_traces = wait_for_result( - lambda: api.trace.list(name=name, limit=10), - is_result_ready=lambda response: ( - {trace.id for trace in response.data} == expected_trace_ids + # Fetch the root observations of all traces with the same name. + roots = wait_for_observations( + filter=trace_name_filter, + is_result_ready=lambda observations: ( + {o.trace_id for o in observations} == expected_trace_ids ), ) # Verify we got all traces - assert len(all_traces.data) == 3 + assert len(roots) == 3 # Verify trace properties - for trace in all_traces.data: - assert trace.name == name - assert trace.session_id == "session-1" - assert trace.input == {"key": "value"} - assert trace.output == "output-value" - - # Test pagination by fetching the first three pages one at a time and - # confirming they collectively cover the created traces. - paginated_ids = set() - for page in range(1, 4): - paginated_response = wait_for_result( - lambda page=page: api.trace.list(name=name, limit=1, page=page), - is_result_ready=lambda response: ( - len(response.data) == 1 and response.data[0].id in expected_trace_ids - ), + for root in roots: + assert root.is_root_observation is True + assert root.trace_name == name + assert root.session_id == "session-1" + assert root.input == {"key": "value"} + assert root.output == "output-value" + + # Test cursor pagination by walking pages of one item and confirming they + # collectively cover the created traces. + paginated_ids = [] + cursor = None + for _ in range(3): + page = api.observations.get_many( + filter=trace_name_filter, limit=1, cursor=cursor + ) + assert len(page.data) == 1 + paginated_ids.append(page.data[0].trace_id) + cursor = page.meta.cursor + + assert set(paginated_ids) == expected_trace_ids + assert len(paginated_ids) == 3 + if cursor is not None: + assert ( + api.observations.get_many( + filter=trace_name_filter, limit=1, cursor=cursor + ).data + == [] ) - paginated_ids.add(paginated_response.data[0].id) - - assert paginated_ids == expected_trace_ids def test_get_observation(): @@ -1890,10 +1772,12 @@ def test_get_observation(): # Ensure data is sent langfuse.flush() - sleep(2) # Fetch the observation using the API - observation = get_api().legacy.observations_v1.get(generation_id) + observations = wait_for_observations(parent_span.trace_id, min_count=2) + matching = [o for o in observations if o.id == generation_id] + assert len(matching) == 1 + observation = matching[0] # Verify observation properties assert observation.id == generation_id @@ -1926,18 +1810,18 @@ def test_get_observations(): # Fetch observations using the API expected_generation_ids = {gen1_id, gen2_id} - observations = wait_for_result( - lambda: api.legacy.observations_v1.get_many(name=name, limit=10), - is_result_ready=lambda response: expected_generation_ids.issubset( - {obs.id for obs in response.data} + observations = wait_for_observations( + name=name, + is_result_ready=lambda observations: expected_generation_ids.issubset( + {obs.id for obs in observations} ), ) # Verify fetched observations - assert len(observations.data) == 2 + assert len(observations) == 2 # Filter for just the generations - generations = [obs for obs in observations.data if obs.type == "GENERATION"] + generations = [obs for obs in observations if obs.type == "GENERATION"] assert len(generations) == 2 # Verify the generation IDs match what we created @@ -1945,91 +1829,32 @@ def test_get_observations(): assert gen1_id in gen_ids assert gen2_id in gen_ids - # Test pagination by confirming both created generations can be reached - # across separate pages. - paginated_ids = set() - for page in range(1, 3): - paginated_response = wait_for_result( - lambda page=page: api.legacy.observations_v1.get_many( - name=name, limit=1, page=page - ), - is_result_ready=lambda response: ( - len(response.data) == 1 - and response.data[0].id in expected_generation_ids - ), - ) - paginated_ids.add(paginated_response.data[0].id) - - assert paginated_ids == expected_generation_ids - - -def test_get_trace_not_found(): - # Attempt to fetch a non-existent trace using the API - with pytest.raises(Exception): - get_api(retry=False).trace.get(create_uuid()) - + # Test cursor pagination by confirming both created generations can be + # reached across separate pages. + first_page = api.observations.get_many(name=name, limit=1) + assert len(first_page.data) == 1 + assert first_page.meta.cursor is not None + second_page = api.observations.get_many( + name=name, limit=1, cursor=first_page.meta.cursor + ) + assert len(second_page.data) == 1 -def test_get_observation_not_found(): - # Attempt to fetch a non-existent observation using the API - with pytest.raises(Exception): - get_api(retry=False).legacy.observations_v1.get(create_uuid()) + assert {first_page.data[0].id, second_page.data[0].id} == expected_generation_ids -def test_get_traces_empty(): - # Fetch traces with a filter that should return no results - response = get_api(retry=False).trace.list(name=create_uuid()) +def test_get_observations_for_unknown_trace_is_empty(): + response = get_api(retry=False).observations.get_many(trace_id=create_uuid()) - assert len(response.data) == 0 - assert response.meta.total_items == 0 + assert response.data == [] + assert response.meta.cursor is None def test_get_observations_empty(): # Fetch observations with a filter that should return no results - response = get_api(retry=False).legacy.observations_v1.get_many(name=create_uuid()) + response = get_api(retry=False).observations.get_many(name=create_uuid()) - assert len(response.data) == 0 - assert response.meta.total_items == 0 - - -def test_get_sessions(): - langfuse = Langfuse() - - # unique name - name = create_uuid() - session1 = create_uuid() - session2 = create_uuid() - session3 = create_uuid() - - # Create multiple traces with different session IDs - # Create first trace - with langfuse.start_as_current_observation(name=name): - with propagate_attributes(trace_name=name, session_id=session1): - pass - - # Create second trace - with langfuse.start_as_current_observation(name=name): - with propagate_attributes(trace_name=name, session_id=session2): - pass - - # Create third trace - with langfuse.start_as_current_observation(name=name): - with propagate_attributes(trace_name=name, session_id=session3): - pass - - langfuse.flush() - - # Fetch sessions - sleep(3) - response = get_api().sessions.list() - - # Assert the structure of the response, cannot check for the exact number of sessions as the table is not cleared between tests - assert hasattr(response, "data") - assert hasattr(response, "meta") - assert isinstance(response.data, list) - - # fetch only one, cannot check for the exact number of sessions as the table is not cleared between tests - response = get_api().sessions.list(limit=1, page=2) - assert len(response.data) == 1 + assert response.data == [] + assert response.meta.cursor is None @pytest.mark.skip( @@ -2037,7 +1862,6 @@ def test_get_sessions(): ) def test_create_trace_sampling_zero(): langfuse = Langfuse(sample_rate=0) - api_wrapper = LangfuseAPI() trace_name = create_uuid() # Create a span with trace properties - with sample_rate=0, this will not be sent to the API @@ -2061,12 +1885,9 @@ def test_create_trace_sampling_zero(): langfuse.flush() sleep(2) - # Try to fetch the trace - should fail as it wasn't sent to the API - fetched_trace = api_wrapper.get_trace(trace_id) - assert fetched_trace == { - "error": "LangfuseNotFoundError", - "message": f"Trace {trace_id} not found within authorized project", - } + # The trace's observations must not exist as they were never sent to the API + assert get_observations(trace_id=trace_id) == [] + assert get_scores(trace_id=trace_id) == [] def test_mask_function(request): @@ -2083,16 +1904,15 @@ def mask_func(data): return data langfuse = Langfuse(mask=mask_func) - api_wrapper = LangfuseAPI() # Create a root span with trace properties with langfuse.start_as_current_observation(name="test-span") as root_span: with propagate_attributes(trace_name="test_trace"): - root_span.set_trace_io(input={"sensitive": "data"}) + root_span.update(input={"sensitive": "data"}) # Get trace ID for later use trace_id = root_span.trace_id # Add output to the trace - root_span.set_trace_io(output={"more": "sensitive"}) + root_span.update(output={"more": "sensitive"}) # Create a generation as child gen = root_span.start_observation( @@ -2112,44 +1932,39 @@ def mask_func(data): # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - fetched_trace = api_wrapper.get_trace(trace_id) - assert fetched_trace["input"] == {"sensitive": "MASKED"} - assert fetched_trace["output"] == {"more": "MASKED"} + observations = wait_for_observations(trace_id, min_count=3) + fetched_root = get_root_observation(observations) + assert fetched_root.input == {"sensitive": "MASKED"} + assert fetched_root.output == {"more": "MASKED"} - fetched_gen = [ - o for o in fetched_trace["observations"] if o["type"] == "GENERATION" - ][0] - assert fetched_gen["input"] == {"prompt": "MASKED"} - assert fetched_gen["output"] == "MASKED" + fetched_gen = [o for o in observations if o.type == "GENERATION"][0] + assert fetched_gen.input == {"prompt": "MASKED"} + assert fetched_gen.output == "MASKED" fetched_span = [ - o - for o in fetched_trace["observations"] - if o["type"] == "SPAN" and o["name"] == "test_span" + o for o in observations if o.type == "SPAN" and o.name == "test_span" ][0] - assert fetched_span["input"] == {"data": "MASKED"} - assert fetched_span["output"] == "MASKED" + assert fetched_span.input == {"data": "MASKED"} + assert fetched_span.output == "MASKED" # Create a root span with trace properties with langfuse.start_as_current_observation(name="test-span") as root_span: with propagate_attributes(trace_name="test_trace"): - root_span.set_trace_io(input={"should_raise": "data"}) + root_span.update(input={"should_raise": "data"}) # Get trace ID for later use trace_id = root_span.trace_id # Add output to the trace - root_span.set_trace_io(output={"should_raise": "sensitive"}) + root_span.update(output={"should_raise": "sensitive"}) # Ensure data is sent langfuse.flush() - sleep(2) # Retrieve and verify - fetched_trace = api_wrapper.get_trace(trace_id) - assert fetched_trace["input"] == "" - assert fetched_trace["output"] == "" + fetched_root = wait_for_root_observation(trace_id) + assert fetched_root.input == "" + assert fetched_root.output == "" def test_get_project_id(): @@ -2218,13 +2033,11 @@ def test_start_as_current_observation_types(): pass langfuse.flush() - sleep(2) - api = get_api() - trace = api.trace.get(trace_id) + observations = wait_for_observations(trace_id, min_count=len(observation_types) + 1) # Check we have all expected observation types - found_types = {obs.type for obs in trace.observations} + found_types = {obs.type for obs in observations} expected_types = {obs_type.upper() for obs_type in observation_types} | { "SPAN" } # includes parent span @@ -2234,12 +2047,12 @@ def test_start_as_current_observation_types(): # Verify each specific observation exists for obs_type in observation_types: - observations = [ + matching = [ obs - for obs in trace.observations + for obs in observations if obs.name == f"test-{obs_type}" and obs.type == obs_type.upper() ] - assert len(observations) == 1, f"Expected one {obs_type.upper()} observation" + assert len(matching) == 1, f"Expected one {obs_type.upper()} observation" def test_that_generation_like_properties_are_actually_created(): @@ -2296,21 +2109,22 @@ def test_that_generation_like_properties_are_actually_created(): langfuse.flush() - api = get_api() - trace = api.trace.get(trace_id) + observations = wait_for_observations( + trace_id, min_count=len(generation_like_types) + 1 + ) # Verify that the properties are persisted in the API for generation-like types for obs_type in generation_like_types: - observations = [ + matching = [ obs - for obs in trace.observations + for obs in observations if obs.name == f"test-{obs_type}" and obs.type == obs_type.upper() ] - assert len(observations) == 1, ( - f"Expected one {obs_type.upper()} observation, but found {len(observations)}" + assert len(matching) == 1, ( + f"Expected one {obs_type.upper()} observation, but found {len(matching)}" ) - obs = observations[0] + obs = matching[0] assert obs.model == test_model, f"{obs_type} should have model property" assert obs.model_parameters == test_model_parameters, ( diff --git a/tests/e2e/test_decorators.py b/tests/e2e/test_decorators.py index 9fac4cf22..a33a4f3d5 100644 --- a/tests/e2e/test_decorators.py +++ b/tests/e2e/test_decorators.py @@ -15,7 +15,11 @@ from langfuse._client.resource_manager import LangfuseResourceManager from langfuse.langchain import CallbackHandler from langfuse.media import LangfuseMedia -from tests.support.utils import get_api, wait_for_trace +from tests.support.utils import ( + get_observations, + user_metadata, + wait_for_trace_snapshot, +) mock_metadata = {"key": "metadata"} mock_deep_metadata = {"key": "mock_deep_metadata"} @@ -100,7 +104,7 @@ def level_1_function(*args, **kwargs): assert result == "level_1" # Wrapped function returns correctly # ID setting for span or trace - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=3) assert len(trace_data.observations) == 3 # trace parameters if set anywhere in the call stack @@ -133,7 +137,7 @@ def level_1_function(*args, **kwargs): assert level_3_observation.name == "level_3" assert level_3_observation.metadata["key"] == mock_deep_metadata["key"] assert level_3_observation.type == "GENERATION" - assert level_3_observation.calculated_total_cost > 0 + assert level_3_observation.total_cost > 0 assert level_3_observation.output == "mock_output" assert level_3_observation.version == "version-1" @@ -182,7 +186,7 @@ def level_1_function(*args, **kwargs): assert result == "level_1" # Wrapped function returns correctly # ID setting for span or trace - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=3) assert len(trace_data.observations) == 3 # trace parameters if set anywhere in the call stack @@ -215,7 +219,7 @@ def level_1_function(*args, **kwargs): assert level_3_observation.name == "level_3" assert level_3_observation.metadata["key"] == mock_deep_metadata["key"] assert level_3_observation.type == "GENERATION" - assert level_3_observation.calculated_total_cost > 0 + assert level_3_observation.total_cost > 0 assert level_3_observation.output == "mock_output" assert level_3_observation.version == "version-1" @@ -260,7 +264,7 @@ def level_1_function(*args, **kwargs): langfuse.flush() - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=3) # trace parameters if set anywhere in the call stack assert trace_data.session_id == mock_session_id @@ -349,7 +353,7 @@ def level_1_function(*args, **kwargs): langfuse.flush() for mock_id in [mock_trace_id_1, mock_trace_id_2]: - trace_data = get_api().trace.get(mock_id) + trace_data = wait_for_trace_snapshot(mock_id, min_observations=3) assert len(trace_data.observations) == 3 # ID setting for span or trace @@ -382,7 +386,7 @@ def level_1_function(*args, **kwargs): assert level_3_observation.metadata["key"] == mock_deep_metadata["key"] assert level_3_observation.type == "GENERATION" - assert level_3_observation.calculated_total_cost > 0 + assert level_3_observation.total_cost > 0 def test_decorators_langchain(): @@ -427,7 +431,7 @@ def level_1_function(*args, **kwargs): langfuse.flush() - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( trace.session_id == mock_session_id @@ -522,8 +526,10 @@ def level_1_function(*args, **kwargs): assert result == "level_3" # Wrapped function returns correctly # ID setting for span or trace - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, + min_observations=3, + min_scores=3, is_result_ready=lambda trace: { "test-observation-score", "test-trace-score", @@ -553,7 +559,7 @@ def level_1_function(*args, **kwargs): assert any( [ score.name == "another-test-trace-score" - and score.string_value == "my_value" + and score.value == "my_value" and score.data_type == "CATEGORICAL" for score in trace_scores ] @@ -599,7 +605,7 @@ def function_with_circular_arg(circular_obj, *args, **kwargs): # Validate that the function executed as expected assert result == "function response" - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=1) assert ( trace_data.observations[0].input["args"][0]["reference"] == "CircularRefObject" @@ -632,7 +638,7 @@ def main(*args, **kwargs): assert result == "function response" - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) # Check that disabled capture_io doesn't capture manually set input/output assert len(trace_data.observations) == 2 @@ -702,7 +708,7 @@ def level_1_function(self, *args, **kwargs): assert result == "level_1" # Wrapped function returns correctly # ID setting for span or trace - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=4) assert len(trace_data.observations) == 4 # trace parameters if set anywhere in the call stack @@ -743,7 +749,7 @@ def level_1_function(self, *args, **kwargs): assert level_3_observation.name == "level_3_function" assert level_3_observation.metadata["key"] == mock_deep_metadata["key"] assert level_3_observation.type == "GENERATION" - assert level_3_observation.calculated_total_cost > 0 + assert level_3_observation.total_cost > 0 assert level_3_observation.output == "mock_output" @@ -779,7 +785,7 @@ def main(**kwargs): assert result == mock_output - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) # Find the main and nested observations adjacencies = defaultdict(list) @@ -833,7 +839,7 @@ async def main_async(**kwargs): assert result == mock_output - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) # Check correct nesting adjacencies = defaultdict(list) @@ -902,7 +908,7 @@ async def level_1_function(*args, **kwargs): assert result == "level_1" # Wrapped function returns correctly # ID setting for span or trace - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( trace.session_id == mock_session_id @@ -946,11 +952,11 @@ async def level_1_function(*args, **kwargs): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.usage.input is not None - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["input"] is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None print(generation) - assert generation.output == 2 + assert generation.output == "2" def test_generator_as_function_input(): @@ -982,7 +988,7 @@ def main(**kwargs): assert result == mock_output - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) nested_obs = next(o for o in trace_data.observations if o.name == "nested") @@ -1019,7 +1025,7 @@ def main(**kwargs): main(langfuse_trace_id=mock_trace_id) langfuse.flush() - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) # Find the observation with name 'nested' nested_observation = next(o for o in trace_data.observations if o.name == "nested") @@ -1052,7 +1058,7 @@ def function(): assert result == mock_output - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: any( observation.name == "function" and observation.output == mock_output @@ -1071,10 +1077,10 @@ def test_media(): media = LangfuseMedia(content_bytes=pdf_bytes, content_type="application/pdf") - @observe() + @observe(capture_input=False, capture_output=False) def main(): sleep(1) - langfuse.set_current_trace_io( + langfuse.update_current_span( input={ "context": { "nested": media, @@ -1099,7 +1105,7 @@ def main(): langfuse.flush() - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( "@@@langfuseMedia:type=application/pdf|id=" @@ -1155,19 +1161,20 @@ def main(): langfuse.flush() - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, - is_result_ready=lambda trace: ( - trace.metadata is not None - and trace.metadata.get("key1") == "value1" - and trace.metadata.get("key2") == "value2" - and trace.tags == ["tag1", "tag2"] - ), + min_observations=2, + is_result_ready=lambda trace: trace.tags == ["tag1", "tag2"], ) - assert trace_data.metadata["key1"] == "value1" - assert trace_data.metadata["key2"] == "value2" + # Trace metadata comes from the root; the nested propagation only reaches + # the nested observation. + assert trace_data.metadata == {"key1": "value1"} + nested_observation = _get_observation_by_name(trace_data, "nested") + assert user_metadata(nested_observation) == {"key1": "value1", "key2": "value2"} + assert nested_observation.tags == ["tag1", "tag2"] + assert trace_data.root.tags == ["tag1"] assert trace_data.tags == ["tag1", "tag2"] @@ -1227,7 +1234,7 @@ def level_1_function(*args, **kwargs): assert result == "level_1" # Verify trace was created properly - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( trace.name == mock_name and len(trace.observations) == 3 @@ -1287,7 +1294,7 @@ def level_1_function(*args, **kwargs): assert result == "level_4" - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( trace.name == mock_name @@ -1362,7 +1369,7 @@ def level_1_function(*args, **kwargs): assert result == "level_1" - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( trace.name == mock_name and len(trace.observations) == 2 @@ -1417,13 +1424,7 @@ def level_1_function(*args, **kwargs): # Should skip tracing entirely in multi-project setup without public key # This is expected behavior to prevent cross-project data leakage - try: - trace_data = get_api().trace.get(mock_trace_id) - # If trace is found, it should have no observations (tracing was skipped) - assert len(trace_data.observations) == 0 - except Exception: - # Trace not found is also expected - tracing was completely disabled - pass + assert get_observations(trace_id=mock_trace_id) == [] # Reset instances to not leak to other test suites removeMockResourceManagerInstances() @@ -1486,7 +1487,7 @@ async def async_level_1_function(*args, **kwargs): assert result == "async_level_3" # Verify trace was created properly - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( trace.name == mock_name @@ -1568,7 +1569,7 @@ async def async_level_1_function(*args, **kwargs): assert result == "sync_level_4" - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=4) assert len(trace_data.observations) == 4 assert trace_data.name == mock_name @@ -1647,8 +1648,8 @@ async def async_level_1_function(task_id, *args, **kwargs): assert result2 == "async_level_3_task_2" # Verify both traces were created correctly and didn't interfere - trace_data_1 = get_api().trace.get(trace_id_1) - trace_data_2 = get_api().trace.get(trace_id_2) + trace_data_1 = wait_for_trace_snapshot(trace_id_1, min_observations=3) + trace_data_2 = wait_for_trace_snapshot(trace_id_2, min_observations=3) assert trace_data_1.name == f"{mock_name}_task_1" assert trace_data_2.name == f"{mock_name}_task_2" @@ -1721,7 +1722,7 @@ async def async_consumer_function(): assert result == "Hello, Async World!" - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) assert len(trace_data.observations) == 2 assert trace_data.name == mock_name @@ -1786,7 +1787,7 @@ async def async_root_function(*args, **kwargs): assert result == "exception_handled" - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=3) assert len(trace_data.observations) == 3 assert trace_data.name == mock_name @@ -1851,11 +1852,11 @@ def root_function(): ) # Verify trace structure - trace_data = wait_for_trace( + trace_data = wait_for_trace_snapshot( mock_trace_id, is_result_ready=lambda trace: ( len(trace.observations) >= 2 - and {"parent_root", "child_stream"}.issubset( + and {"root", "sync_generator"}.issubset( { observation.name for observation in trace.observations @@ -1928,7 +1929,7 @@ async def root_function(): ) # Verify trace structure - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) assert len(trace_data.observations) == 2 # Verify both observations are present @@ -1996,7 +1997,7 @@ async def parent_function(): ) # Verify trace structure - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) assert len(trace_data.observations) == 2 # Check both observations exist @@ -2043,7 +2044,7 @@ async def root_function(): assert items == ["first_item"] # Verify trace structure - should have both observations despite exception - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) assert len(trace_data.observations) == 2 # Check that the failing generator observation has ERROR level @@ -2083,7 +2084,7 @@ def root_function(): assert items == [] # Verify trace structure - trace_data = get_api().trace.get(mock_trace_id) + trace_data = wait_for_trace_snapshot(mock_trace_id, min_observations=2) assert len(trace_data.observations) == 2 # Verify empty generator observation diff --git a/tests/e2e/test_media.py b/tests/e2e/test_media.py index d322e1788..ad8a8cb37 100644 --- a/tests/e2e/test_media.py +++ b/tests/e2e/test_media.py @@ -4,7 +4,7 @@ from langfuse._client.client import Langfuse from langfuse.media import LangfuseMedia -from tests.support.utils import wait_for_trace +from tests.support.utils import wait_for_observations def test_replace_media_reference_string_in_object(): @@ -30,25 +30,26 @@ def test_replace_media_reference_string_in_object(): langfuse.flush() - fetched_trace = wait_for_trace( + fetched_observations = wait_for_observations( span.trace_id, - is_result_ready=lambda trace: ( - bool(trace.observations) - and re.match( + is_result_ready=lambda observations: ( + re.match( r"^@@@langfuseMedia:type=audio/wav\|id=.+\|source=base64_data_uri@@@$", - trace.observations[0].metadata.get("context", {}).get("nested", ""), + observations[0].metadata.get("context", {}).get("nested", ""), ) is not None ), ) - media_ref = fetched_trace.observations[0].metadata["context"]["nested"] + assert len(fetched_observations) == 1 + fetched_observation = fetched_observations[0] + media_ref = fetched_observation.metadata["context"]["nested"] assert re.match( r"^@@@langfuseMedia:type=audio/wav\|id=.+\|source=base64_data_uri@@@$", media_ref, ) resolved_obs = langfuse.resolve_media_references( - obj=fetched_trace.observations[0], resolve_with="base64_data_uri" + obj=fetched_observation, resolve_with="base64_data_uri" ) expected_base64 = f"data:audio/wav;base64,{base64_audio}" @@ -61,15 +62,11 @@ def test_replace_media_reference_string_in_object(): langfuse.flush() - fetched_trace2 = wait_for_trace( + fetched_observations2 = wait_for_observations( span2.trace_id, - is_result_ready=lambda trace: ( - bool(trace.observations) - and trace.observations[0].metadata.get("context", {}).get("nested") - == fetched_trace.observations[0].metadata["context"]["nested"] + is_result_ready=lambda observations: ( + observations[0].metadata.get("context", {}).get("nested") == media_ref ), ) - assert ( - fetched_trace2.observations[0].metadata["context"]["nested"] - == fetched_trace.observations[0].metadata["context"]["nested"] - ) + assert len(fetched_observations2) == 1 + assert fetched_observations2[0].metadata["context"]["nested"] == media_ref diff --git a/tests/e2e/test_prompt.py b/tests/e2e/test_prompt.py index 6e113cb41..5b611d1fb 100644 --- a/tests/e2e/test_prompt.py +++ b/tests/e2e/test_prompt.py @@ -1,7 +1,7 @@ import pytest from langfuse._client.client import Langfuse -from tests.support.utils import create_uuid, get_api +from tests.support.utils import create_uuid, wait_for_observations def test_create_prompt(): @@ -429,18 +429,17 @@ def test_prompt_end_to_end(): langfuse.flush() - api = get_api() - - trace = api.trace.get(generation.trace_id) - - assert len(trace.observations) == 1 - - generation = trace.observations[0] - assert generation.prompt_id is not None + observations = wait_for_observations( + generation.trace_id, + is_result_ready=lambda observations: observations[0].prompt_id is not None, + ) - observation = api.legacy.observations_v1.get(generation.id) + assert len(observations) == 1 + observation = observations[0] assert observation.prompt_id is not None + assert observation.prompt_name == "test" + assert observation.prompt_version == prompt.version def test_do_not_return_fallback_if_fetch_success(): @@ -525,11 +524,10 @@ def test_do_not_link_observation_if_fallback(): ).end() langfuse.flush() - api = get_api() - trace = api.trace.get(generation.trace_id) + observations = wait_for_observations(generation.trace_id) - assert len(trace.observations) == 1 - assert trace.observations[0].prompt_id is None + assert len(observations) == 1 + assert observations[0].prompt_id is None def test_variable_names_on_content_with_variable_names(): diff --git a/tests/live_provider/test_langchain.py b/tests/live_provider/test_langchain.py index d04ddb81c..c1ebe8566 100644 --- a/tests/live_provider/test_langchain.py +++ b/tests/live_provider/test_langchain.py @@ -1,3 +1,4 @@ +import json import random import string import time @@ -18,7 +19,12 @@ from langfuse._client.client import Langfuse from langfuse.langchain import CallbackHandler -from tests.support.utils import create_uuid, encode_file_to_base64, get_api +from tests.support.utils import ( + create_uuid, + encode_file_to_base64, + wait_for_observations, + wait_for_trace_snapshot, +) def test_callback_generated_from_trace_chat(): @@ -44,7 +50,7 @@ def test_callback_generated_from_trace_chat(): langfuse.flush() - trace = get_api().trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) assert trace.input is None assert trace.output is None @@ -92,7 +98,7 @@ def test_callback_generated_from_lcel_chain(): langfuse.flush() - trace = get_api().trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) assert trace.input is None assert trace.output is None @@ -151,10 +157,10 @@ def test_basic_chat_openai(): # Ensure data is flushed to API sleep(2) - # Retrieve trace by name - traces = get_api().trace.list(name=test_name) - assert len(traces.data) > 0 - trace = get_api().trace.get(traces.data[0].id) + # Retrieve trace by its root observation, which carries the run name + roots = wait_for_observations(name=test_name) + assert len(roots) > 0 + trace = wait_for_trace_snapshot(roots[0].trace_id, min_observations=2) # Assertions assert trace.name == test_name @@ -193,7 +199,7 @@ def test_callback_simple_openai(): sleep(2) # Retrieve trace - trace = get_api().trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) # Assertions - add 1 for the wrapping span assert len(trace.observations) > 1 @@ -241,7 +247,7 @@ def test_callback_multiple_invocations_on_different_traces(): sleep(2) # Retrieve trace - trace = get_api().trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=3) # Add 1 to account for the wrapping span assert len(trace.observations) > 2 @@ -300,7 +306,7 @@ def test_openai_instruct_usage(): lf_handler._langfuse_client.flush() - observations = get_api().trace.get(trace_id).observations + observations = wait_for_observations(trace_id, min_count=3) assert len(observations) >= 3 assert any( @@ -318,7 +324,7 @@ def test_openai_instruct_usage(): assert observation.output != "" assert observation.input is not None assert observation.input != "" - assert observation.usage is not None + assert observation.usage_details is not None assert observation.usage_details["input"] is not None assert observation.usage_details["output"] is not None assert observation.usage_details["total"] is not None @@ -479,7 +485,7 @@ def test_link_langfuse_prompts_invoke(): langfuse_handler._langfuse_client.flush() sleep(2) - trace = get_api().trace.get(trace_id=trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=4) observations = trace.observations @@ -567,7 +573,12 @@ def test_link_langfuse_prompts_stream(): langfuse_handler._langfuse_client.flush() sleep(2) - trace = get_api().trace.get(trace_id=trace_id) + trace = wait_for_trace_snapshot( + trace_id, + is_result_ready=lambda trace: ( + len([o for o in trace.observations if o.type == "GENERATION"]) >= 4 + ), + ) observations = trace.observations @@ -653,11 +664,26 @@ def test_link_langfuse_prompts_batch(): langfuse_handler._langfuse_client.flush() - traces = get_api().trace.list(name=trace_name).data + trace_name_filter = json.dumps( + [ + { + "type": "string", + "column": "traceName", + "operator": "=", + "value": trace_name, + } + ] + ) + traced_observations = wait_for_observations(filter=trace_name_filter) - assert len(traces) == 1 + assert {o.trace_id for o in traced_observations} == {trace_id} - trace = get_api().trace.get(trace_id=trace_id) + trace = wait_for_trace_snapshot( + trace_id, + is_result_ready=lambda trace: ( + len([o for o in trace.observations if o.type == "GENERATION"]) >= 10 + ), + ) observations = trace.observations @@ -783,7 +809,7 @@ class GetWeather(BaseModel): handler._langfuse_client.flush() - trace = get_api().trace.get(trace_id=trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generations = list(filter(lambda x: x.type == "GENERATION", trace.observations)) assert len(generations) > 0 @@ -882,7 +908,7 @@ def test_multimodal(): handler._langfuse_client.flush() - trace = get_api().trace.get(trace_id=trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) assert len(trace.observations) >= 2 assert any( @@ -979,7 +1005,7 @@ def call_model(state: MessagesState): print(final_state["messages"][-1].content) handler._langfuse_client.flush() - trace = get_api().trace.get(trace_id=trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) assert len(trace.observations) > 0 @@ -1012,7 +1038,7 @@ def test_cached_token_usage(): handler._langfuse_client.flush() - trace = get_api().trace.get(handler.get_trace_id()) + trace = wait_for_trace_snapshot(handler.get_trace_id()) generation = next((o for o in trace.observations if o.type == "GENERATION")) diff --git a/tests/live_provider/test_langchain_integration.py b/tests/live_provider/test_langchain_integration.py index edb5455c4..dd6b166b2 100644 --- a/tests/live_provider/test_langchain_integration.py +++ b/tests/live_provider/test_langchain_integration.py @@ -7,7 +7,7 @@ from langfuse import Langfuse from langfuse.langchain import CallbackHandler -from tests.support.utils import create_uuid, get_api +from tests.support.utils import create_uuid, wait_for_trace_snapshot def _is_streaming_response(response): @@ -45,8 +45,7 @@ def test_stream_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -61,15 +60,15 @@ def test_stream_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output["content"] is not None assert generation.output["role"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -100,8 +99,7 @@ def test_stream_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -116,14 +114,14 @@ def test_stream_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -150,8 +148,7 @@ def test_invoke_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -165,15 +162,15 @@ def test_invoke_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output["content"] is not None assert generation.output["role"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -201,8 +198,7 @@ def test_invoke_in_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -216,14 +212,14 @@ def test_invoke_in_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert test_phrase in generation.output assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -251,8 +247,7 @@ def test_batch_in_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=3) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -266,13 +261,13 @@ def test_batch_in_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -300,8 +295,7 @@ def test_batch_in_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=3) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -314,13 +308,13 @@ def test_batch_in_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -354,8 +348,7 @@ async def test_astream_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -369,15 +362,15 @@ async def test_astream_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output["content"] is not None assert generation.output["role"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -411,8 +404,7 @@ async def test_astream_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -427,14 +419,14 @@ async def test_astream_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert test_phrase in generation.output assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -463,8 +455,7 @@ async def test_ainvoke_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -478,15 +469,15 @@ async def test_ainvoke_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output["content"] is not None assert generation.output["role"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -514,8 +505,7 @@ async def test_ainvoke_in_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=2) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -529,14 +519,14 @@ async def test_ainvoke_in_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert generation.metadata["tags"] == tags - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert test_phrase in generation.output assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -571,8 +561,7 @@ def test_chains_batch_in_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=9) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -585,13 +574,13 @@ def test_chains_batch_in_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -622,8 +611,7 @@ def test_chains_batch_in_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=9) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -636,13 +624,13 @@ def test_chains_batch_in_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -675,8 +663,7 @@ async def test_chains_abatch_in_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=9) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -689,13 +676,13 @@ async def test_chains_abatch_in_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -725,8 +712,7 @@ async def test_chains_abatch_in_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=9) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 assert len(trace.observations) == 9 @@ -738,13 +724,13 @@ async def test_chains_abatch_in_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -778,8 +764,7 @@ async def test_chains_ainvoke_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=5) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -792,15 +777,15 @@ async def test_chains_ainvoke_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output["content"] is not None assert generation.output["role"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -834,8 +819,7 @@ async def test_chains_ainvoke_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=5) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -848,13 +832,13 @@ async def test_chains_ainvoke_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -894,8 +878,7 @@ async def test_chains_astream_chat_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=5) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -910,15 +893,15 @@ async def test_chains_astream_chat_models(model_name): assert generation.model_parameters.get("max_completion_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.output["content"] is not None assert generation.output["role"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None @@ -956,8 +939,7 @@ async def test_chains_astream_completions_models(model_name): langfuse_client.flush() assert handler.runs == {} - api = get_api() - trace = api.trace.get(trace_id) + trace = wait_for_trace_snapshot(trace_id, min_observations=5) generationList = list(filter(lambda o: o.type == "GENERATION", trace.observations)) assert len(generationList) != 0 @@ -972,11 +954,11 @@ async def test_chains_astream_completions_models(model_name): assert generation.model_parameters.get("max_tokens") is not None assert generation.model_parameters.get("temperature") is not None assert all(x in generation.metadata["tags"] for x in tags) - assert generation.usage.output is not None - assert generation.usage.total is not None + assert generation.usage_details["output"] is not None + assert generation.usage_details["total"] is not None assert generation.input_price is not None assert generation.output_price is not None - assert generation.calculated_input_cost is not None - assert generation.calculated_output_cost is not None - assert generation.calculated_total_cost is not None + assert generation.cost_details["input"] is not None + assert generation.cost_details["output"] is not None + assert generation.total_cost is not None assert generation.latency is not None diff --git a/tests/live_provider/test_openai.py b/tests/live_provider/test_openai.py index 43a07be23..18dc61b81 100644 --- a/tests/live_provider/test_openai.py +++ b/tests/live_provider/test_openai.py @@ -6,7 +6,11 @@ from pydantic import BaseModel from langfuse._client.client import Langfuse -from tests.support.utils import create_uuid, encode_file_to_base64, get_api +from tests.support.utils import ( + create_uuid, + encode_file_to_base64, + wait_for_observations, +) langfuse: Langfuse | None = None @@ -58,27 +62,25 @@ def test_openai_chat_completion(openai): sleep(1) - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" assert len(completion.choices) != 0 - assert generation.data[0].input == [ + assert generation[0].input == [ { "content": "You are an expert mathematician", "role": "assistant", }, {"content": "1 + 1 = ", "role": "user"}, ] - assert generation.data[0].type == "GENERATION" - assert "gpt-3.5-turbo-0125" in generation.data[0].model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].type == "GENERATION" + assert "gpt-3.5-turbo-0125" in generation[0].model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 0, "top_p": 1, @@ -86,11 +88,11 @@ def test_openai_chat_completion(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert "2" in generation.data[0].output["content"] - assert generation.data[0].output["role"] == "assistant" + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert "2" in generation[0].output["content"] + assert generation[0].output["role"] == "assistant" def test_openai_chat_completion_stream(openai): @@ -116,21 +118,19 @@ def test_openai_chat_completion_stream(openai): langfuse.flush() sleep(3) - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert "gpt-3.5-turbo-0125" in generation.data[0].model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert "gpt-3.5-turbo-0125" in generation[0].model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 0, "top_p": 1, @@ -138,16 +138,16 @@ def test_openai_chat_completion_stream(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].output == 2 - assert generation.data[0].completion_start_time is not None + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert generation[0].output == "2" + assert generation[0].completion_start_time is not None # Completion start time for time-to-first-token - assert generation.data[0].completion_start_time is not None - assert generation.data[0].completion_start_time >= generation.data[0].start_time - assert generation.data[0].completion_start_time <= generation.data[0].end_time + assert generation[0].completion_start_time is not None + assert generation[0].completion_start_time >= generation[0].start_time + assert generation[0].completion_start_time <= generation[0].end_time def test_openai_chat_completion_stream_with_next_iteration(openai): @@ -177,21 +177,19 @@ def test_openai_chat_completion_stream_with_next_iteration(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-3.5-turbo-0125" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-3.5-turbo-0125" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 0, "top_p": 1, @@ -199,16 +197,16 @@ def test_openai_chat_completion_stream_with_next_iteration(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].output == 2 - assert generation.data[0].completion_start_time is not None + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert generation[0].output == "2" + assert generation[0].completion_start_time is not None # Completion start time for time-to-first-token - assert generation.data[0].completion_start_time is not None - assert generation.data[0].completion_start_time >= generation.data[0].start_time - assert generation.data[0].completion_start_time <= generation.data[0].end_time + assert generation[0].completion_start_time is not None + assert generation[0].completion_start_time >= generation[0].start_time + assert generation[0].completion_start_time <= generation[0].end_time def test_openai_chat_completion_stream_fail(openai): @@ -227,33 +225,28 @@ def test_openai_chat_completion_stream_fail(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "fake" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" + + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert generation[0].model == "fake" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "temperature": 0, "top_p": 1, "frequency_penalty": 0, "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].level == "ERROR" - assert generation.data[0].status_message is not None - assert generation.data[0].output is None + assert generation[0].level == "ERROR" + assert generation[0].status_message is not None + assert generation[0].output is None openai.api_key = os.environ["OPENAI_API_KEY"] @@ -277,13 +270,11 @@ def test_openai_chat_completion_with_langfuse_prompt(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert isinstance(generation.data[0].prompt_id, str) + assert len(generation) != 0 + assert generation[0].name == generation_name + assert isinstance(generation[0].prompt_id, str) def test_openai_chat_completion_fail(openai): @@ -300,29 +291,27 @@ def test_openai_chat_completion_fail(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) - - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "fake" - assert generation.data[0].level == "ERROR" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].status_message is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert generation[0].model == "fake" + assert generation[0].level == "ERROR" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].status_message is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "temperature": 0, "top_p": 1, "frequency_penalty": 0, "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].output is None + assert generation[0].output is None openai.api_key = os.environ["OPENAI_API_KEY"] @@ -360,25 +349,21 @@ def test_openai_chat_completion_two_calls(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name + assert len(generation) != 0 + assert generation[0].name == generation_name assert len(completion.choices) != 0 - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] - generation_2 = get_api().legacy.observations_v1.get_many( - name=generation_name_2, type="GENERATION" - ) + generation_2 = wait_for_observations(name=generation_name_2, type="GENERATION") - assert len(generation_2.data) != 0 - assert generation_2.data[0].name == generation_name_2 + assert len(generation_2) != 0 + assert generation_2[0].name == generation_name_2 assert len(completion_2.choices) != 0 - assert generation_2.data[0].input == [{"content": "2 + 2 = ", "role": "user"}] + assert generation_2[0].input == [{"content": "2 + 2 = ", "role": "user"}] def test_openai_chat_completion_with_seed(openai): @@ -394,11 +379,9 @@ def test_openai_chat_completion_with_seed(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert generation.data[0].model_parameters == { + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 0, "top_p": 1, @@ -423,32 +406,30 @@ def test_openai_completion(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" assert len(completion.choices) != 0 - assert completion.choices[0].text == generation.data[0].output - assert generation.data[0].input == "1 + 1 = " - assert generation.data[0].type == "GENERATION" - assert "gpt-3.5-turbo-instruct" in generation.data[0].model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert completion.choices[0].text == generation[0].output + assert generation[0].input == "1 + 1 = " + assert generation[0].type == "GENERATION" + assert "gpt-3.5-turbo-instruct" in generation[0].model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "temperature": 0, "top_p": 1, "frequency_penalty": 0, "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].output == "2\n\n1 + 2 = 3\n\n2 + 3 = " + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert generation[0].output == "2\n\n1 + 2 = 3\n\n2 + 3 = " @requires_legacy_completion_model @@ -472,37 +453,35 @@ def test_openai_completion_stream(openai): assert len(content) > 0 - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - - assert generation.data[0].input == "1 + 1 = " - assert generation.data[0].type == "GENERATION" - assert "gpt-3.5-turbo-instruct" in generation.data[0].model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" + + assert generation[0].input == "1 + 1 = " + assert generation[0].type == "GENERATION" + assert "gpt-3.5-turbo-instruct" in generation[0].model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "temperature": 0, "top_p": 1, "frequency_penalty": 0, "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].output == "2\n\n1 + 2 = 3\n\n2 + 3 = " - assert generation.data[0].completion_start_time is not None + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert generation[0].output == "2\n\n1 + 2 = 3\n\n2 + 3 = " + assert generation[0].completion_start_time is not None # Completion start time for time-to-first-token - assert generation.data[0].completion_start_time is not None - assert generation.data[0].completion_start_time >= generation.data[0].start_time - assert generation.data[0].completion_start_time <= generation.data[0].end_time + assert generation[0].completion_start_time is not None + assert generation[0].completion_start_time >= generation[0].start_time + assert generation[0].completion_start_time <= generation[0].end_time def test_openai_completion_fail(openai): @@ -521,29 +500,27 @@ def test_openai_completion_fail(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) - - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - assert generation.data[0].input == "1 + 1 = " - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "fake" - assert generation.data[0].level == "ERROR" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].status_message is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" + assert generation[0].input == "1 + 1 = " + assert generation[0].type == "GENERATION" + assert generation[0].model == "fake" + assert generation[0].level == "ERROR" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].status_message is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "temperature": 0, "top_p": 1, "frequency_penalty": 0, "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].output is None + assert generation[0].output is None openai.api_key = os.environ["OPENAI_API_KEY"] @@ -564,33 +541,28 @@ def test_openai_completion_stream_fail(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - - assert generation.data[0].input == "1 + 1 = " - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-3.5-turbo" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].input == "1 + 1 = " + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-3.5-turbo" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "temperature": 0, "top_p": 1, "frequency_penalty": 0, "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].level == "ERROR" - assert generation.data[0].status_message is not None - assert generation.data[0].output is None + assert generation[0].level == "ERROR" + assert generation[0].status_message is not None + assert generation[0].output is None openai.api_key = os.environ["OPENAI_API_KEY"] @@ -614,13 +586,11 @@ def test_openai_completion_with_langfuse_prompt(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert isinstance(generation.data[0].prompt_id, str) + assert len(generation) != 0 + assert generation[0].name == generation_name + assert isinstance(generation[0].prompt_id, str) def test_fails_wrong_name(openai): @@ -656,21 +626,19 @@ async def test_async_chat(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name + assert len(generation) != 0 + assert generation[0].name == generation_name assert len(completion.choices) != 0 - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-3.5-turbo-0125" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-3.5-turbo-0125" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 1, "top_p": 1, @@ -678,11 +646,11 @@ async def test_async_chat(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert "2" in generation.data[0].output["content"] - assert generation.data[0].output["role"] == "assistant" + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert "2" in generation[0].output["content"] + assert generation[0].output["role"] == "assistant" @pytest.mark.asyncio @@ -703,19 +671,17 @@ async def test_async_chat_stream(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) - - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-3.5-turbo-0125" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-3.5-turbo-0125" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 1, "top_p": 1, @@ -723,15 +689,15 @@ async def test_async_chat_stream(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert "2" in str(generation.data[0].output) + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert "2" in str(generation[0].output) # Completion start time for time-to-first-token - assert generation.data[0].completion_start_time is not None - assert generation.data[0].completion_start_time >= generation.data[0].start_time - assert generation.data[0].completion_start_time <= generation.data[0].end_time + assert generation[0].completion_start_time is not None + assert generation[0].completion_start_time >= generation[0].start_time + assert generation[0].completion_start_time <= generation[0].end_time @pytest.mark.asyncio @@ -762,21 +728,19 @@ async def test_async_chat_stream_with_anext(openai): print(result) - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].input == [ + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].input == [ {"content": "Give me a one-liner joke", "role": "user"} ] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-3.5-turbo-0125" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-3.5-turbo-0125" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 1, "top_p": 1, @@ -784,14 +748,14 @@ async def test_async_chat_stream_with_anext(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None # Completion start time for time-to-first-token - assert generation.data[0].completion_start_time is not None - assert generation.data[0].completion_start_time >= generation.data[0].start_time - assert generation.data[0].completion_start_time <= generation.data[0].end_time + assert generation[0].completion_start_time is not None + assert generation[0].completion_start_time >= generation[0].start_time + assert generation[0].completion_start_time <= generation[0].end_time def test_openai_function_call(openai): @@ -825,14 +789,12 @@ class StepByStepAIResponse(BaseModel): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].output is not None - assert "function_call" in generation.data[0].output + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].output is not None + assert "function_call" in generation[0].output assert output["title"] is not None @@ -869,14 +831,12 @@ class StepByStepAIResponse(BaseModel): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].output is not None - assert "function_call" in generation.data[0].output + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].output is not None + assert "function_call" in generation[0].output def test_openai_tool_call(openai): @@ -913,21 +873,17 @@ def test_openai_tool_call(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name + assert len(generation) != 0 + assert generation[0].name == generation_name assert ( - generation.data[0].output["tool_calls"][0]["function"]["name"] + generation[0].output["tool_calls"][0]["function"]["name"] == "get_current_weather" ) - assert ( - generation.data[0].output["tool_calls"][0]["function"]["arguments"] is not None - ) - assert generation.data[0].input["tools"] == tools - assert generation.data[0].input["messages"] == messages + assert generation[0].output["tool_calls"][0]["function"]["arguments"] is not None + assert generation[0].input["tools"] == tools + assert generation[0].input["messages"] == messages def test_openai_tool_call_streamed(openai): @@ -969,22 +925,18 @@ def test_openai_tool_call_streamed(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name + assert len(generation) != 0 + assert generation[0].name == generation_name assert ( - generation.data[0].output["tool_calls"][0]["function"]["name"] + generation[0].output["tool_calls"][0]["function"]["name"] == "get_current_weather" ) - assert ( - generation.data[0].output["tool_calls"][0]["function"]["arguments"] is not None - ) - assert generation.data[0].input["tools"] == tools - assert generation.data[0].input["messages"] == messages + assert generation[0].output["tool_calls"][0]["function"]["arguments"] is not None + assert generation[0].input["tools"] == tools + assert generation[0].input["messages"] == messages def test_langchain_integration(openai): @@ -1047,28 +999,28 @@ def test_structured_output_response_format_kwarg(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" + generation = wait_for_observations( + name=generation_name, type="GENERATION", expand_metadata="response_format" ) - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].metadata["someKey"] == "someResponse" - assert generation.data[0].metadata["response_format"] == { + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].metadata["someKey"] == "someResponse" + assert generation[0].metadata["response_format"] == { "type": "json_schema", "json_schema": json_schema, } - assert generation.data[0].input == [ + assert generation[0].input == [ {"role": "system", "content": "You are a helpful math tutor."}, {"content": "solve 8x + 31 = 2", "role": "user"}, ] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-4o-2024-08-06" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-4o-2024-08-06" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 1, "top_p": 1, @@ -1076,10 +1028,10 @@ def test_structured_output_response_format_kwarg(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert generation.data[0].output["role"] == "assistant" + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert generation[0].output["role"] == "assistant" def test_structured_output_beta_completions_parse(openai): @@ -1117,31 +1069,29 @@ class CalendarEvent(BaseModel): if Version(openai.__version__) >= Version("1.50.0"): # Check the trace and observation properties - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) == 1 - assert generation.data[0].name == generation_name - assert generation.data[0].type == "GENERATION" - assert "gpt-4o" in generation.data[0].model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time + assert len(generation) == 1 + assert generation[0].name == generation_name + assert generation[0].type == "GENERATION" + assert "gpt-4o" in generation[0].model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time # Check input and output - assert len(generation.data[0].input) == 2 - assert generation.data[0].input[0]["role"] == "system" - assert generation.data[0].input[1]["role"] == "user" - assert isinstance(generation.data[0].output, dict) - assert "name" in generation.data[0].output["content"] - assert "date" in generation.data[0].output["content"] - assert "participants" in generation.data[0].output["content"] + assert len(generation[0].input) == 2 + assert generation[0].input[0]["role"] == "system" + assert generation[0].input[1]["role"] == "user" + assert isinstance(generation[0].output, dict) + assert "name" in generation[0].output["content"] + assert "date" in generation[0].output["content"] + assert "participants" in generation[0].output["content"] # Check usage - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None @pytest.mark.asyncio @@ -1163,19 +1113,17 @@ async def test_close_async_stream(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) - - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].input == [{"content": "1 + 1 = ", "role": "user"}] - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == "gpt-3.5-turbo-0125" - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].model_parameters == { + generation = wait_for_observations(name=generation_name, type="GENERATION") + + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].input == [{"content": "1 + 1 = ", "role": "user"}] + assert generation[0].type == "GENERATION" + assert generation[0].model == "gpt-3.5-turbo-0125" + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].model_parameters == { "service_tier": "default", "temperature": 1, "top_p": 1, @@ -1183,15 +1131,15 @@ async def test_close_async_stream(openai): "max_tokens": "Infinity", "presence_penalty": 0, } - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert "2" in str(generation.data[0].output) + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert "2" in str(generation[0].output) # Completion start time for time-to-first-token - assert generation.data[0].completion_start_time is not None - assert generation.data[0].completion_start_time >= generation.data[0].start_time - assert generation.data[0].completion_start_time <= generation.data[0].end_time + assert generation[0].completion_start_time is not None + assert generation[0].completion_start_time >= generation[0].start_time + assert generation[0].completion_start_time <= generation[0].end_time def test_base_64_image_input(openai): @@ -1225,26 +1173,24 @@ def test_base_64_image_input(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name - assert generation.data[0].input[0]["content"][0]["text"] == "What’s in this image?" + assert len(generation) != 0 + assert generation[0].name == generation_name + assert generation[0].input[0]["content"][0]["text"] == "What’s in this image?" assert ( f"@@@langfuseMedia:type={content_type}|id=" - in generation.data[0].input[0]["content"][1]["image_url"]["url"] + in generation[0].input[0]["content"][1]["image_url"]["url"] ) - assert generation.data[0].type == "GENERATION" - assert "gpt-4o-mini" in generation.data[0].model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - assert "dog" in generation.data[0].output["content"] + assert generation[0].type == "GENERATION" + assert "gpt-4o-mini" in generation[0].model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + assert "dog" in generation[0].output["content"] def test_audio_input_and_output(openai): @@ -1277,32 +1223,28 @@ def test_audio_input_and_output(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - assert generation.data[0].name == generation_name + assert len(generation) != 0 + assert generation[0].name == generation_name assert ( - generation.data[0].input[0]["content"][0]["text"] - == "Do what this recording says." + generation[0].input[0]["content"][0]["text"] == "Do what this recording says." ) assert ( "@@@langfuseMedia:type=audio/wav|id=" - in generation.data[0].input[0]["content"][1]["input_audio"]["data"] - ) - assert generation.data[0].type == "GENERATION" - assert generation.data[0].model == model - assert generation.data[0].start_time is not None - assert generation.data[0].end_time is not None - assert generation.data[0].start_time < generation.data[0].end_time - assert generation.data[0].usage.input is not None - assert generation.data[0].usage.output is not None - assert generation.data[0].usage.total is not None - print(generation.data[0].output) + in generation[0].input[0]["content"][1]["input_audio"]["data"] + ) + assert generation[0].type == "GENERATION" + assert generation[0].model == model + assert generation[0].start_time is not None + assert generation[0].end_time is not None + assert generation[0].start_time < generation[0].end_time + assert generation[0].usage_details["input"] is not None + assert generation[0].usage_details["output"] is not None + assert generation[0].usage_details["total"] is not None + print(generation[0].output) assert ( - "@@@langfuseMedia:type=audio/wav|id=" - in generation.data[0].output["audio"]["data"] + "@@@langfuseMedia:type=audio/wav|id=" in generation[0].output["audio"]["data"] ) @@ -1317,25 +1259,22 @@ def test_response_api_text_input(openai): ) langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - generationData = generation.data[0] + assert len(generation) != 0 + generationData = generation[0] assert generationData.name == generation_name assert ( - generation.data[0].input - == "Tell me a three sentence bedtime story about a unicorn." + generation[0].input == "Tell me a three sentence bedtime story about a unicorn." ) assert generationData.type == "GENERATION" assert "gpt-4o" in generationData.model assert generationData.start_time is not None assert generationData.end_time is not None assert generationData.start_time < generationData.end_time - assert generationData.usage.input is not None - assert generationData.usage.output is not None - assert generationData.usage.total is not None + assert generationData.usage_details["input"] is not None + assert generationData.usage_details["output"] is not None + assert generationData.usage_details["total"] is not None assert generationData.output is not None @@ -1363,22 +1302,20 @@ def test_response_api_image_input(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - generationData = generation.data[0] + assert len(generation) != 0 + generationData = generation[0] assert generationData.name == generation_name - assert generation.data[0].input[0]["content"][0]["text"] == "what is in this image?" + assert generation[0].input[0]["content"][0]["text"] == "what is in this image?" assert generationData.type == "GENERATION" assert "gpt-4o" in generationData.model assert generationData.start_time is not None assert generationData.end_time is not None assert generationData.start_time < generationData.end_time - assert generationData.usage.input is not None - assert generationData.usage.output is not None - assert generationData.usage.total is not None + assert generationData.usage_details["input"] is not None + assert generationData.usage_details["output"] is not None + assert generationData.usage_details["total"] is not None assert generationData.output is not None @@ -1395,12 +1332,10 @@ def test_response_api_web_search(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - generationData = generation.data[0] + assert len(generation) != 0 + generationData = generation[0] assert generationData.name == generation_name assert generationData.input == { "input": "What was a positive news story from today?", @@ -1411,9 +1346,9 @@ def test_response_api_web_search(openai): assert generationData.start_time is not None assert generationData.end_time is not None assert generationData.start_time < generationData.end_time - assert generationData.usage.input is not None - assert generationData.usage.output is not None - assert generationData.usage.total is not None + assert generationData.usage_details["input"] is not None + assert generationData.usage_details["output"] is not None + assert generationData.usage_details["total"] is not None assert generationData.output is not None assert generationData.metadata is not None @@ -1435,14 +1370,12 @@ def test_response_api_streaming(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - generationData = generation.data[0] + assert len(generation) != 0 + generationData = generation[0] assert generationData.name == generation_name - assert generation.data[0].input == [ + assert generation[0].input == [ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Hello!"}, ] @@ -1451,9 +1384,9 @@ def test_response_api_streaming(openai): assert generationData.start_time is not None assert generationData.end_time is not None assert generationData.start_time < generationData.end_time - assert generationData.usage.input is not None - assert generationData.usage.output is not None - assert generationData.usage.total is not None + assert generationData.usage_details["input"] is not None + assert generationData.usage_details["output"] is not None + assert generationData.usage_details["total"] is not None assert generationData.output is not None assert generationData.metadata is not None assert generationData.metadata["instructions"] == "You are a helpful assistant." @@ -1492,14 +1425,12 @@ def test_response_api_functions(openai): langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - generationData = generation.data[0] + assert len(generation) != 0 + generationData = generation[0] assert generationData.name == generation_name - assert generation.data[0].input == { + assert generation[0].input == { "input": "What is the weather like in Boston today?", "tools": tools, "tool_choice": "auto", @@ -1509,9 +1440,9 @@ def test_response_api_functions(openai): assert generationData.start_time is not None assert generationData.end_time is not None assert generationData.start_time < generationData.end_time - assert generationData.usage.input is not None - assert generationData.usage.output is not None - assert generationData.usage.total is not None + assert generationData.usage_details["input"] is not None + assert generationData.usage_details["output"] is not None + assert generationData.usage_details["total"] is not None assert generationData.output is not None assert generationData.metadata is not None @@ -1528,22 +1459,20 @@ def test_response_api_reasoning(openai): ) langfuse.flush() - generation = get_api().legacy.observations_v1.get_many( - name=generation_name, type="GENERATION" - ) + generation = wait_for_observations(name=generation_name, type="GENERATION") - assert len(generation.data) != 0 - generationData = generation.data[0] + assert len(generation) != 0 + generationData = generation[0] assert generationData.name == generation_name - assert generation.data[0].input == "How much wood would a woodchuck chuck?" + assert generation[0].input == "How much wood would a woodchuck chuck?" assert generationData.type == "GENERATION" assert "o3-mini" in generationData.model assert generationData.start_time is not None assert generationData.end_time is not None assert generationData.start_time < generationData.end_time - assert generationData.usage.input is not None - assert generationData.usage.output is not None - assert generationData.usage.total is not None + assert generationData.usage_details["input"] is not None + assert generationData.usage_details["output"] is not None + assert generationData.usage_details["total"] is not None assert generationData.output is not None assert generationData.metadata is not None @@ -1560,12 +1489,10 @@ def test_openai_embeddings(openai): langfuse.flush() sleep(1) - embedding = get_api().legacy.observations_v1.get_many( - name=embedding_name, type="EMBEDDING" - ) + embedding = wait_for_observations(name=embedding_name, type="EMBEDDING") - assert len(embedding.data) != 0 - embedding_data = embedding.data[0] + assert len(embedding) != 0 + embedding_data = embedding[0] assert embedding_data.name == embedding_name assert embedding_data.metadata["test_key"] == "test_value" assert embedding_data.input == "The quick brown fox jumps over the lazy dog" @@ -1574,8 +1501,8 @@ def test_openai_embeddings(openai): assert embedding_data.start_time is not None assert embedding_data.end_time is not None assert embedding_data.start_time < embedding_data.end_time - assert embedding_data.usage.input is not None - assert embedding_data.usage.total is not None + assert embedding_data.usage_details["input"] is not None + assert embedding_data.usage_details["total"] is not None assert embedding_data.output is not None assert "dimensions" in embedding_data.output assert "count" in embedding_data.output @@ -1596,18 +1523,16 @@ def test_openai_embeddings_multiple_inputs(openai): langfuse.flush() sleep(1) - embedding = get_api().legacy.observations_v1.get_many( - name=embedding_name, type="EMBEDDING" - ) + embedding = wait_for_observations(name=embedding_name, type="EMBEDDING") - assert len(embedding.data) != 0 - embedding_data = embedding.data[0] + assert len(embedding) != 0 + embedding_data = embedding[0] assert embedding_data.name == embedding_name assert embedding_data.input == inputs assert embedding_data.type == "EMBEDDING" assert "text-embedding-ada-002" in embedding_data.model - assert embedding_data.usage.input is not None - assert embedding_data.usage.total is not None + assert embedding_data.usage_details["input"] is not None + assert embedding_data.usage_details["total"] is not None assert embedding_data.output["count"] == len(inputs) @@ -1629,16 +1554,14 @@ async def test_async_openai_embeddings(openai): langfuse.flush() sleep(1) - embedding = get_api().legacy.observations_v1.get_many( - name=embedding_name, type="EMBEDDING" - ) + embedding = wait_for_observations(name=embedding_name, type="EMBEDDING") - assert len(embedding.data) != 0 - embedding_data = embedding.data[0] + assert len(embedding) != 0 + embedding_data = embedding[0] assert embedding_data.name == embedding_name assert embedding_data.input == "Async embedding test" assert embedding_data.type == "EMBEDDING" assert "text-embedding-ada-002" in embedding_data.model assert embedding_data.metadata["async"] is True - assert embedding_data.usage.input is not None - assert embedding_data.usage.total is not None + assert embedding_data.usage_details["input"] is not None + assert embedding_data.usage_details["total"] is not None diff --git a/tests/support/api_wrapper.py b/tests/support/api_wrapper.py deleted file mode 100644 index c4519252f..000000000 --- a/tests/support/api_wrapper.py +++ /dev/null @@ -1,130 +0,0 @@ -import os - -import httpx - -from langfuse.api.commons.errors.not_found_error import NotFoundError -from tests.support.retry import ( - DEFAULT_RETRY_INTERVAL_SECONDS, - DEFAULT_RETRY_TIMEOUT_SECONDS, - is_not_found_payload, - retry_until_ready, -) - - -class LangfuseAPI: - def __init__(self, username=None, password=None, base_url=None): - username = username if username else os.environ["LANGFUSE_PUBLIC_KEY"] - password = password if password else os.environ["LANGFUSE_SECRET_KEY"] - self.auth = (username, password) - self.BASE_URL = base_url if base_url else os.environ["LANGFUSE_BASE_URL"] - - def _get_json( - self, - url, - params=None, - *, - retry=True, - is_result_ready=None, - timeout_seconds=DEFAULT_RETRY_TIMEOUT_SECONDS, - interval_seconds=DEFAULT_RETRY_INTERVAL_SECONDS, - ): - def _request(): - response = httpx.get(url, params=params, auth=self.auth) - payload = response.json() - - if response.status_code == 404 and is_not_found_payload(payload): - raise NotFoundError(body=payload, headers=dict(response.headers)) - - return payload - - if not retry: - return _request() - - return retry_until_ready( - _request, - is_result_ready=is_result_ready, - timeout_seconds=timeout_seconds, - interval_seconds=interval_seconds, - ) - - def get_observation( - self, - observation_id, - *, - retry=True, - is_result_ready=None, - timeout_seconds=DEFAULT_RETRY_TIMEOUT_SECONDS, - interval_seconds=DEFAULT_RETRY_INTERVAL_SECONDS, - ): - url = f"{self.BASE_URL}/api/public/observations/{observation_id}" - return self._get_json( - url, - retry=retry, - is_result_ready=is_result_ready, - timeout_seconds=timeout_seconds, - interval_seconds=interval_seconds, - ) - - def get_scores( - self, - page=None, - limit=None, - user_id=None, - name=None, - *, - retry=True, - is_result_ready=None, - timeout_seconds=DEFAULT_RETRY_TIMEOUT_SECONDS, - interval_seconds=DEFAULT_RETRY_INTERVAL_SECONDS, - ): - params = {"page": page, "limit": limit, "userId": user_id, "name": name} - url = f"{self.BASE_URL}/api/public/scores" - return self._get_json( - url, - params=params, - retry=retry, - is_result_ready=is_result_ready, - timeout_seconds=timeout_seconds, - interval_seconds=interval_seconds, - ) - - def get_traces( - self, - page=None, - limit=None, - user_id=None, - name=None, - *, - retry=True, - is_result_ready=None, - timeout_seconds=DEFAULT_RETRY_TIMEOUT_SECONDS, - interval_seconds=DEFAULT_RETRY_INTERVAL_SECONDS, - ): - params = {"page": page, "limit": limit, "userId": user_id, "name": name} - url = f"{self.BASE_URL}/api/public/traces" - return self._get_json( - url, - params=params, - retry=retry, - is_result_ready=is_result_ready, - timeout_seconds=timeout_seconds, - interval_seconds=interval_seconds, - ) - - def get_trace( - self, - trace_id, - *, - retry=True, - is_result_ready=None, - timeout_seconds=DEFAULT_RETRY_TIMEOUT_SECONDS, - interval_seconds=DEFAULT_RETRY_INTERVAL_SECONDS, - ): - url = f"{self.BASE_URL}/api/public/traces/{trace_id}" - return self._get_json( - url, - retry=retry, - is_result_ready=is_result_ready, - timeout_seconds=timeout_seconds, - interval_seconds=interval_seconds, - ) diff --git a/tests/support/utils.py b/tests/support/utils.py index a29274d3a..191eeae05 100644 --- a/tests/support/utils.py +++ b/tests/support/utils.py @@ -1,19 +1,42 @@ import base64 +import json import os -from typing import Any, Callable, TypeVar +from dataclasses import dataclass +from typing import Any, Callable, Sequence, TypeVar from uuid import uuid4 -from langfuse.api import LangfuseAPI +from langfuse.api import LangfuseAPI, ObservationV2, ScoreV3 from tests.support.retry import ( DEFAULT_RETRY_INTERVAL_SECONDS, DEFAULT_RETRY_TIMEOUT_SECONDS, retry_until_ready, ) -READ_METHOD_NAMES = {"get", "get_by_id", "get_many", "get_run", "list"} -PAGINATION_ARGUMENTS = {"limit", "page"} +READ_METHOD_NAMES = {"get", "get_by_id", "get_many", "get_many_v3", "get_run", "list"} +PAGINATION_ARGUMENTS = {"limit", "page", "cursor", "fields", "expand_metadata"} T = TypeVar("T") +ALL_OBSERVATION_FIELDS = ( + "core,basic,time,io,metadata,model,usage,prompt,metrics,trace_context" +) +SCORE_FIELDS = "details,subject" + +# The v2 observations API returns "" instead of null for unset string fields. +_EMPTY_AS_NONE_FIELDS = ( + "name", + "status_message", + "version", + "user_id", + "session_id", + "model", + "internal_model_id", + "prompt_id", + "prompt_name", + "trace_name", + "release", +) +_SDK_METADATA_KEY_PREFIXES = ("scope.", "resourceAttributes.") + def _has_filters(kwargs: dict[str, Any]) -> bool: return any( @@ -51,7 +74,9 @@ def _call(*args: Any, **kwargs: Any) -> Any: def _result_ready(method_name: str, kwargs: dict[str, Any]): - if method_name not in {"get_many", "list"} or not _has_filters(kwargs): + if method_name not in {"get_many", "get_many_v3", "list"} or not _has_filters( + kwargs + ): return None def _has_data(result: Any) -> bool: @@ -89,17 +114,315 @@ def wait_for_result( ) -def wait_for_trace( +def _parse_json_string(value: Any) -> Any: + # Only objects and arrays: the SDK sends string IO unquoted, so text such as + # "2" or "true" must stay a string. + if not isinstance(value, str) or not value.lstrip().startswith(("{", "[")): + return value + + try: + return json.loads(value) + except ValueError: + return value + + +def normalize_observation(observation: ObservationV2) -> ObservationV2: + """Map a v2 observation to the shape the SDK sent. + + The v2 API returns input/output as raw strings and unset string fields as + "". This parses JSON object/array input/output and maps "" to None. + """ + update: dict[str, Any] = { + "input": _parse_json_string(observation.input), + "output": _parse_json_string(observation.output), + } + for field in _EMPTY_AS_NONE_FIELDS: + if getattr(observation, field, None) == "": + update[field] = None + + return observation.model_copy(update=update) + + +def user_metadata(observation: ObservationV2) -> dict[str, Any]: + """Observation metadata without the scope/resource keys the server adds.""" + metadata = observation.metadata or {} + assert isinstance(metadata, dict), metadata + + return { + key: value + for key, value in metadata.items() + if not key.startswith(_SDK_METADATA_KEY_PREFIXES) + } + + +def get_observations( + *, + fields: str = ALL_OBSERVATION_FIELDS, + api: Any = None, + **filters: Any, +) -> list[ObservationV2]: + """Fetch all observations matching `filters` (all pages), oldest first.""" + api = api or get_api(retry=False) + observations: list[ObservationV2] = [] + cursor = None + + while True: + response = api.observations.get_many( + fields=fields, limit=1000, cursor=cursor, **filters + ) + observations.extend(response.data) + cursor = response.meta.cursor + if not cursor or not response.data: + break + + # The events table can briefly return several rows for one span until + # ClickHouse merges them; keep the most recently updated row per id. + latest_by_id: dict[str, ObservationV2] = {} + for observation in observations: + current = latest_by_id.get(observation.id) + if current is None or (observation.updated_at or observation.start_time) >= ( + current.updated_at or current.start_time + ): + latest_by_id[observation.id] = observation + + return sorted( + (normalize_observation(observation) for observation in latest_by_id.values()), + key=lambda observation: observation.start_time, + ) + + +def wait_for_observations( + trace_id: str | None = None, + *, + min_count: int = 1, + is_result_ready: Callable[[list[ObservationV2]], bool] | None = None, + fields: str = ALL_OBSERVATION_FIELDS, + timeout_seconds: float = DEFAULT_RETRY_TIMEOUT_SECONDS, + interval_seconds: float = DEFAULT_RETRY_INTERVAL_SECONDS, + **filters: Any, +) -> list[ObservationV2]: + """Poll the v2 observations API until at least `min_count` observations + match (and `is_result_ready` holds), then return them oldest first.""" + if trace_id is not None: + filters["trace_id"] = trace_id + assert filters, "wait_for_observations needs at least one filter" + + def _ready(observations: list[ObservationV2]) -> bool: + return len(observations) >= min_count and ( + is_result_ready is None or is_result_ready(observations) + ) + + return wait_for_result( + lambda: get_observations(fields=fields, **filters), + is_result_ready=_ready, + timeout_seconds=timeout_seconds, + interval_seconds=interval_seconds, + ) + + +def get_root_observation(observations: Sequence[ObservationV2]) -> ObservationV2: + """Return the single root observation, which carries the trace-level + attributes (trace name, user, session, tags, public) and trace IO.""" + roots = [ + observation for observation in observations if observation.is_root_observation + ] + assert len(roots) == 1, ( + f"expected exactly one root observation, got {[r.name for r in roots]}" + ) + + return roots[0] + + +def wait_for_root_observation( + trace_id: str, + *, + min_count: int = 1, + is_result_ready: Callable[[ObservationV2], bool] | None = None, + timeout_seconds: float = DEFAULT_RETRY_TIMEOUT_SECONDS, + interval_seconds: float = DEFAULT_RETRY_INTERVAL_SECONDS, +) -> ObservationV2: + def _ready(observations: list[ObservationV2]) -> bool: + roots = [o for o in observations if o.is_root_observation] + return len(roots) == 1 and ( + is_result_ready is None or is_result_ready(roots[0]) + ) + + return get_root_observation( + wait_for_observations( + trace_id, + min_count=min_count, + is_result_ready=_ready, + timeout_seconds=timeout_seconds, + interval_seconds=interval_seconds, + ) + ) + + +@dataclass(frozen=True) +class TraceSnapshot: + """A trace as v4 exposes it: its observations plus (optionally) scores. + + Mirrors how the platform aggregates events into a trace: input, output + and metadata come from the root observation; name, user, session, + version, release and environment are the latest non-empty value across + all observations; tags are the union; public is true if any observation + is public. + """ + + id: str + observations: list[ObservationV2] + scores: list[ScoreV3] | None = None + + @property + def root(self) -> ObservationV2: + return get_root_observation(self.observations) + + def _latest(self, field: str) -> Any: + values = [ + getattr(observation, field) + for observation in self.observations + if getattr(observation, field) + ] + return values[-1] if values else None + + @property + def name(self) -> str | None: + # The API falls back to the root's own name when no trace name was + # set, so a root whose trace_name equals its name carries no signal. + explicit_names = [ + observation.trace_name + for observation in self.observations + if observation.trace_name + and not ( + observation.is_root_observation + and observation.trace_name == observation.name + ) + ] + return explicit_names[-1] if explicit_names else self.root.trace_name + + @property + def user_id(self) -> str | None: + return self._latest("user_id") + + @property + def session_id(self) -> str | None: + return self._latest("session_id") + + @property + def tags(self) -> list[str]: + return sorted( + {tag for observation in self.observations for tag in observation.tags or []} + ) + + @property + def public(self) -> bool: + return any(observation.public for observation in self.observations) + + @property + def version(self) -> str | None: + return self._latest("version") + + @property + def release(self) -> str | None: + return self._latest("release") + + @property + def environment(self) -> str | None: + return self._latest("environment") + + @property + def input(self) -> Any: + return self.root.input + + @property + def output(self) -> Any: + return self.root.output + + @property + def metadata(self) -> dict[str, Any]: + return user_metadata(self.root) + + +def wait_for_trace_snapshot( trace_id: str, *, - is_result_ready: Callable[[Any], bool] | None = None, + min_observations: int = 1, + min_scores: int | None = None, + is_result_ready: Callable[[TraceSnapshot], bool] | None = None, timeout_seconds: float = DEFAULT_RETRY_TIMEOUT_SECONDS, interval_seconds: float = DEFAULT_RETRY_INTERVAL_SECONDS, -): - api = get_api(retry=False) +) -> TraceSnapshot: + """Poll until the trace has `min_observations` observations (and + `min_scores` scores, which are only fetched when given).""" + + def _fetch() -> TraceSnapshot: + return TraceSnapshot( + id=trace_id, + observations=get_observations(trace_id=trace_id), + scores=None if min_scores is None else get_scores(trace_id=trace_id), + ) + + def _ready(snapshot: TraceSnapshot) -> bool: + if len(snapshot.observations) < min_observations: + return False + if min_scores is not None and len(snapshot.scores or []) < min_scores: + return False + if is_result_ready is None: + return True + try: + return is_result_ready(snapshot) + except AssertionError: + # Root-derived attributes are unavailable until the root arrives. + return False + return wait_for_result( - lambda: api.trace.get(trace_id), - is_result_ready=is_result_ready, + _fetch, + is_result_ready=_ready, + timeout_seconds=timeout_seconds, + interval_seconds=interval_seconds, + ) + + +def get_scores( + *, fields: str = SCORE_FIELDS, api: Any = None, **filters: Any +) -> list[ScoreV3]: + api = api or get_api(retry=False) + scores: list[ScoreV3] = [] + cursor = None + + while True: + response = api.scores_v3.get_many_v3( + fields=fields, limit=100, cursor=cursor, **filters + ) + scores.extend(response.data) + cursor = response.meta.cursor + if not cursor or not response.data: + break + + return scores + + +def wait_for_scores( + *, + min_count: int = 1, + is_result_ready: Callable[[list[ScoreV3]], bool] | None = None, + fields: str = SCORE_FIELDS, + timeout_seconds: float = DEFAULT_RETRY_TIMEOUT_SECONDS, + interval_seconds: float = DEFAULT_RETRY_INTERVAL_SECONDS, + **filters: Any, +) -> list[ScoreV3]: + """Poll the v3 scores API (filters: trace_id, session_id, observation_id, + name, ...) until at least `min_count` scores match.""" + assert filters, "wait_for_scores needs at least one filter" + + def _ready(scores: list[ScoreV3]) -> bool: + return len(scores) >= min_count and ( + is_result_ready is None or is_result_ready(scores) + ) + + return wait_for_result( + lambda: get_scores(fields=fields, **filters), + is_result_ready=_ready, timeout_seconds=timeout_seconds, interval_seconds=interval_seconds, ) diff --git a/tests/unit/test_batch_evaluation.py b/tests/unit/test_batch_evaluation.py new file mode 100644 index 000000000..1ac1751fa --- /dev/null +++ b/tests/unit/test_batch_evaluation.py @@ -0,0 +1,380 @@ +"""Unit tests for batch evaluation on the v2 observations API.""" + +import inspect +import json +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from typing import Any, Dict, List, Optional +from unittest.mock import MagicMock + +import pytest + +from langfuse import Langfuse +from langfuse.api import ObservationsV2Response, ObservationV2 +from langfuse.batch_evaluation import ( + DEFAULT_BATCH_EVALUATION_FIELDS, + BatchEvaluationResumeToken, + BatchEvaluationRunner, + EvaluatorInputs, +) +from langfuse.experiment import Evaluation + +BASE_TIME = datetime(2026, 1, 1, tzinfo=timezone.utc) + + +def make_observation( + index: int, + *, + trace_id: Optional[str] = "default", + is_root: bool = False, + start_time: Optional[datetime] = None, + observation_id: Optional[str] = None, + parent_observation_id: Optional[str] = "default", +) -> ObservationV2: + return ObservationV2( + id=observation_id or f"obs-{index}", + trace_id=f"trace-{index}" if trace_id == "default" else trace_id, + start_time=start_time or BASE_TIME + timedelta(seconds=index), + end_time=None, + project_id="project", + parent_observation_id=( + (None if is_root else f"parent-{index}") + if parent_observation_id == "default" + else parent_observation_id + ), + type="SPAN", + is_root_observation=is_root, + input=json.dumps({"question": index}), + output=f"answer {index}", + ) + + +class FakeObservationsApi: + """Serves observations newest first with an index cursor, like the v2 API.""" + + def __init__(self, observations: List[ObservationV2]): + self.observations = sorted( + observations, key=lambda o: (o.start_time, o.id), reverse=True + ) + self.calls: List[Dict[str, Any]] = [] + self.fail_on_call: Optional[int] = None + + def get_many(self, **kwargs: Any) -> ObservationsV2Response: + self.calls.append(kwargs) + if self.fail_on_call is not None and len(self.calls) == self.fail_on_call: + raise ConnectionError("fetch failed") + + matching = [o for o in self.observations if self._matches(o, kwargs["filter"])] + start = int(kwargs["cursor"]) if kwargs["cursor"] else 0 + page = matching[start : start + kwargs["limit"]] + has_more = start + kwargs["limit"] < len(matching) + meta = {"cursor": str(start + kwargs["limit"])} if has_more else {} + return ObservationsV2Response(data=page, meta=meta) + + @staticmethod + def _matches(observation: ObservationV2, filter_json: Optional[str]) -> bool: + for condition in json.loads(filter_json) if filter_json else []: + if condition["column"] == "isRootObservation": + if observation.is_root_observation is not condition["value"]: + return False + elif condition["column"] == "startTime": + assert condition["operator"] == "<=" + if observation.start_time > datetime.fromisoformat(condition["value"]): + return False + return True + + +def make_runner(observations: List[ObservationV2]): + api = FakeObservationsApi(observations) + client = SimpleNamespace( + api=SimpleNamespace(observations=api), + create_score=MagicMock(), + flush=MagicMock(), + ) + return BatchEvaluationRunner(client), api, client # type: ignore[arg-type] + + +def mapper(*, item: ObservationV2) -> EvaluatorInputs: + return EvaluatorInputs(input=item.input, output=item.output, metadata={}) + + +def length_evaluator(*, input, output, **kwargs): + return Evaluation(name="length", value=len(output)) + + +async def run(runner: BatchEvaluationRunner, **kwargs: Any): + kwargs.setdefault("mapper", mapper) + kwargs.setdefault("evaluators", [length_evaluator]) + return await runner.run_async(**kwargs) + + +@pytest.mark.asyncio +async def test_paginates_with_cursor_and_scores_observations(): + runner, api, client = make_runner([make_observation(i) for i in range(5)]) + + result = await run(runner, fetch_batch_size=2) + + assert [call["cursor"] for call in api.calls] == [None, "2", "4"] + assert all("page" not in call for call in api.calls) + assert all( + call["fields"] == DEFAULT_BATCH_EVALUATION_FIELDS and call["filter"] is None + for call in api.calls + ) + assert result.completed is True + assert result.has_more_items is False + assert result.resume_token is None + assert result.total_items_fetched == 5 + assert result.total_items_processed == 5 + assert result.total_scores_created == 5 + assert set(result.item_evaluations) == {f"obs-{i}" for i in range(5)} + client.create_score.assert_any_call( + observation_id="obs-3", + trace_id="trace-3", + name="length", + value=len("answer 3"), + comment=None, + metadata={}, + data_type=None, + config_id=None, + ) + client.flush.assert_called_once() + + +@pytest.mark.asyncio +async def test_mapper_receives_raw_string_io_and_requested_fields(): + runner, api, _ = make_runner([make_observation(1)]) + seen: List[ObservationV2] = [] + + def recording_mapper(*, item): + seen.append(item) + return mapper(item=item) + + await run(runner, mapper=recording_mapper, fields="core,io") + + assert api.calls[0]["fields"] == "core,io" + assert isinstance(seen[0], ObservationV2) + assert seen[0].input == '{"question": 1}' + + +def test_client_and_runner_share_default_fields(): + client_default = ( + inspect.signature(Langfuse.run_batched_evaluation).parameters["fields"].default + ) + runner_default = ( + inspect.signature(BatchEvaluationRunner.run_async).parameters["fields"].default + ) + + assert client_default == runner_default == DEFAULT_BATCH_EVALUATION_FIELDS + assert "io" in DEFAULT_BATCH_EVALUATION_FIELDS.split(",") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("kwargs", "message"), + [ + ({"filter": "not json"}, "JSON array"), + ({"filter": '{"tags": ["a"]}'}, "JSON array"), + ], +) +async def test_invalid_arguments_raise_before_fetching(kwargs, message): + runner, api, _ = make_runner([make_observation(1)]) + + with pytest.raises(ValueError, match=message): + await run(runner, **kwargs) + + assert api.calls == [] + + +@pytest.mark.asyncio +async def test_max_items_aligns_pages_and_resume_continues_without_gaps(): + # Shared start times exercise ties that a timestamp-only resume would skip. + observations = [ + make_observation(i, start_time=BASE_TIME + timedelta(seconds=i // 3)) + for i in range(10) + ] + runner, api, _ = make_runner(observations) + + first = await run(runner, max_items=5, fetch_batch_size=3) + + assert [call["limit"] for call in api.calls] == [3, 2] + assert first.completed is True + assert first.has_more_items is True + assert first.total_items_fetched == 5 + token = first.resume_token + assert token is not None + assert token.cursor == "5" + assert token.items_processed == 5 + + second = await run(runner, resume_from=token, fetch_batch_size=3) + + assert second.resume_token is None + assert second.has_more_items is False + assert set(first.item_evaluations).isdisjoint(second.item_evaluations) + assert set(first.item_evaluations) | set(second.item_evaluations) == { + f"obs-{i}" for i in range(10) + } + + +@pytest.mark.asyncio +async def test_max_items_reached_on_last_page_has_no_resume_token(): + runner, _, _ = make_runner([make_observation(i) for i in range(4)]) + + result = await run(runner, max_items=4, fetch_batch_size=2) + + assert result.has_more_items is False + assert result.resume_token is None + + +@pytest.mark.asyncio +async def test_fetch_failure_returns_resume_token_for_failed_page(): + runner, api, _ = make_runner([make_observation(i) for i in range(6)]) + api.fail_on_call = 2 + + failed = await run(runner, fetch_batch_size=2, max_retries=0) + + assert failed.completed is False + assert failed.total_items_processed == 2 + token = failed.resume_token + assert token is not None + assert token.cursor == "2" + assert token.last_processed_id == "obs-4" + + api.fail_on_call = None + resumed = await run(runner, fetch_batch_size=2, resume_from=token) + + assert resumed.completed is True + assert set(resumed.item_evaluations) == {"obs-3", "obs-2", "obs-1", "obs-0"} + assert api.calls[-1]["request_options"] == {"max_retries": 3} + + +@pytest.mark.asyncio +async def test_resume_without_cursor_falls_back_to_start_time_bound(): + tied_time = BASE_TIME + timedelta(seconds=3) + runner, api, _ = make_runner( + [make_observation(i) for i in range(5)] + + [make_observation(9, observation_id="obs-tied", start_time=tied_time)] + ) + token = BatchEvaluationResumeToken( + filter=None, + last_processed_timestamp=(BASE_TIME + timedelta(seconds=3)).isoformat(), + last_processed_id="obs-3", + items_processed=2, + ) + + result = await run(runner, resume_from=token) + + assert json.loads(api.calls[0]["filter"]) == [ + { + "type": "datetime", + "column": "startTime", + "operator": "<=", + "value": token.last_processed_timestamp, + } + ] + assert api.calls[0]["cursor"] is None + assert set(result.item_evaluations) == {"obs-tied", "obs-2", "obs-1", "obs-0"} + + +@pytest.mark.asyncio +async def test_resume_reuses_the_token_filter_and_rejects_a_different_one(): + runner, api, _ = make_runner([make_observation(i) for i in range(4)]) + token_filter = json.dumps( + [{"type": "string", "column": "name", "operator": "=", "value": "x"}] + ) + token = BatchEvaluationResumeToken( + filter=token_filter, + cursor="2", + last_processed_timestamp=BASE_TIME.isoformat(), + last_processed_id="obs-2", + items_processed=2, + ) + api._matches = lambda observation, filter_json: True # type: ignore[method-assign] + + await run(runner, resume_from=token) + assert json.loads(api.calls[0]["filter"]) == json.loads(token_filter) + + with pytest.raises(ValueError, match="different filter"): + await run(runner, resume_from=token, filter="[]") + + +@pytest.mark.asyncio +async def test_observation_scores_can_be_duplicated_onto_trace(): + runner, _, client = make_runner([make_observation(1)]) + + result = await run(runner, _add_observation_scores_to_trace=True) + + assert result.total_scores_created == 2 + targets = [ + (call.kwargs.get("observation_id"), call.kwargs["trace_id"]) + for call in client.create_score.call_args_list + ] + assert targets == [("obs-1", "trace-1"), (None, "trace-1")] + + +@pytest.mark.asyncio +async def test_item_failures_and_evaluator_failures_are_tracked(): + observations = [make_observation(1), make_observation(2, trace_id=None)] + runner, _, client = make_runner(observations) + + def failing_evaluator(**kwargs): + raise RuntimeError("boom") + + def composite(*, evaluations, **kwargs): + return Evaluation(name="composite", value=len(evaluations)) + + result = await run( + runner, + evaluators=[length_evaluator, failing_evaluator], + composite_evaluator=composite, + ) + + assert result.total_items_processed == 1 + assert result.failed_item_ids == ["obs-2"] + assert result.error_summary == {"ValueError": 1} + assert result.total_evaluations_failed == 1 + assert result.total_composite_scores_created == 1 + assert [e.name for e in result.item_evaluations["obs-1"]] == [ + "length", + "composite", + ] + stats = {s.name: s for s in result.evaluator_stats} + assert stats["failing_evaluator"].failed_runs == 1 + assert stats["length_evaluator"].successful_runs == 1 + assert client.create_score.call_count == 2 + + +@pytest.mark.asyncio +async def test_async_mapper_and_evaluator_are_awaited(): + runner, _, _ = make_runner([make_observation(1)]) + + async def async_mapper(*, item): + return mapper(item=item) + + async def async_evaluator(*, output, **kwargs): + return [Evaluation(name="a", value=1), Evaluation(name="b", value=2)] + + result = await run(runner, mapper=async_mapper, evaluators=[async_evaluator]) + + assert result.total_scores_created == 2 + + +@pytest.mark.asyncio +async def test_duplicate_rows_for_one_observation_are_evaluated_once(): + # The events table can briefly return several rows for the same observation. + stale = make_observation(1) + fresh = make_observation(1).model_copy( + update={"output": "fresh answer", "updated_at": BASE_TIME + timedelta(hours=1)} + ) + runner, _, client = make_runner([make_observation(0), stale, fresh]) + seen: List[ObservationV2] = [] + + def recording_mapper(*, item): + seen.append(item) + return mapper(item=item) + + result = await run(runner, mapper=recording_mapper, fetch_batch_size=2) + + assert sorted(item.id for item in seen) == ["obs-0", "obs-1"] + assert [item.output for item in seen if item.id == "obs-1"] == ["fresh answer"] + assert result.total_items_processed == 2 + assert client.create_score.call_count == 2 diff --git a/tests/unit/test_e2e_support.py b/tests/unit/test_e2e_support.py index 8320bd2fe..98d25127c 100644 --- a/tests/unit/test_e2e_support.py +++ b/tests/unit/test_e2e_support.py @@ -1,144 +1,260 @@ +from datetime import datetime, timedelta, timezone from types import SimpleNamespace +from langfuse.api import ObservationV2 from langfuse.api.commons.errors.not_found_error import NotFoundError -from tests.support.api_wrapper import LangfuseAPI as SupportLangfuseAPI from tests.support.retry import retry_until_ready -from tests.support.utils import get_api, wait_for_trace +from tests.support.utils import ( + TraceSnapshot, + get_api, + get_observations, + normalize_observation, + user_metadata, + wait_for_observations, + wait_for_scores, + wait_for_trace_snapshot, +) +START = datetime(2024, 1, 1, tzinfo=timezone.utc) -def test_get_api_retries_not_found(monkeypatch): - monkeypatch.setattr("tests.support.retry.sleep", lambda _: None) - attempts = {"count": 0} +def _observation(index: int = 0, **fields) -> ObservationV2: + defaults = { + "id": f"obs-{index}", + "trace_id": "trace-123", + "start_time": START + timedelta(seconds=index), + "project_id": "project", + "type": "SPAN", + "is_root_observation": False, + } + return ObservationV2(**{**defaults, **fields}) - class FakeTraceService: - def get(self, trace_id): - attempts["count"] += 1 - if attempts["count"] < 3: - raise NotFoundError( - body={ - "error": "LangfuseNotFoundError", - "message": f"Trace {trace_id} not found within authorized project", - } - ) +def _page(data, cursor=None): + return SimpleNamespace(data=data, meta=SimpleNamespace(cursor=cursor)) - return {"id": trace_id} - class FakeClient: - trace = FakeTraceService() +def _install_client(monkeypatch, **services): + monkeypatch.setattr("tests.support.retry.sleep", lambda _: None) + client = SimpleNamespace(**services) + monkeypatch.setattr("tests.support.utils.LangfuseAPI", lambda **_: client) - monkeypatch.setattr("tests.support.utils.LangfuseAPI", lambda **_: FakeClient()) - trace = get_api().trace.get("trace-123") +def test_get_api_retries_not_found(monkeypatch): + attempts = {"count": 0} - assert trace == {"id": "trace-123"} - assert attempts["count"] == 3 + def get_many(**kwargs): + attempts["count"] += 1 + if attempts["count"] < 3: + raise NotFoundError( + body={ + "error": "LangfuseNotFoundError", + "message": "Observations not found within authorized project", + } + ) -def test_get_api_retries_filtered_lists(monkeypatch): - monkeypatch.setattr("tests.support.retry.sleep", lambda _: None) + return _page([kwargs["trace_id"]]) - attempts = {"count": 0} + _install_client(monkeypatch, observations=SimpleNamespace(get_many=get_many)) - class FakeTraceService: - def list(self, **kwargs): - attempts["count"] += 1 + response = get_api().observations.get_many(trace_id="trace-123") - if attempts["count"] < 3: - return SimpleNamespace(data=[]) + assert response.data == ["trace-123"] + assert attempts["count"] == 3 - return SimpleNamespace(data=[kwargs["name"]]) - class FakeClient: - trace = FakeTraceService() +def test_get_api_retries_filtered_lists(monkeypatch): + attempts = {"count": 0} - monkeypatch.setattr("tests.support.utils.LangfuseAPI", lambda **_: FakeClient()) + def get_many(**kwargs): + attempts["count"] += 1 + return _page([] if attempts["count"] < 3 else [kwargs["name"]]) - response = get_api().trace.list(name="ready-trace") + _install_client(monkeypatch, observations=SimpleNamespace(get_many=get_many)) - assert response.data == ["ready-trace"] + response = get_api().observations.get_many(name="ready-observation") + + assert response.data == ["ready-observation"] assert attempts["count"] == 3 def test_get_api_retry_can_be_disabled(monkeypatch): attempts = {"count": 0} - class FakeTraceService: - def list(self, **kwargs): - attempts["count"] += 1 - return SimpleNamespace(data=[]) - - class FakeClient: - trace = FakeTraceService() + def get_many(**kwargs): + attempts["count"] += 1 + return _page([]) - monkeypatch.setattr("tests.support.utils.LangfuseAPI", lambda **_: FakeClient()) + _install_client(monkeypatch, observations=SimpleNamespace(get_many=get_many)) - response = get_api(retry=False).trace.list(name="missing-trace") + response = get_api(retry=False).observations.get_many(name="missing") assert response.data == [] assert attempts["count"] == 1 -def test_raw_api_wrapper_retries_not_found_payload(monkeypatch): - monkeypatch.setattr("tests.support.retry.sleep", lambda _: None) +def test_normalize_observation_parses_io_and_maps_empty_strings_to_none(): + observation = normalize_observation( + _observation( + input='{"question": "hi"}', + output="plain text", + name="", + session_id="", + user_id="user-1", + ) + ) - attempts = {"count": 0} + assert observation.input == {"question": "hi"} + assert observation.output == "plain text" + assert observation.name is None + assert observation.session_id is None + assert observation.user_id == "user-1" - class FakeResponse: - def __init__(self, status_code, payload): - self.status_code = status_code - self._payload = payload - self.headers = {} - def json(self): - return self._payload +def test_user_metadata_drops_server_added_keys(): + observation = _observation( + metadata={ + "key": "value", + "scope.name": "langfuse-sdk", + "resourceAttributes.service.name": "test", + } + ) - def fake_get(*args, **kwargs): - attempts["count"] += 1 + assert user_metadata(observation) == {"key": "value"} - if attempts["count"] < 3: - return FakeResponse( - 404, - { - "error": "LangfuseNotFoundError", - "message": "Trace trace-123 not found within authorized project", - }, - ) - return FakeResponse(200, {"id": "trace-123", "observations": []}) +def test_get_observations_follows_cursor_and_sorts_by_start_time(monkeypatch): + calls = [] + pages = { + None: _page([_observation(2), _observation(0)], cursor="next"), + "next": _page([_observation(1)]), + } - monkeypatch.setattr("tests.support.api_wrapper.httpx.get", fake_get) + def get_many(**kwargs): + calls.append(kwargs) + return pages[kwargs["cursor"]] - api = SupportLangfuseAPI(username="user", password="pass", base_url="http://test") - trace = api.get_trace("trace-123") + _install_client(monkeypatch, observations=SimpleNamespace(get_many=get_many)) - assert trace["id"] == "trace-123" - assert attempts["count"] == 3 + observations = get_observations(trace_id="trace-123") + assert [o.id for o in observations] == ["obs-0", "obs-1", "obs-2"] + assert [call["cursor"] for call in calls] == [None, "next"] + assert all(call["trace_id"] == "trace-123" for call in calls) -def test_wait_for_trace_retries_until_predicate_matches(monkeypatch): - monkeypatch.setattr("tests.support.retry.sleep", lambda _: None) +def test_wait_for_observations_polls_until_min_count(monkeypatch): attempts = {"count": 0} - class FakeTraceService: - def get(self, trace_id): - attempts["count"] += 1 - return {"id": trace_id, "observations": [1] * attempts["count"]} + def get_many(**kwargs): + attempts["count"] += 1 + return _page([_observation(i) for i in range(attempts["count"])]) + + _install_client(monkeypatch, observations=SimpleNamespace(get_many=get_many)) - class FakeClient: - trace = FakeTraceService() + observations = wait_for_observations("trace-123", min_count=3) + + assert len(observations) == 3 + assert attempts["count"] == 3 - monkeypatch.setattr("tests.support.utils.LangfuseAPI", lambda **_: FakeClient()) - trace = wait_for_trace( - "trace-123", is_result_ready=lambda trace: len(trace["observations"]) == 3 +def test_wait_for_trace_snapshot_waits_for_root_and_scores(monkeypatch): + attempts = {"observations": 0, "scores": 0} + + def get_many(**kwargs): + attempts["observations"] += 1 + observations = [_observation(1, name="child")] + if attempts["observations"] >= 2: + observations.append(_observation(0, name="root", is_root_observation=True)) + return _page(observations) + + def get_many_v3(**kwargs): + attempts["scores"] += 1 + return _page(["score"] if attempts["scores"] >= 3 else []) + + _install_client( + monkeypatch, + observations=SimpleNamespace(get_many=get_many), + scores_v3=SimpleNamespace(get_many_v3=get_many_v3), ) - assert trace["id"] == "trace-123" - assert len(trace["observations"]) == 3 - assert attempts["count"] == 3 + snapshot = wait_for_trace_snapshot( + "trace-123", + min_scores=1, + is_result_ready=lambda trace: trace.root.name == "root", + ) + + assert snapshot.root.name == "root" + assert snapshot.scores == ["score"] + assert attempts["scores"] == 3 + + +def test_trace_snapshot_aggregates_trace_attributes_like_the_platform(): + snapshot = TraceSnapshot( + id="trace-123", + observations=[ + _observation( + 0, + name="root", + trace_name="root", + is_root_observation=True, + input={"q": 1}, + metadata={"root_key": "root", "scope.name": "sdk"}, + tags=["b"], + ), + _observation( + 1, + name="child", + trace_name="explicit-name", + session_id="session-1", + user_id="user-1", + tags=["a"], + public=True, + ), + _observation(2, name="grandchild", session_id="session-2"), + ], + ) + + assert snapshot.name == "explicit-name" + assert snapshot.session_id == "session-2" + assert snapshot.user_id == "user-1" + assert snapshot.tags == ["a", "b"] + assert snapshot.public is True + assert snapshot.input == {"q": 1} + assert snapshot.metadata == {"root_key": "root"} + + +def test_trace_snapshot_name_falls_back_to_root_trace_name(): + snapshot = TraceSnapshot( + id="trace-123", + observations=[ + _observation(0, name="root", trace_name="root", is_root_observation=True), + _observation(1, name="child", trace_name="root"), + ], + ) + + assert snapshot.name == "root" + assert snapshot.session_id is None + assert snapshot.public is False + + +def test_wait_for_scores_follows_cursor_and_polls(monkeypatch): + attempts = {"count": 0} + + def get_many_v3(**kwargs): + attempts["count"] += 1 + if attempts["count"] < 2: + return _page([]) + if kwargs["cursor"] is None: + return _page(["score-1"], cursor="next") + return _page(["score-2"]) + + _install_client(monkeypatch, scores_v3=SimpleNamespace(get_many_v3=get_many_v3)) + + scores = wait_for_scores(min_count=2, trace_id="trace-123") + + assert scores == ["score-1", "score-2"] def test_retry_until_ready_clears_stale_error_after_success(monkeypatch): @@ -171,3 +287,37 @@ def operation(): assert trace["id"] == "trace-123" assert trace["attempt"] == 3 + + +def test_normalize_observation_keeps_scalar_text_io_as_strings(): + observation = normalize_observation( + ObservationV2( + id="obs", + trace_id="trace", + start_time=datetime(2026, 1, 1, tzinfo=timezone.utc), + project_id="project", + type="SPAN", + input="2", + output="true", + ) + ) + + assert observation.input == "2" + assert observation.output == "true" + + +def test_get_observations_keeps_the_latest_row_per_observation_id(monkeypatch): + stale = _observation(0, name="stale", updated_at=START) + fresh = _observation(0, name="fresh", updated_at=START + timedelta(seconds=1)) + + def get_many(**kwargs): + return _page([fresh, stale, _observation(1)]) + + _install_client(monkeypatch, observations=SimpleNamespace(get_many=get_many)) + + observations = get_observations(trace_id="trace-123") + + assert [(o.id, o.name) for o in observations] == [ + ("obs-0", "fresh"), + ("obs-1", None), + ] diff --git a/tests/unit/test_mask_otel_spans.py b/tests/unit/test_mask_otel_spans.py index 5d1d52ee4..d9a23befb 100644 --- a/tests/unit/test_mask_otel_spans.py +++ b/tests/unit/test_mask_otel_spans.py @@ -293,7 +293,6 @@ def test_export_stage_media_processes_string_sequence_attributes(): @pytest.mark.parametrize( ("attribute_key", "expected_field"), [ - ("langfuse.trace.input", "input"), ("langfuse.observation.input", "input"), ("ai.prompt.messages", "input"), ("gcp.vertex.agent.tool_call_args", "input"), @@ -302,7 +301,6 @@ def test_export_stage_media_processes_string_sequence_attributes(): ("gen_ai.input.messages", "input"), ("gen_ai.prompt.0.content", "input"), ("llm.input_messages.0.message.content", "input"), - ("langfuse.trace.output", "output"), ("langfuse.observation.output", "output"), ("ai.response.toolCalls", "output"), ("gcp.vertex.agent.tool_response", "output"), diff --git a/tests/unit/test_otel.py b/tests/unit/test_otel.py index ccdec39fd..6c4b23d3e 100644 --- a/tests/unit/test_otel.py +++ b/tests/unit/test_otel.py @@ -498,8 +498,8 @@ def test_generation_name_update(self, langfuse_client, memory_exporter): def test_trace_update(self, langfuse_client, memory_exporter): """Test updating trace level attributes.""" - # Create a span and set trace attributes using propagate_attributes and set_trace_io - with langfuse_client.start_as_current_observation(name="trace-span") as span: + # Create a span and set trace attributes using propagate_attributes + with langfuse_client.start_as_current_observation(name="trace-span"): with propagate_attributes( trace_name="updated-trace-name", user_id="test-user", @@ -507,7 +507,7 @@ def test_trace_update(self, langfuse_client, memory_exporter): tags=["tag1", "tag2"], metadata={"trace-meta": "data"}, ): - span.set_trace_io(input={"trace-input": "value"}) + pass # Get the span data spans = self.get_spans_by_name(memory_exporter, "trace-span") @@ -526,12 +526,10 @@ def test_trace_update(self, langfuse_client, memory_exporter): else: tags = list(attributes[LangfuseOtelSpanAttributes.TRACE_TAGS]) - input_data = json.loads(attributes[LangfuseOtelSpanAttributes.TRACE_INPUT]) metadata = attributes[f"{LangfuseOtelSpanAttributes.TRACE_METADATA}.trace-meta"] # Check attribute values assert sorted(tags) == sorted(["tag1", "tag2"]) - assert input_data == {"trace-input": "value"} assert metadata == "data" def test_complex_scenario(self, langfuse_client, memory_exporter): diff --git a/tests/unit/test_resource_manager.py b/tests/unit/test_resource_manager.py index f66a1e052..f48e61555 100644 --- a/tests/unit/test_resource_manager.py +++ b/tests/unit/test_resource_manager.py @@ -1,5 +1,6 @@ """Test the LangfuseResourceManager and get_client() function.""" +import logging from queue import Queue from types import SimpleNamespace from typing import Sequence @@ -410,6 +411,36 @@ def test_at_fork_reinit_new_httpx_client_uses_configured_timeout_and_headers( client.shutdown() +def test_create_score_with_invalid_input_logs_error_without_enqueueing( + monkeypatch, caplog +): + monkeypatch.setenv("LANGFUSE_MEDIA_UPLOAD_ENABLED", "false") + + with LangfuseResourceManager._lock: + LangfuseResourceManager._instances.clear() + + client = Langfuse( + public_key="pk-invalid-score", + secret_key="sk-invalid-score", + span_exporter=NoOpSpanExporter(), + ) + rm = client._resources + assert rm is not None + enqueued = [] + monkeypatch.setattr(rm, "add_score_task", lambda event, **_: enqueued.append(event)) + + with caplog.at_level(logging.ERROR, logger="langfuse"): + client.create_score(name="invalid", value=object(), trace_id="a" * 32) + + assert enqueued == [] + assert rm._score_ingestion_queue.empty() + error_records = [r for r in caplog.records if r.levelno == logging.ERROR] + assert len(error_records) == 1 + assert "Error creating score" in error_records[0].getMessage() + + client.shutdown() + + def test_stop_and_join_consumer_threads_broadcasts_media_shutdown_after_pausing_all(): events = []