Skip to content
Closed
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
5 changes: 5 additions & 0 deletions .env.sample
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,11 @@ JOINTFM_SCHEMA_VERSION=v1
# 1. Deployment ID: SDK builds the hosted predictionsUnstructured URL.
# JOINTFM_DEPLOYMENT_ID=

# Optional same-checkpoint backup. On transport or 502/503/504 failures the client
# switches to this deployment for the rest of its lifetime when health model_version
# matches the primary.
# JOINTFM_BACKUP_DEPLOYMENT_ID=

# 2. Deployment URL: SDK appends predictionsUnstructured.
# JOINTFM_DEPLOYMENT_URL=https://app.datarobot.com/api/v2/deployments/<deployment-id>

Expand Down
2 changes: 2 additions & 0 deletions config.sample.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ environment:
datarobot_endpoint: DATAROBOT_ENDPOINT
datarobot_api_token: DATAROBOT_API_TOKEN
deployment_id: JOINTFM_DEPLOYMENT_ID
backup_deployment_id: JOINTFM_BACKUP_DEPLOYMENT_ID
deployment_url: JOINTFM_DEPLOYMENT_URL
predict_url: JOINTFM_PREDICT_URL
deployment_target: JOINTFM_DEPLOYMENT_TARGET
Expand All @@ -17,6 +18,7 @@ deployment:
datarobot_endpoint: null
datarobot_api_token: null
deployment_id: null
backup_deployment_id: null
deployment_url: null
predict_url: null
deployment_target: null
Expand Down
2 changes: 2 additions & 0 deletions src/jointfm_client/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,6 +107,7 @@
from jointfm_client.settings import (
DATAROBOT_API_TOKEN_ENV,
DATAROBOT_ENDPOINT_ENV,
JOINTFM_BACKUP_DEPLOYMENT_ID_ENV,
JOINTFM_DEPLOYMENT_ID_ENV,
JOINTFM_DEPLOYMENT_TARGET_ENV,
JOINTFM_DEPLOYMENT_URL_ENV,
Expand Down Expand Up @@ -169,6 +170,7 @@
"HostedDeploymentConfig",
"IMPORT_NAMESPACE",
"MeanForecastResult",
"JOINTFM_BACKUP_DEPLOYMENT_ID_ENV",
"JOINTFM_DEPLOYMENT_ID_ENV",
"JOINTFM_DEPLOYMENT_TARGET_ENV",
"JOINTFM_DEPLOYMENT_URL_ENV",
Expand Down
99 changes: 86 additions & 13 deletions src/jointfm_client/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@
from jointfm_client.exceptions import (
JointFMConfigurationError,
JointFMHTTPStatusError,
JointFMRequestError,
JointFMServiceError,
UnsupportedModelVersionError,
)
Expand All @@ -73,6 +74,7 @@
r"n_samples exceeds the configured container cap:\s*"
r"requested\s+(?P<requested>[0-9]+),\s*max\s+(?P<cap>[0-9]+)"
)
_FAILOVER_HTTP_STATUS_CODES = frozenset({502, 503, 504})


class JointFMClient:
Expand Down Expand Up @@ -108,6 +110,7 @@ def __init__(
self._datarobot_request_id_headers = datarobot_request_id_headers
self._health_metadata: HealthMetadata | None = None
self._sample_batch_cap: int | None = None
self._using_backup = False

@classmethod
def from_env(
Expand Down Expand Up @@ -152,10 +155,32 @@ def health(self, *, cache: bool = False, refresh: bool = False) -> HealthMetadat
deployment gateway only proxies the unstructured prediction route; the
container short-circuits that body before any schema or model version
validation and returns the same typed health payload.

When ``JOINTFM_BACKUP_DEPLOYMENT_ID`` is configured and the primary probe
fails with a transport or gateway error, the client switches to the backup
deployment for the rest of its lifetime, but only when the backup advertises
the same ``model_version`` already pinned for this client (or any version on
the first successful probe).
"""
if cache and not refresh and self._health_metadata is not None:
return self._health_metadata

try:
return self._probe_health(cache=cache)
except Exception as error:
if not self._should_failover(error):
raise
pinned_model_version = (
None
if self._health_metadata is None
else self._health_metadata.model_version
)
self._activate_backup()
metadata = self._probe_health(cache=cache)
self._ensure_backup_model_version(metadata, pinned_model_version)
return metadata

def _probe_health(self, *, cache: bool) -> HealthMetadata:
if self._uses_predict_route_for_health():
payload = self._fetch_hosted_health_payload()
else:
Expand Down Expand Up @@ -195,14 +220,14 @@ def _fetch_hosted_health_payload(self) -> Mapping[str, Any]:

def predict(self, payload: Mapping[str, Any]) -> Mapping[str, Any]:
"""Submit one V1 JSON prediction payload to the configured endpoint."""
predict_url = self._require_predict_url("predict")
self._require_predict_url("predict")
model_version = payload.get("model_version")
if not isinstance(model_version, str):
raise JointFMConfigurationError(
"JointFMClient.predict() requires payload['model_version']"
)
self._resolve_model_version(model_version=model_version)
response_payload = self._transport_for_request().post_json(predict_url, payload)
response_payload = self._post_predict_json(payload)
ForecastResponse.raise_for_errors(response_payload)
return response_payload

Expand Down Expand Up @@ -243,7 +268,7 @@ def forecast(
| None = None,
) -> ForecastResponse:
"""Build and submit a forecast request from tabular history inputs."""
predict_url = self._require_predict_url("forecast")
self._require_predict_url("forecast")
resolved_model_version = self._resolve_model_version(
model_version=model_version,
)
Expand Down Expand Up @@ -303,18 +328,16 @@ def forecast(
)
sample_cap = self._resolve_sample_batch_cap(payload)
if sample_cap is not None:
return self._forecast_sample_batches(predict_url, payload, sample_cap)
return self._forecast_sample_batches(payload, sample_cap)

try:
response_payload = self._transport_for_request().post_json(
predict_url, payload
)
response_payload = self._post_predict_json(payload)
except JointFMHTTPStatusError as error:
sample_cap = _sample_batch_cap_from_error(error, payload)
if sample_cap is None:
raise
self._sample_batch_cap = sample_cap
return self._forecast_sample_batches(predict_url, payload, sample_cap)
return self._forecast_sample_batches(payload, sample_cap)

return _forecast_response_from_payload(response_payload, payload)

Expand Down Expand Up @@ -484,7 +507,6 @@ def _resolve_sample_batch_cap(self, payload: Mapping[str, Any]) -> int | None:

def _forecast_sample_batches(
self,
predict_url: str,
payload: Mapping[str, Any],
sample_cap: int,
) -> SampleForecastResult:
Expand All @@ -498,10 +520,7 @@ def _forecast_sample_batches(
batch_payload = dict(payload)
batch_payload["n_samples"] = batch_samples
_set_batch_seed(batch_payload, batch_index)
response_payload = self._transport_for_request().post_json(
predict_url,
batch_payload,
)
response_payload = self._post_predict_json(batch_payload)
batch_result = _forecast_response_from_payload(
response_payload,
batch_payload,
Expand All @@ -522,6 +541,60 @@ def _forecast_sample_batches(
f"JointFM forecast response violated the V1 contract: {error}"
) from error

def _post_predict_json(self, payload: Mapping[str, Any]) -> Mapping[str, Any]:
predict_url = self._require_predict_url("predict")
try:
return self._transport_for_request().post_json(predict_url, payload)
except Exception as error:
if not self._should_failover(error):
raise
pinned_model_version = (
None
if self._health_metadata is None
else self._health_metadata.model_version
)
self._activate_backup()
metadata = self._probe_health(cache=True)
self._ensure_backup_model_version(metadata, pinned_model_version)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Sticky failover before verify

High Severity

Failover calls _activate_backup before the backup health probe and _ensure_backup_model_version succeed. On a version mismatch or unreachable backup, _using_backup stays true and URLs already point at the backup, so the client never returns to the primary and later calls keep failing or use the wrong deployment.

Suggested change
self._ensure_backup_model_version(metadata, pinned_model_version)
primary_predict_url = self.predict_url
primary_health_url = self.health_url
primary_health_metadata = self._health_metadata
primary_sample_batch_cap = self._sample_batch_cap
self._activate_backup()
try:
metadata = self._probe_health(cache=True)
self._ensure_backup_model_version(metadata, pinned_model_version)
except Exception:
self.predict_url = primary_predict_url
self.health_url = primary_health_url
self._using_backup = False
self._health_metadata = primary_health_metadata
self._sample_batch_cap = primary_sample_batch_cap
raise
Additional Locations (1)
Fix in Cursor Fix in Web

Triggered by team rule: suggestion rule

Reviewed by Cursor Bugbot for commit 0d026e0. Configure here.

return self._transport_for_request().post_json(
self._require_predict_url("predict"), payload
)

def _should_failover(self, error: BaseException) -> bool:
if self._using_backup:
return False
if self.settings is None or self.settings.backup_predict_url is None:
return False
if isinstance(error, JointFMRequestError):
return True
return (
isinstance(error, JointFMHTTPStatusError)
and error.status_code in _FAILOVER_HTTP_STATUS_CODES
)

def _activate_backup(self) -> None:
if self.settings is None or self.settings.backup_predict_url is None:
raise JointFMConfigurationError(
"JointFMClient backup deployment is not configured"
)
self.predict_url = self.settings.backup_predict_url
self.health_url = self.settings.backup_predict_url
self._using_backup = True
self._health_metadata = None
self._sample_batch_cap = None

def _ensure_backup_model_version(
self, metadata: HealthMetadata, pinned_model_version: str | None
) -> None:
if (
pinned_model_version is not None
and metadata.model_version != pinned_model_version
):
raise UnsupportedModelVersionError(
"Backup JointFM deployment model_version differs from the primary: "
f"expected {pinned_model_version!r}, got {metadata.model_version!r}"
)

def _require_settings(self, method_name: str) -> JointFMSettings:
if self.settings is None:
raise JointFMConfigurationError(
Expand Down
9 changes: 9 additions & 0 deletions src/jointfm_client/configuration.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,7 @@ class EnvironmentVariableConfig(_ConfigModel):
datarobot_endpoint: str = "DATAROBOT_ENDPOINT"
datarobot_api_token: str = "DATAROBOT_API_TOKEN"
deployment_id: str = "JOINTFM_DEPLOYMENT_ID"
backup_deployment_id: str = "JOINTFM_BACKUP_DEPLOYMENT_ID"
deployment_url: str = "JOINTFM_DEPLOYMENT_URL"
predict_url: str = "JOINTFM_PREDICT_URL"
deployment_target: str = "JOINTFM_DEPLOYMENT_TARGET"
Expand Down Expand Up @@ -92,6 +93,7 @@ class HostedDeploymentConfig(_ConfigModel):
datarobot_endpoint: str | None = None
datarobot_api_token: str | None = Field(default=None, repr=False)
deployment_id: str | None = None
backup_deployment_id: str | None = None
deployment_url: str | None = None
predict_url: str | None = None
deployment_target: str | None = None
Expand All @@ -113,6 +115,9 @@ def to_environment_values(
values, environment.datarobot_api_token, self.datarobot_api_token
)
_set_if_configured(values, environment.deployment_id, self.deployment_id)
_set_if_configured(
values, environment.backup_deployment_id, self.backup_deployment_id
)
_set_if_configured(values, environment.deployment_url, self.deployment_url)
_set_if_configured(values, environment.predict_url, self.predict_url)
_set_if_configured(
Expand Down Expand Up @@ -355,6 +360,9 @@ def _set_if_configured(values: dict[str, str], name: str, value: str | None) ->
DATAROBOT_ENDPOINT_ENV: Final = DEFAULT_ENVIRONMENT_CONFIG.datarobot_endpoint
DATAROBOT_API_TOKEN_ENV: Final = DEFAULT_ENVIRONMENT_CONFIG.datarobot_api_token
JOINTFM_DEPLOYMENT_ID_ENV: Final = DEFAULT_ENVIRONMENT_CONFIG.deployment_id
JOINTFM_BACKUP_DEPLOYMENT_ID_ENV: Final = (
DEFAULT_ENVIRONMENT_CONFIG.backup_deployment_id
)
JOINTFM_DEPLOYMENT_URL_ENV: Final = DEFAULT_ENVIRONMENT_CONFIG.deployment_url
JOINTFM_PREDICT_URL_ENV: Final = DEFAULT_ENVIRONMENT_CONFIG.predict_url
JOINTFM_DEPLOYMENT_TARGET_ENV: Final = DEFAULT_ENVIRONMENT_CONFIG.deployment_target
Expand Down Expand Up @@ -432,6 +440,7 @@ def _set_if_configured(values: dict[str, str], name: str, value: str | None) ->
"ForecastConfig",
"ForecastCsvConfig",
"HostedDeploymentConfig",
"JOINTFM_BACKUP_DEPLOYMENT_ID_ENV",
"JOINTFM_DEPLOYMENT_ID_ENV",
"JOINTFM_DEPLOYMENT_TARGET_ENV",
"JOINTFM_DEPLOYMENT_URL_ENV",
Expand Down
Loading
Loading