Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 40 additions & 8 deletions truss/remote/baseten/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import logging
import pathlib
import textwrap
import time
from typing import IO, TYPE_CHECKING, Any, Dict, List, NamedTuple, Optional, Tuple, Type

import requests
Expand Down Expand Up @@ -36,6 +37,11 @@
MAX_ITERATIONS = 10_000
MIN_BATCH_SIZE = 100

# Wall-clock bound on the whole forward scan. A chatty job needs one request per
# MAX_BATCH_SIZE lines, so without a deadline the scan outlives any CI step
# timeout and the caller reports its own timeout instead of the job's logs.
DEFAULT_LOG_FETCH_TIMEOUT_SEC = 60.0

# LIMIT for the number of logs to fetch per request defined by the server
MAX_BATCH_SIZE = 1000

Expand Down Expand Up @@ -625,16 +631,13 @@ def _process_batch_logs(
Tuple of (should_continue, next_start_time, next_end_time)
"""

# If no logs returned, we're done
# An empty window ends the scan. A short-but-non-empty batch does not: the
# window is capped at 2h, so a job that outlives one window still has logs
# past its end.
if not batch_logs:
logging.info(f"No logs returned for job {job_id} at iteration {iteration}")
return False, None, None

# If we got fewer logs than the batch size, we've reached the end
if len(batch_logs) == 0:
logging.info(f"Reached end of logs for job {job_id} at iteration {iteration}")
return False, None, None

# Timestamp returned in nanoseconds for the last log in this batch converted
# to milliseconds to use as start for next iteration
last_log_timestamp = int(batch_logs[-1]["timestamp"]) // NANOSECONDS_PER_MILLISECOND
Expand Down Expand Up @@ -662,6 +665,7 @@ def __init__(
project_id: str,
job_id: str,
batch_size: int = MAX_BATCH_SIZE,
timeout_sec: float = DEFAULT_LOG_FETCH_TIMEOUT_SEC,
):
self.api = api
self.project_id = project_id
Expand All @@ -670,6 +674,8 @@ def __init__(
self.iteration = 0
self.current_start_time = None
self.current_end_time = None
self.truncated = False
self._deadline = time.monotonic() + timeout_sec
self._initialize_time_window()

def _initialize_time_window(self):
Expand All @@ -684,12 +690,22 @@ def __iter__(self):

def __next__(self) -> List[Any]:
if self.iteration >= MAX_ITERATIONS:
self.truncated = True
logging.warning(
f"Reached maximum iteration limit ({MAX_ITERATIONS}) while paginating "
f"training job logs for project_id={self.project_id}, job_id={self.job_id}."
)
raise StopIteration

if time.monotonic() >= self._deadline:
self.truncated = True
logging.warning(
f"Timed out after {self.iteration} batches while paginating training job "
f"logs for project_id={self.project_id}, job_id={self.job_id}. "
"Returning the logs fetched so far."
)
raise StopIteration

query_params = _build_log_query_params(
self.current_start_time, self.current_end_time, self.batch_size
)
Expand Down Expand Up @@ -741,23 +757,39 @@ def __next__(self) -> List[Any]:


def get_training_job_logs_with_pagination(
api: BasetenApi, project_id: str, job_id: str, batch_size: int = MAX_BATCH_SIZE
api: BasetenApi,
project_id: str,
job_id: str,
batch_size: int = MAX_BATCH_SIZE,
timeout_sec: float = DEFAULT_LOG_FETCH_TIMEOUT_SEC,
) -> List[Any]:
"""
This method implements forward time-based pagination by starting from the earliest
available log and working forward in time. It uses the timestamp of the newest log in
each batch as the start time for the next request.

The scan is bounded by ``timeout_sec``; on expiry it returns the logs fetched so far
and warns, rather than blocking the caller indefinitely.

Returns:
List of all logs in chronological order (oldest first)
"""
all_logs = []

logs_iterator = BatchedTrainingLogsFetcher(api, project_id, job_id, batch_size)
logs_iterator = BatchedTrainingLogsFetcher(
api, project_id, job_id, batch_size, timeout_sec=timeout_sec
)

for batch_logs in logs_iterator:
all_logs.extend(batch_logs)

logging.info(f"Completed pagination for job {job_id}. Total logs: {len(all_logs)}")

if logs_iterator.truncated:
console.print(
f"Warning: stopped fetching logs early; showing the first "
f"{len(all_logs)} lines only.",
style="yellow",
)

return all_logs
31 changes: 31 additions & 0 deletions truss/tests/remote/baseten/test_core.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import itertools
import json
from tempfile import NamedTemporaryFile
from unittest import mock
Expand Down Expand Up @@ -1100,3 +1101,33 @@ def test_create_bis_llm_service_creates_model_version():
api.create_bis_llm_model_version.assert_called_once_with(
model_id="bis-llm-model-id", body=body
)


def test_get_training_job_logs_with_pagination_stops_at_timeout(baseten_api):
"""A chatty job needs one request per batch of lines, so an unbounded scan
outlives the caller. The deadline must return the logs fetched so far."""
counter = itertools.count()
baseten_api._fetch_log_batch = mock.Mock(
side_effect=lambda *_: [
{
"timestamp": str(1640995200000000000 + 60_000_000_000 * next(counter)),
"message": "Log",
}
]
)
baseten_api.get_training_job = mock.Mock(
return_value={"training_job": {"created_at": "2022-01-01T00:00:00Z"}}
)

# Deadline is set at construction and still unbreached for the first batch;
# every later reading is past it. Batches never run out on their own.
readings = iter([0.0, 0.0])
with mock.patch(
"truss.remote.baseten.core.time.monotonic", lambda: next(readings, 999.0)
):
result = get_training_job_logs_with_pagination(
baseten_api, "project-123", "job-456", batch_size=1, timeout_sec=30.0
)

assert len(result) == 1
assert baseten_api._fetch_log_batch.call_count == 1
Loading