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
89 changes: 67 additions & 22 deletions sagemaker-core/src/sagemaker/core/helper/session_helper.py
Original file line number Diff line number Diff line change
Expand Up @@ -1274,9 +1274,15 @@ def wait_for_endpoint(self, endpoint, poll=DEFAULT_EP_POLL, live_logging=False):
cloudwatch_client = self.boto_session.client("logs")
paginator = cloudwatch_client.get_paginator("filter_log_events")
paginator_config = create_paginator_config()
not_found_budget = _EndpointNotFoundBudget()
desc = _wait_until(
lambda: _live_logging_deploy_done(
self.sagemaker_client, endpoint, paginator, paginator_config, EP_LOGGER_POLL
self.sagemaker_client,
endpoint,
paginator,
paginator_config,
EP_LOGGER_POLL,
not_found_budget=not_found_budget,
),
poll=EP_LOGGER_POLL,
)
Expand Down Expand Up @@ -2979,30 +2985,73 @@ def _deploy_done(sagemaker_client, endpoint_name):
return None if status in in_progress_statuses else desc


def _live_logging_deploy_done(sagemaker_client, endpoint_name, paginator, paginator_config, poll):
"""Placeholder docstring"""
# DescribeEndpoint can briefly report "Could not find endpoint" right after
# CreateEndpoint returns. An endpoint still missing after this many consecutive
# polls was deleted out of band, and waiting on it can never succeed.
_MAX_ENDPOINT_NOT_FOUND_POLLS = 10


class _EndpointNotFoundBudget(object):
"""Bounds how many consecutive polls a deploy wait tolerates a missing endpoint."""

def __init__(self, max_polls=_MAX_ENDPOINT_NOT_FOUND_POLLS):
"""Allow up to ``max_polls`` consecutive "endpoint not found" polls."""
self.max_polls = max_polls
self.misses = 0

def miss(self, error):
"""Record one "not found" poll, re-raising ``error`` once the budget is spent."""
self.misses += 1
if self.misses > self.max_polls:
raise error

def reset(self):
"""Clear the count once the endpoint is visible."""
self.misses = 0


def _live_logging_deploy_done(
sagemaker_client, endpoint_name, paginator, paginator_config, poll, not_found_budget=None
):
"""Return the ``DescribeEndpoint`` response once the endpoint leaves ``Creating``.

Streams the endpoint's CloudWatch logs on every poll and returns ``None`` while
the endpoint is still being created. The log group only exists once a container
has started, so an endpoint that fails before any instance is provisioned (for
example on InsufficientInstanceCapacity) never gets one: a missing log group
must not keep a finished deployment waiting.

Args:
not_found_budget (_EndpointNotFoundBudget): Optional bound on how many
consecutive polls a missing endpoint is tolerated. Without it a
missing endpoint is waited on indefinitely.
"""
stop = False
endpoint_status = None
try:
desc = sagemaker_client.describe_endpoint(EndpointName=endpoint_name)
endpoint_status = desc["EndpointStatus"]
except ClientError as e:
if e.response["Error"]["Code"] == "ValidationException":
if not_found_budget is not None:
not_found_budget.miss(e)
LOGGER.debug("Waiting for endpoint to become visible")
return None
raise e
if not_found_budget is not None:
not_found_budget.reset()

try:
# if endpoint is in an invalid state -> set stop to true, sleep, and flush the logs
if endpoint_status != "Creating":
stop = True
if endpoint_status == "InService":
LOGGER.info(
"Created endpoint with name %s. Waiting for it to be InService", endpoint_name
)
else:
time.sleep(poll)
# if endpoint is in an invalid state -> set stop to true, sleep, and flush the logs
if endpoint_status != "Creating":
stop = True
if endpoint_status == "InService":
LOGGER.info(
"Created endpoint with name %s. Waiting for it to be InService", endpoint_name
)
else:
time.sleep(poll)

try:
pages = paginator.paginate(
logGroupName=f"/aws/sagemaker/Endpoints/{endpoint_name}",
logStreamNamePrefix="AllTraffic/",
Expand All @@ -3016,17 +3065,13 @@ def _live_logging_deploy_done(sagemaker_client, endpoint_name, paginator, pagina
LOGGER.info(event["message"])
else:
LOGGER.debug("No log events available")

# if stop is true -> return the describe response and stop polling
if stop:
return desc
except ClientError as e:
if e.response["Error"]["Code"] == "ResourceNotFoundException":
LOGGER.debug("Waiting for endpoint log group to appear")
return None
raise e
if e.response["Error"]["Code"] != "ResourceNotFoundException":
raise e
LOGGER.debug("Waiting for endpoint log group to appear")

return None
# if stop is true -> return the describe response and stop polling
return desc if stop else None


def _deployment_entity_exists(describe_fn):
Expand Down
38 changes: 22 additions & 16 deletions sagemaker-core/src/sagemaker/core/utils/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -505,22 +505,28 @@ def __next__(self) -> T:

# Otherwise, get the next page of summaries by calling the list method with the next token if available
else:
if self.next_token:
response = getattr(self.client, self.list_method)(
NextToken=self.next_token, **self.list_method_kwargs
)
else:
response = getattr(self.client, self.list_method)(**self.list_method_kwargs)

self.summary_list = response.get(self.summaries_key, [])
self.next_token = response.get("NextToken", None)
self.index = 0

# If list_method returned an empty list, raise StopIteration
if len(self.summary_list) == 0:
raise StopIteration

return self.__next__()
while True:
previous_token = self.next_token
if self.next_token:
response = getattr(self.client, self.list_method)(
NextToken=self.next_token, **self.list_method_kwargs
)
else:
response = getattr(self.client, self.list_method)(**self.list_method_kwargs)

self.summary_list = response.get(self.summaries_key, [])
self.next_token = response.get("NextToken", None)
self.index = 0

if len(self.summary_list) > 0:
return self.__next__()

# A page can be empty and still carry a NextToken (the service
# filters results after paginating), so an empty page only ends
# the iteration when there is no further page. A NextToken that
# repeats would page forever, so it also ends the iteration.
if not self.next_token or self.next_token == previous_token:
raise StopIteration


def serialize(value: Any) -> Any:
Expand Down
35 changes: 35 additions & 0 deletions sagemaker-core/tests/unit/generated/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,6 +174,41 @@ def test_next_client_returns_empty_list(resource_iterator):
next(iterator)


def test_next_follows_next_token_past_empty_pages(resource_iterator):
iterator, client, _ = resource_iterator
client.list_training_jobs.side_effect = [
{"TrainingJobSummaries": [], "NextToken": "token-1"},
{"TrainingJobSummaries": [], "NextToken": "token-2"},
LIST_TRAINING_JOB_RESPONSE_WITHOUT_NEXT_TOKEN,
]

with patch.object(TrainingJob, "refresh"):
names = [job.training_job_name for job in iterator]

assert names == [
summary["TrainingJobName"]
for summary in LIST_TRAINING_JOB_RESPONSE_WITHOUT_NEXT_TOKEN["TrainingJobSummaries"]
]
assert client.list_training_jobs.call_args_list == [
call(),
call(NextToken="token-1"),
call(NextToken="token-2"),
]


def test_next_stops_when_empty_pages_repeat_the_next_token(resource_iterator):
iterator, client, _ = resource_iterator
client.list_training_jobs.side_effect = [
{"TrainingJobSummaries": [], "NextToken": "token-1"},
{"TrainingJobSummaries": [], "NextToken": "token-1"},
]

with pytest.raises(StopIteration):
next(iterator)

assert client.list_training_jobs.call_count == 2


def test_next_without_next_token(resource_iterator):
iterator, client, _ = resource_iterator
client.list_training_jobs.return_value = LIST_TRAINING_JOB_RESPONSE_WITHOUT_NEXT_TOKEN
Expand Down
120 changes: 119 additions & 1 deletion sagemaker-core/tests/unit/helper/test_session_helper.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,12 @@
from unittest.mock import Mock, patch
from botocore.exceptions import ClientError

from sagemaker.core.helper.session_helper import Session
from sagemaker.core.exceptions import UnexpectedStatusException
from sagemaker.core.helper.session_helper import (
Session,
_EndpointNotFoundBudget,
_live_logging_deploy_done,
)
from sagemaker.core.session_settings import SessionSettings


Expand Down Expand Up @@ -375,6 +380,119 @@ def test_wait_for_endpoint_failure(self, mock_boto_session, mock_sagemaker_clien
with pytest.raises(Exception, match="Error hosting endpoint"):
session.wait_for_endpoint("my-endpoint")

@patch("sagemaker.core.helper.session_helper._has_permission_for_live_logging")
@patch("time.sleep")
def test_live_logging_wait_raises_when_endpoint_fails_without_log_group(
self, mock_sleep, mock_permission, mock_boto_session, mock_sagemaker_client
):
"""A Failed endpoint that never got a log group ends the wait instead of hanging."""
mock_permission.return_value = True
failed = {
"EndpointStatus": "Failed",
"FailureReason": "Unable to provision requested ML compute capacity due to "
"InsufficientInstanceCapacity error.",
}
# A finite side_effect so a regression surfaces as a failure, not a hang.
mock_sagemaker_client.describe_endpoint.side_effect = [failed, failed]
paginator = mock_boto_session.client.return_value.get_paginator.return_value
paginator.paginate.side_effect = ClientError(
{"Error": {"Code": "ResourceNotFoundException"}}, "FilterLogEvents"
)
session = Session(boto_session=mock_boto_session, sagemaker_client=mock_sagemaker_client)

with pytest.raises(UnexpectedStatusException, match="InsufficientInstanceCapacity"):
session.wait_for_endpoint("my-endpoint", live_logging=True)

assert mock_sagemaker_client.describe_endpoint.call_count == 1


class TestLiveLoggingDeployDone:
"""Test _live_logging_deploy_done."""

RESOURCE_NOT_FOUND = ClientError(
{"Error": {"Code": "ResourceNotFoundException"}}, "FilterLogEvents"
)
ENDPOINT_NOT_FOUND = ClientError(
{"Error": {"Code": "ValidationException", "Message": "Could not find endpoint"}},
"DescribeEndpoint",
)

@patch("time.sleep")
def test_finished_endpoint_without_log_group_returns_desc(self, mock_sleep):
"""InService or Failed with no log group is finished, not "still waiting"."""
for status in ("InService", "Failed"):
client = Mock()
desc = {"EndpointStatus": status}
client.describe_endpoint.return_value = desc
paginator = Mock()
paginator.paginate.side_effect = self.RESOURCE_NOT_FOUND

assert _live_logging_deploy_done(client, "my-endpoint", paginator, {}, 5) == desc

def test_creating_endpoint_without_log_group_keeps_waiting(self):
"""Creating with no log group yet is still in progress."""
client = Mock()
client.describe_endpoint.return_value = {"EndpointStatus": "Creating"}
paginator = Mock()
paginator.paginate.side_effect = self.RESOURCE_NOT_FOUND

assert _live_logging_deploy_done(client, "my-endpoint", paginator, {}, 5) is None

def test_other_log_errors_are_raised(self):
"""Log-fetch errors other than a missing log group still propagate."""
client = Mock()
client.describe_endpoint.return_value = {"EndpointStatus": "InService"}
paginator = Mock()
paginator.paginate.side_effect = ClientError(
{"Error": {"Code": "ThrottlingException"}}, "FilterLogEvents"
)

with pytest.raises(ClientError):
_live_logging_deploy_done(client, "my-endpoint", paginator, {}, 5)

def test_missing_endpoint_is_waited_on_without_budget(self):
"""Without a budget a missing endpoint keeps the legacy "keep waiting" result."""
client = Mock()
client.describe_endpoint.side_effect = self.ENDPOINT_NOT_FOUND

assert _live_logging_deploy_done(client, "my-endpoint", Mock(), {}, 5) is None

def test_missing_endpoint_raises_once_budget_is_spent(self):
"""A missing endpoint is tolerated only for the budgeted number of polls."""
client = Mock()
client.describe_endpoint.side_effect = self.ENDPOINT_NOT_FOUND
budget = _EndpointNotFoundBudget(max_polls=2)

for _ in range(2):
assert (
_live_logging_deploy_done(
client, "my-endpoint", Mock(), {}, 5, not_found_budget=budget
)
is None
)
with pytest.raises(ClientError, match="Could not find endpoint"):
_live_logging_deploy_done(client, "my-endpoint", Mock(), {}, 5, not_found_budget=budget)

def test_budget_resets_once_endpoint_is_found(self):
"""Only consecutive "not found" polls count against the budget."""
client = Mock()
client.describe_endpoint.side_effect = [
self.ENDPOINT_NOT_FOUND,
{"EndpointStatus": "Creating"},
self.ENDPOINT_NOT_FOUND,
]
paginator = Mock()
paginator.paginate.return_value = []
budget = _EndpointNotFoundBudget(max_polls=1)

for _ in range(3):
assert (
_live_logging_deploy_done(
client, "my-endpoint", paginator, {}, 5, not_found_budget=budget
)
is None
)


class TestUpdateEndpoint:
"""Test update_endpoint method."""
Expand Down
Loading
Loading