From 326ffd002002a3c6b9e935031e12a97972c4d756 Mon Sep 17 00:00:00 2001 From: Yahor Kavaliou Date: Tue, 25 Aug 2026 11:35:47 +0200 Subject: [PATCH 1/3] failover infra --- .env.sample | 5 ++ config.sample.yaml | 2 + src/jointfm_client/__init__.py | 2 + src/jointfm_client/client.py | 92 +++++++++++++++++++++++++---- src/jointfm_client/configuration.py | 9 +++ src/jointfm_client/settings.py | 89 ++++++++++++++++++++++++++++ 6 files changed, 186 insertions(+), 13 deletions(-) diff --git a/.env.sample b/.env.sample index d72da9a..a852e5a 100644 --- a/.env.sample +++ b/.env.sample @@ -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/ diff --git a/config.sample.yaml b/config.sample.yaml index ecf4bea..96bb049 100644 --- a/config.sample.yaml +++ b/config.sample.yaml @@ -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 @@ -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 diff --git a/src/jointfm_client/__init__.py b/src/jointfm_client/__init__.py index a5810b5..6e86d5a 100644 --- a/src/jointfm_client/__init__.py +++ b/src/jointfm_client/__init__.py @@ -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, @@ -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", diff --git a/src/jointfm_client/client.py b/src/jointfm_client/client.py index 05a1a3a..2a1eda2 100644 --- a/src/jointfm_client/client.py +++ b/src/jointfm_client/client.py @@ -47,6 +47,7 @@ from jointfm_client.exceptions import ( JointFMConfigurationError, JointFMHTTPStatusError, + JointFMRequestError, JointFMServiceError, UnsupportedModelVersionError, ) @@ -73,6 +74,7 @@ r"n_samples exceeds the configured container cap:\s*" r"requested\s+(?P[0-9]+),\s*max\s+(?P[0-9]+)" ) +_FAILOVER_HTTP_STATUS_CODES = frozenset({502, 503, 504}) class JointFMClient: @@ -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( @@ -152,10 +155,30 @@ 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: @@ -195,14 +218,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 @@ -243,7 +266,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, ) @@ -303,18 +326,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) @@ -484,7 +505,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: @@ -498,10 +518,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, @@ -522,6 +539,55 @@ 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) + 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( diff --git a/src/jointfm_client/configuration.py b/src/jointfm_client/configuration.py index c2205f1..72a482f 100644 --- a/src/jointfm_client/configuration.py +++ b/src/jointfm_client/configuration.py @@ -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" @@ -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 @@ -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( @@ -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 @@ -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", diff --git a/src/jointfm_client/settings.py b/src/jointfm_client/settings.py index ce3b0c4..fd297f0 100644 --- a/src/jointfm_client/settings.py +++ b/src/jointfm_client/settings.py @@ -31,6 +31,7 @@ DATAROBOT_ENDPOINT_ENV, DEFAULT_CONFIG_PATH, EnvironmentVariableConfig, + JOINTFM_BACKUP_DEPLOYMENT_ID_ENV, JOINTFM_DEPLOYMENT_ID_ENV, JOINTFM_DEPLOYMENT_TARGET_ENV, JOINTFM_DEPLOYMENT_URL_ENV, @@ -74,6 +75,8 @@ class JointFMSettings: deployment_url: str | None = None deployment_target: str | None = None local_base_url: str | None = None + backup_deployment_id: str | None = None + backup_predict_url: str | None = None def load_settings( @@ -102,6 +105,7 @@ def load_settings( local_base_url = normalize_local_service_base_url( _required_env(env_values, environment.local_base_url) ) + _reject_backup_for_local_service(env_values, environment) return JointFMSettings( datarobot_endpoint=None, datarobot_api_token=None, @@ -126,6 +130,12 @@ def load_settings( ) deployment_url = build_hosted_deployment_url(datarobot_endpoint, deployment_id) predict_url = build_hosted_predict_url(datarobot_endpoint, deployment_id) + backup_deployment_id, backup_predict_url = _optional_backup_deployment( + env_values, + environment, + datarobot_endpoint=datarobot_endpoint, + primary_deployment_id=deployment_id, + ) return JointFMSettings( datarobot_endpoint=datarobot_endpoint, datarobot_api_token=datarobot_api_token, @@ -136,6 +146,8 @@ def load_settings( model_version=model_version, deployment_id=deployment_id, deployment_url=deployment_url, + backup_deployment_id=backup_deployment_id, + backup_predict_url=backup_predict_url, ) if selector_name == environment.deployment_url: @@ -144,6 +156,12 @@ def load_settings( ) deployment_id = deployment_id_from_hosted_deployment_url(deployment_url) predict_url = build_hosted_predict_url_from_deployment_url(deployment_url) + backup_deployment_id, backup_predict_url = _optional_backup_deployment( + env_values, + environment, + datarobot_endpoint=datarobot_endpoint, + primary_deployment_id=deployment_id, + ) return JointFMSettings( datarobot_endpoint=datarobot_endpoint, datarobot_api_token=datarobot_api_token, @@ -154,6 +172,8 @@ def load_settings( model_version=model_version, deployment_id=deployment_id, deployment_url=deployment_url, + backup_deployment_id=backup_deployment_id, + backup_predict_url=backup_predict_url, ) if selector_name == environment.predict_url: @@ -162,6 +182,12 @@ def load_settings( ) deployment_url = deployment_url_from_hosted_predict_url(predict_url) deployment_id = deployment_id_from_hosted_deployment_url(deployment_url) + backup_deployment_id, backup_predict_url = _optional_backup_deployment( + env_values, + environment, + datarobot_endpoint=datarobot_endpoint, + primary_deployment_id=deployment_id, + ) return JointFMSettings( datarobot_endpoint=datarobot_endpoint, datarobot_api_token=datarobot_api_token, @@ -172,6 +198,8 @@ def load_settings( model_version=model_version, deployment_id=deployment_id, deployment_url=deployment_url, + backup_deployment_id=backup_deployment_id, + backup_predict_url=backup_predict_url, ) deployment_target = _required_env(env_values, environment.deployment_target) @@ -188,6 +216,12 @@ def load_settings( datarobot_endpoint, normalized_deployment_id, ) + backup_deployment_id, backup_predict_url = _optional_backup_deployment( + env_values, + environment, + datarobot_endpoint=datarobot_endpoint, + primary_deployment_id=normalized_deployment_id, + ) return JointFMSettings( datarobot_endpoint=datarobot_endpoint, datarobot_api_token=datarobot_api_token, @@ -199,6 +233,8 @@ def load_settings( deployment_id=normalized_deployment_id, deployment_url=deployment_url, deployment_target=deployment_target, + backup_deployment_id=backup_deployment_id, + backup_predict_url=backup_predict_url, ) deployment_url_output = _optional_string_output(target_outputs, "deployment_url") @@ -206,6 +242,12 @@ def load_settings( deployment_url = normalize_hosted_deployment_url(deployment_url_output) deployment_id = deployment_id_from_hosted_deployment_url(deployment_url) predict_url = build_hosted_predict_url_from_deployment_url(deployment_url) + backup_deployment_id, backup_predict_url = _optional_backup_deployment( + env_values, + environment, + datarobot_endpoint=datarobot_endpoint, + primary_deployment_id=deployment_id, + ) return JointFMSettings( datarobot_endpoint=datarobot_endpoint, datarobot_api_token=datarobot_api_token, @@ -217,6 +259,8 @@ def load_settings( deployment_id=deployment_id, deployment_url=deployment_url, deployment_target=deployment_target, + backup_deployment_id=backup_deployment_id, + backup_predict_url=backup_predict_url, ) predict_url_output = _optional_string_output(target_outputs, "predict_url") @@ -224,6 +268,12 @@ def load_settings( predict_url = normalize_hosted_predict_url(predict_url_output) deployment_url = deployment_url_from_hosted_predict_url(predict_url) deployment_id = deployment_id_from_hosted_deployment_url(deployment_url) + backup_deployment_id, backup_predict_url = _optional_backup_deployment( + env_values, + environment, + datarobot_endpoint=datarobot_endpoint, + primary_deployment_id=deployment_id, + ) return JointFMSettings( datarobot_endpoint=datarobot_endpoint, datarobot_api_token=datarobot_api_token, @@ -235,6 +285,8 @@ def load_settings( deployment_id=deployment_id, deployment_url=deployment_url, deployment_target=deployment_target, + backup_deployment_id=backup_deployment_id, + backup_predict_url=backup_predict_url, ) raise JointFMConfigurationError( @@ -494,6 +546,43 @@ def _resolve_single_deployment_selector( return selector_names[0] +def _optional_backup_deployment( + env: Mapping[str, str], + environment: EnvironmentVariableConfig, + *, + datarobot_endpoint: str, + primary_deployment_id: str, +) -> tuple[str | None, str | None]: + """Return backup deployment identity when JOINTFM_BACKUP_DEPLOYMENT_ID is set.""" + value = env.get(environment.backup_deployment_id) + if value is None or value == "": + return None, None + backup_deployment_id = _normalize_non_whitespace_string( + value, JOINTFM_BACKUP_DEPLOYMENT_ID_ENV + ) + if "/" in backup_deployment_id: + raise JointFMConfigurationError( + f"{JOINTFM_BACKUP_DEPLOYMENT_ID_ENV} must be a deployment ID, not a URL" + ) + if backup_deployment_id == primary_deployment_id: + raise JointFMConfigurationError( + f"{JOINTFM_BACKUP_DEPLOYMENT_ID_ENV} must differ from the primary deployment" + ) + backup_predict_url = build_hosted_predict_url(datarobot_endpoint, backup_deployment_id) + return backup_deployment_id, backup_predict_url + + +def _reject_backup_for_local_service( + env: Mapping[str, str], environment: EnvironmentVariableConfig +) -> None: + value = env.get(environment.backup_deployment_id) + if value is None or value == "": + return + raise JointFMConfigurationError( + f"{JOINTFM_BACKUP_DEPLOYMENT_ID_ENV} requires a hosted DataRobot deployment" + ) + + def _load_pulumi_target_outputs( outputs_path: str, deployment_target: str, From 4ca2c63844f5c3d7a7730a1eecf49323f5106870 Mon Sep 17 00:00:00 2001 From: Yahor Kavaliou Date: Tue, 25 Aug 2026 11:42:52 +0200 Subject: [PATCH 2/3] uv format --- src/jointfm_client/client.py | 15 +++++++++++---- src/jointfm_client/settings.py | 4 +++- 2 files changed, 14 insertions(+), 5 deletions(-) diff --git a/src/jointfm_client/client.py b/src/jointfm_client/client.py index 2a1eda2..a519520 100644 --- a/src/jointfm_client/client.py +++ b/src/jointfm_client/client.py @@ -171,7 +171,9 @@ def health(self, *, cache: bool = False, refresh: bool = False) -> HealthMetadat if not self._should_failover(error): raise pinned_model_version = ( - None if self._health_metadata is None else self._health_metadata.model_version + None + if self._health_metadata is None + else self._health_metadata.model_version ) self._activate_backup() metadata = self._probe_health(cache=cache) @@ -547,7 +549,9 @@ def _post_predict_json(self, payload: Mapping[str, Any]) -> Mapping[str, Any]: if not self._should_failover(error): raise pinned_model_version = ( - None if self._health_metadata is None else self._health_metadata.model_version + None + if self._health_metadata is None + else self._health_metadata.model_version ) self._activate_backup() metadata = self._probe_health(cache=True) @@ -570,7 +574,9 @@ def _should_failover(self, error: BaseException) -> bool: 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") + 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 @@ -581,7 +587,8 @@ 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 + pinned_model_version is not None + and metadata.model_version != pinned_model_version ): raise UnsupportedModelVersionError( "Backup JointFM deployment model_version differs from the primary: " diff --git a/src/jointfm_client/settings.py b/src/jointfm_client/settings.py index fd297f0..f7a43b3 100644 --- a/src/jointfm_client/settings.py +++ b/src/jointfm_client/settings.py @@ -568,7 +568,9 @@ def _optional_backup_deployment( raise JointFMConfigurationError( f"{JOINTFM_BACKUP_DEPLOYMENT_ID_ENV} must differ from the primary deployment" ) - backup_predict_url = build_hosted_predict_url(datarobot_endpoint, backup_deployment_id) + backup_predict_url = build_hosted_predict_url( + datarobot_endpoint, backup_deployment_id + ) return backup_deployment_id, backup_predict_url From 0d026e02f1b048b17da4ffc8f2127713c1c3b008 Mon Sep 17 00:00:00 2001 From: Yahor Kavaliou Date: Tue, 25 Aug 2026 11:43:34 +0200 Subject: [PATCH 3/3] tests --- tests/test_settings.py | 39 +++++++++++ tests/test_transport.py | 150 ++++++++++++++++++++++++++++++++++++++++ 2 files changed, 189 insertions(+) diff --git a/tests/test_settings.py b/tests/test_settings.py index fa2623e..850cbd6 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -21,6 +21,7 @@ from jointfm_client 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, @@ -73,9 +74,47 @@ def test_load_settings_from_environment_with_deployment_id_builds_hosted_url() - "deployment-id/predictionsUnstructured" ) assert settings.health_url == settings.predict_url + assert settings.backup_deployment_id is None + assert settings.backup_predict_url is None assert "secret-token" not in repr(settings) +def test_load_settings_with_backup_deployment_id_builds_backup_predict_url() -> None: + """Load settings with backup deployment id builds backup predict url.""" + settings = load_settings( + env=_hosted_env(**{JOINTFM_BACKUP_DEPLOYMENT_ID_ENV: "backup-deployment-id"}), + dotenv_path=None, + ) + + assert settings.backup_deployment_id == "backup-deployment-id" + assert settings.backup_predict_url == ( + "https://app.datarobot.com/api/v2/deployments/" + "backup-deployment-id/predictionsUnstructured" + ) + + +def test_load_settings_rejects_backup_equal_to_primary_deployment() -> None: + """Load settings rejects backup equal to primary deployment.""" + with pytest.raises(JointFMConfigurationError, match="must differ"): + load_settings( + env=_hosted_env(**{JOINTFM_BACKUP_DEPLOYMENT_ID_ENV: "deployment-id"}), + dotenv_path=None, + ) + + +def test_load_settings_rejects_backup_with_local_service() -> None: + """Load settings rejects backup with local service.""" + with pytest.raises(JointFMConfigurationError, match="hosted DataRobot"): + load_settings( + env={ + JOINTFM_LOCAL_BASE_URL_ENV: "http://127.0.0.1:8080/", + JOINTFM_SCHEMA_VERSION_ENV: "v1", + JOINTFM_BACKUP_DEPLOYMENT_ID_ENV: "backup-deployment-id", + }, + dotenv_path=None, + ) + + def test_load_settings_with_local_service_base_url_builds_direct_urls() -> None: """Load settings with local service base url builds direct urls.""" settings = load_settings( diff --git a/tests/test_transport.py b/tests/test_transport.py index 4eb2d7d..a95ee7b 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -1086,6 +1086,156 @@ def test_client_forecast_requires_model_version_without_settings_or_health() -> ) +def _failover_settings() -> JointFMSettings: + """Hosted settings with a backup deployment for failover tests.""" + primary = ( + "https://app.datarobot.com/api/v2/deployments/" + "primary-id/predictionsUnstructured" + ) + backup = ( + "https://app.datarobot.com/api/v2/deployments/backup-id/predictionsUnstructured" + ) + return JointFMSettings( + datarobot_endpoint="https://app.datarobot.com/api/v2", + datarobot_api_token="secret-token", + health_url=primary, + predict_url=primary, + deployment_selector="deployment_id", + schema_version="v1", + model_version="jointfm-inference:0.2.0+ckpt.sdk-test", + deployment_id="primary-id", + backup_deployment_id="backup-id", + backup_predict_url=backup, + ) + + +class FailoverTransport: + """Transport that fails the primary URL and serves the backup.""" + + def __init__( + self, + *, + backup_model_version: str = "jointfm-inference:0.2.0+ckpt.sdk-test", + fail_primary_health: bool = False, + fail_primary_predict: bool = False, + ) -> None: + self.backup_model_version = backup_model_version + self.fail_primary_health = fail_primary_health + self.fail_primary_predict = fail_primary_predict + self.urls: list[str] = [] + self.post_count = 0 + + def get_json(self, url: str) -> Mapping[str, Any]: + """Get json.""" + raise AssertionError("hosted health must POST to the predict URL") + + def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: + """Post json.""" + self.post_count += 1 + self.urls.append(url) + is_primary = "primary-id" in url + if payload.get("request_type") == "health": + if is_primary and self.fail_primary_health: + raise JointFMRequestError("primary unreachable") + model_version = ( + "jointfm-inference:0.2.0+ckpt.sdk-test" + if is_primary + else self.backup_model_version + ) + return _health_payload(model_version=model_version) + if is_primary and self.fail_primary_predict: + raise JointFMHTTPStatusError( + "primary unavailable", + status_code=503, + response_body_excerpt="unavailable", + ) + return _forecast_response_payload() + + +def test_client_health_uses_primary_without_calling_backup() -> None: + """Client health uses primary without calling backup.""" + transport = FailoverTransport() + client = JointFMClient(settings=_failover_settings(), transport=transport) + + metadata = client.health(cache=True) + + assert metadata.model_version == "jointfm-inference:0.2.0+ckpt.sdk-test" + assert transport.urls == [_failover_settings().predict_url] + assert client._using_backup is False + + +def test_client_health_fails_over_to_backup_on_primary_transport_error() -> None: + """Client health fails over to backup on primary transport error.""" + transport = FailoverTransport(fail_primary_health=True) + client = JointFMClient(settings=_failover_settings(), transport=transport) + + metadata = client.health(cache=True) + + assert metadata.model_version == "jointfm-inference:0.2.0+ckpt.sdk-test" + assert client._using_backup is True + assert client.predict_url == _failover_settings().backup_predict_url + assert transport.urls == [ + _failover_settings().predict_url, + _failover_settings().backup_predict_url, + ] + + +def test_client_predict_fails_over_once_and_stays_on_backup() -> None: + """Client predict fails over once and stays on backup.""" + transport = FailoverTransport(fail_primary_predict=True) + client = JointFMClient(settings=_failover_settings(), transport=transport) + client.health(cache=True) + payload = { + "schema_version": "v1", + "model_version": "jointfm-inference:0.2.0+ckpt.sdk-test", + } + + first = client.predict(payload) + second = client.predict(payload) + + assert first == _forecast_response_payload() + assert second == _forecast_response_payload() + assert client._using_backup is True + assert transport.urls == [ + _failover_settings().predict_url, + _failover_settings().predict_url, + _failover_settings().backup_predict_url, + _failover_settings().backup_predict_url, + _failover_settings().backup_predict_url, + ] + + +def test_client_failover_rejects_backup_with_different_model_version() -> None: + """Client failover rejects backup with different model version.""" + settings = _failover_settings() + settings = JointFMSettings( + datarobot_endpoint=settings.datarobot_endpoint, + datarobot_api_token=settings.datarobot_api_token, + health_url=settings.health_url, + predict_url=settings.predict_url, + deployment_selector=settings.deployment_selector, + schema_version=settings.schema_version, + model_version=None, + deployment_id=settings.deployment_id, + backup_deployment_id=settings.backup_deployment_id, + backup_predict_url=settings.backup_predict_url, + ) + transport = FailoverTransport( + fail_primary_predict=True, + backup_model_version="jointfm-inference:9.9.9+ckpt.other", + ) + client = JointFMClient(settings=settings, transport=transport) + client.health(cache=True) + + with pytest.raises(UnsupportedModelVersionError, match="Backup"): + client.predict( + { + "schema_version": "v1", + "model_version": "jointfm-inference:0.2.0+ckpt.sdk-test", + } + ) + + def _start_json_server( statuses: list[HTTPStatus], *,