diff --git a/tests/integration/cost_calculation/conftest.py b/tests/integration/cost_calculation/conftest.py index 9c6ffe3f7d0..7a75320a70f 100644 --- a/tests/integration/cost_calculation/conftest.py +++ b/tests/integration/cost_calculation/conftest.py @@ -3,7 +3,7 @@ from __future__ import annotations import functools import json import os -from collections.abc import Mapping +from collections.abc import Callable, Mapping from dataclasses import dataclass from hashlib import sha256 from typing import Final @@ -134,8 +134,16 @@ def read_rows_now(key: str) -> tuple[CostRow, ...]: def poll_rows(key: str, count: int) -> tuple[CostRow, ...]: + return poll_rows_where(key, count, lambda _row: True) + + +def poll_rows_where( + key: str, + count: int, + predicate: Callable[[CostRow], bool], +) -> tuple[CostRow, ...]: result: Final = eventually( - lambda: read_rows_now(key), + lambda: tuple(row for row in read_rows_now(key) if predicate(row)), lambda rows: len(rows) >= count, seconds=60, ) diff --git a/tests/integration/cost_calculation/test_batch_realtime_cost.py b/tests/integration/cost_calculation/test_batch_realtime_cost.py index 3bf807059b4..1941e9ad153 100644 --- a/tests/integration/cost_calculation/test_batch_realtime_cost.py +++ b/tests/integration/cost_calculation/test_batch_realtime_cost.py @@ -12,7 +12,7 @@ import websockets from integration._support.client import JSON_OBJECT, Gateway, Scenario, object_value, string_value from integration._support.upstream import delete_scenario, register_scenario from integration.cost_calculation.assertions import assert_exact -from integration.cost_calculation.conftest import poll_rows, read_rows_now +from integration.cost_calculation.conftest import poll_rows, poll_rows_where, read_rows_now from integration.cost_calculation.cost_tracking_case import ( BATCH_CASES, REALTIME_CASES, @@ -67,6 +67,8 @@ def _batch_response(case: BatchCostCase) -> JsonResponse | RoutedResponse: "completed": case.completed_count, "failed": case.failed_count, } + has_output: Final = any(line.status_code == 200 for line in case.output_lines) + has_failed: Final = any(line.status_code != 200 for line in case.output_lines) batch: Final = { "id": "batch-$REQUEST_ID", "object": "batch", @@ -75,8 +77,8 @@ def _batch_response(case: BatchCostCase) -> JsonResponse | RoutedResponse: "input_file_id": "file-in-$REQUEST_ID", "completion_window": "24h", "status": "completed", - "output_file_id": "file-out-$REQUEST_ID" if case.output_lines else None, - "error_file_id": None if case.output_lines else "file-err-$REQUEST_ID", + "output_file_id": "file-out-$REQUEST_ID" if has_output else None, + "error_file_id": "file-err-$REQUEST_ID" if has_failed else None, "created_at": 1, "in_progress_at": 1, "completed_at": 1, @@ -84,39 +86,46 @@ def _batch_response(case: BatchCostCase) -> JsonResponse | RoutedResponse: "request_counts": counts, "metadata": None, } + routes: Final = { + "POST /files": JsonResponse( + content_type="application/json", + body={ + "id": "file-in-$REQUEST_ID", + "object": "file", + "purpose": "batch", + "bytes": 100, + "created_at": 1, + "filename": "in.jsonl", + "status": "processed", + }, + ), + "POST /batches": JsonResponse( + content_type="application/json", + body={ + **batch, + "status": "validating", + "output_file_id": None, + "error_file_id": None, + }, + ), + "GET /batches/batch-$REQUEST_ID": JsonResponse( + content_type="application/json", + body=batch, + ), + **( + { + "GET /files/file-out-$REQUEST_ID/content": TextResponse( + content_type="application/jsonl", + body="\n".join(lines) + ("\n" if lines else ""), + ) + } + if has_output + else {} + ), + } return RoutedResponse( content_type="application/x-routed", - routes={ - "POST /files": JsonResponse( - content_type="application/json", - body={ - "id": "file-in-$REQUEST_ID", - "object": "file", - "purpose": "batch", - "bytes": 100, - "created_at": 1, - "filename": "in.jsonl", - "status": "processed", - }, - ), - "POST /batches": JsonResponse( - content_type="application/json", - body={ - **batch, - "status": "validating", - "output_file_id": None, - "error_file_id": None, - }, - ), - "GET /batches/batch-$REQUEST_ID": JsonResponse( - content_type="application/json", - body=batch, - ), - "GET /files/file-out-$REQUEST_ID/content": TextResponse( - content_type="application/jsonl", - body="\n".join(lines) + ("\n" if lines else ""), - ), - }, + routes=routes, ) @@ -164,7 +173,6 @@ def test_batch_costs(gateway: Gateway, case: BatchCostCase) -> None: ) assert file_response.is_success, file_response.text file_body: Final = JSON_OBJECT.validate_json(file_response.content) - time.sleep(2) batch_response: Final = gateway.request( "POST", "/v1/batches", @@ -183,9 +191,9 @@ def test_batch_costs(gateway: Gateway, case: BatchCostCase) -> None: second_retrieval: Final = gateway.request("GET", f"/v1/batches/{batch_id}", key=key) assert first_retrieval.is_success, first_retrieval.text assert second_retrieval.is_success, second_retrieval.text - rows: Final = poll_rows(key, 1) - retrieval_rows: Final = tuple(row for row in rows if row.call_type == "aretrieve_batch") + retrieval_rows: Final = poll_rows_where(key, 1, lambda row: row.call_type == "aretrieve_batch") assert len(retrieval_rows) == 1 + rows: Final = read_rows_now(key) assert all(row.spend == 0.0 for row in rows if row.call_type != "aretrieve_batch") row: Final = retrieval_rows[0] assert row.status == "success"