diff --git a/.env.sample b/.env.sample index e98b8c8..896f455 100644 --- a/.env.sample +++ b/.env.sample @@ -4,7 +4,7 @@ DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= # JointFM compatibility metadata for the selected deployment. -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.2.0+ckpt.fin-2026-05-22 diff --git a/AGENTS.md b/AGENTS.md index 823357d..faeeb9d 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -1 +1 @@ -Agents must ground their thoughts in facts, not assumptions: before planning, claiming, or editing, read the relevant material — source, configs, data, docs, test output — and only act on beliefs backed by something just read or run. Agents must run ad-hoc Python and CLI commands through `task run -- ` (which wraps `uv run` with the project's canonical parameters), never bare `python` or hand-written `uv run` invocations — this keeps every invocation from accidentally re-resolving or mutating the `.venv` and `uv.lock`. Agents must explain what each command will do and why it is being run before running it. Before asking the user a question, agents must first explain the corresponding context and terminology — what the question concerns, why it arises, and what any project-specific terms mean — so the user can answer without digging through the code themselves. Agents must run `task pre-commit` and fix all reported issues before reporting success to the user. After `task pre-commit` succeeds, show the diff and explain why each change is necessary before reporting success to the user. When what an agent reports is a set of items measured or classified along the same few dimensions — files and their status, config keys and their values, metrics before and after, options and their trade-offs — it must be presented as a Markdown table instead of prose or a bullet list, with short cells and only the columns that carry information; prose stays the right form for a single item, a narrative explanation, or an ordered procedure, because a one-row table, or one whose cells are sentences, is harder to read than the paragraph it replaced. Agents must never perform destructive or state-changing git operations unless the user explicitly instructs the agent to run that specific operation — no `git push --force`, no `git reset --hard`, no `git stash` (which hides uncommitted work), no branch/tag deletion, no history rewrites (`rebase`, `commit --amend` on published commits, `filter-branch`), no `git clean -fdx`, no `--no-verify` to bypass hooks, and no discarding of uncommitted work. Read-only inspection commands (`git status`, `git diff`, `git log`, `git show`) are always allowed. If a task seems to require a state-changing git operation and the user has not explicitly asked for it, stop and ask the user to run it. Agents must never mention temporary planning identifiers — phase names or numbers such as "Phase 3g", milestone, sprint, or ticket codes — in docstrings, comments, help or description strings, error messages, or test docstrings; those labels are deleted when the plan is retired and leave readers with a dangling reference nobody can decode, so describe the concept by its lasting behavior or stable configuration key instead (phase names belong only in roadmap and planning docs, and linking to such a doc by its actual filename is fine). Agents must not mention or recommend key rotation (rotating API keys, tokens, or other credentials) — the user takes care of key rotation themselves. Agents must keep comments in configuration files value-independent: a comment attached to a configuration parameter — in `config.yaml`/`config.sample.yaml`, a `Taskfile`, `pyproject.toml`, or any other settings file — may only state what remains true whatever the value is (what the parameter controls, what its options mean, which invariant or trade-off it participates in), and must never restate the current value, derive from it, or do arithmetic on it (`# 64 of 576 rows`, `# 2x the 288-step horizon`, `# keeps 90% of draws at full context`); configuration is retuned constantly, so such a comment turns false the moment somebody edits the value on the line below it and nothing catches the drift, so express a needed derivation as a rule over the parameter or record concrete numbers in a dated notes doc where they read as a historical observation. ripgrep (`rg`) is available (installed by `task setup` via `task install:ripgrep`); prefer it for fast code and text search. +Agents must ground their thoughts in facts, not assumptions: before planning, claiming, or editing, read the relevant material — source, configs, data, docs, test output — and only act on beliefs backed by something just read or run. Agents must run ad-hoc Python and CLI commands through `task run -- ` (which wraps `uv run` with the project's canonical parameters), never bare `python` or hand-written `uv run` invocations — this keeps every invocation from accidentally re-resolving or mutating the `.venv` and `uv.lock`. Agents must explain what each command will do and why it is being run before running it. Before asking the user a question, agents must first explain the corresponding context and terminology — what the question concerns, why it arises, and what any project-specific terms mean — so the user can answer without digging through the code themselves. Agents must run `task pre-commit` and fix all reported issues before reporting success to the user. After `task pre-commit` succeeds, show the diff and explain why each change is necessary before reporting success to the user. Agents must lead a response with its outcome — what now works, what the answer is, and anything the user must decide or act on — within the first few lines, so a reader who stops there misses nothing that changes what they do next; the evidence and the mechanism follow underneath, ordered by what the user cannot reconstruct alone, and what the user can already see is left out: no restatement of the request, no narration of the steps taken or the files visited, no walk through a diff that is shown anyway, and no closing summary that repeats the opening. That focus governs only what is *not* written — it never removes a `module.py:120` citation, a measured number, a side-effect or migration impact, an unrelated defect the agent noticed, or a question the user must answer; when focus and justification collide, the prose goes and the evidence stays. When what an agent reports is a set of items measured or classified along the same few dimensions — files and their status, config keys and their values, metrics before and after, options and their trade-offs — it must be presented as a Markdown table instead of prose or a bullet list, with short cells and only the columns that carry information; prose stays the right form for a single item, a narrative explanation, or an ordered procedure, because a one-row table, or one whose cells are sentences, is harder to read than the paragraph it replaced. Agents must never perform destructive or state-changing git operations unless the user explicitly instructs the agent to run that specific operation — no `git push --force`, no `git reset --hard`, no `git stash` (which hides uncommitted work), no branch/tag deletion, no history rewrites (`rebase`, `commit --amend` on published commits, `filter-branch`), no `git clean -fdx`, no `--no-verify` to bypass hooks, and no discarding of uncommitted work. Read-only inspection commands (`git status`, `git diff`, `git log`, `git show`) are always allowed. If a task seems to require a state-changing git operation and the user has not explicitly asked for it, stop and ask the user to run it. Agents must never mention temporary planning identifiers — phase names or numbers such as "Phase 3g", milestone, sprint, or ticket codes — in docstrings, comments, help or description strings, error messages, or test docstrings; those labels are deleted when the plan is retired and leave readers with a dangling reference nobody can decode, so describe the concept by its lasting behavior or stable configuration key instead (phase names belong only in roadmap and planning docs, and linking to such a doc by its actual filename is fine). Agents must not mention or recommend key rotation (rotating API keys, tokens, or other credentials) — the user takes care of key rotation themselves. Agents must keep comments in configuration files value-independent: a comment attached to a configuration parameter — in `config.yaml`/`config.sample.yaml`, a `Taskfile`, `pyproject.toml`, or any other settings file — may only state what remains true whatever the value is (what the parameter controls, what its options mean, which invariant or trade-off it participates in), and must never restate the current value, derive from it, or do arithmetic on it (`# 64 of 576 rows`, `# 2x the 288-step horizon`, `# keeps 90% of draws at full context`); configuration is retuned constantly, so such a comment turns false the moment somebody edits the value on the line below it and nothing catches the drift, so express a needed derivation as a rule over the parameter or record concrete numbers in a dated notes doc where they read as a historical observation. ripgrep (`rg`) is available (installed by `task setup` via `task install:ripgrep`); prefer it for fast code and text search. diff --git a/README.md b/README.md index 50eab26..2b1e7f0 100644 --- a/README.md +++ b/README.md @@ -10,7 +10,7 @@ The SDK targets the DataRobot-hosted unstructured prediction route and the same - Import namespace: `jointfm_client` - Supported Python: `>=3.11` - Current SDK package version: `0.8.0` -- Current JointFM service schema: `schema_version="v3"` +- Current JointFM service schema: `schema_version="v4"` The public API shape is a synchronous low-level `JointFMClient` with `health()`, `health_instances()`, and `predict(payload)` methods plus high-level `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, and `forecast_quantiles(...)` helpers. The SDK is not a proxy service; callers use it as a local Python library that talks to the hosted or local JointFM endpoint. @@ -50,7 +50,7 @@ Example deployment configuration: deployment: datarobot_endpoint: https://app.datarobot.com/api/v2 datarobot_api_token: - schema_version: v3 + schema_version: v4 deployment_id: # Optional model-version pin; the SDK discovers it from /healthz when unset: # model_version: jointfm-inference:0.3.0+ckpt.fin-2026-05-22 @@ -68,7 +68,7 @@ Equivalent `.env` deployment configuration: ```dotenv DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 JOINTFM_DEPLOYMENT_ID= # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.fin-2026-05-22 @@ -78,7 +78,7 @@ Equivalent local REST configuration for a service started from the `joint` repos ```dotenv JOINTFM_LOCAL_BASE_URL=http://127.0.0.1:8080 -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: # JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.fin_i504_o63_f0_t10_h16l16_mam7_af_t3r1_cnn_k3l4_hpst_h16l2_studentt_m4cr2df8skew ``` @@ -167,7 +167,7 @@ The bootstrap helper resolves the nearest src-layout Python project root, switch The current forecast request contract is: -- `schema_version`: exactly `"v3"`, configured as `JOINTFM_SCHEMA_VERSION` for `from_env()` clients +- `schema_version`: exactly `"v4"`, configured as `JOINTFM_SCHEMA_VERSION` for `from_env()` clients - `model_version`: exact model version advertised by `/healthz` or otherwise selected by the caller. Optional for `from_env()` clients: when `JOINTFM_MODEL_VERSION` is unset the SDK reads it from `/healthz` on first use; when set it acts as a drift-detection pin - `query_mode`: `"forecast"` for the unconditional forecast, or `"condition"` for a conditional query at one future position; the high-level helpers set it from whether a `condition` block was passed - `return_mode`: one of `"mean"`, `"samples"`, or `"quantiles"` @@ -184,10 +184,10 @@ The `condition` query mode asks for the model's joint distribution at one future A column carries at most one condition, of either kind: -- `EqualityCondition(column, value)` pins the column to a finite value. The pinned column leaves the read-out set, so it must not appear in `requested_columns`. +- `EqualityCondition(column, value)` pins the column to a finite value. It stays nameable in `requested_columns` and reads back the value the request supplied, so a scenario answer lines up column for column with an unconditioned one. - `IntervalCondition(column, lower=None, upper=None)` confines the column to a range; `None` leaves that side open, and at least one side must be bounded. The column stays readable, and what comes back is its distribution inside the range. -Every column without a condition is a read-out column, and at least one must remain. Pass the block to `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, or `forecast_quantiles(...)`: +Every column without a condition is a read-out column, and at least one must remain. `requested_columns` chooses what the response carries independently of that, defaulting to every declared column in declared order. Pass the block to `forecast(...)`, `forecast_mean(...)`, `forecast_samples(...)`, or `forecast_quantiles(...)`: ```python from jointfm_client import ConditionBlock, EqualityCondition, IntervalCondition @@ -227,7 +227,7 @@ Successful forecast responses preserve `schema_version`, `image_version`, `model ```json { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", @@ -242,7 +242,7 @@ Known error codes are `VALIDATION_ERROR`, `UNSUPPORTED_HEAD_QUERY_COMBINATION`, ## Compatibility Policy -The SDK supports only `schema_version="v3"`. `validate_service_metadata()` checks `/healthz` metadata and raises typed compatibility errors before prediction if the service advertises a different schema, an unexpected model version, mode capabilities outside the recorded service contract, or an unsupported `decoding_strategy`. Return modes and time-index modes must match the SDK's lists exactly. Query modes and condition kinds are derived by the service from the mounted head, so a deployment may advertise fewer of them than the SDK knows; it must advertise at least one query mode and nothing the SDK does not know. +The SDK supports only `schema_version="v4"`. `validate_service_metadata()` checks `/healthz` metadata and raises typed compatibility errors before prediction if the service advertises a different schema, an unexpected model version, mode capabilities outside the recorded service contract, or an unsupported `decoding_strategy`. Return modes and time-index modes must match the SDK's lists exactly. Query modes and condition kinds are derived by the service from the mounted head, so a deployment may advertise fewer of them than the SDK knows; it must advertise at least one query mode and nothing the SDK does not know. Callers should pass an expected `model_version` when they already know which deployment artifact they intend to use. A mismatch is treated as a hard compatibility error rather than silently downgrading, guessing, or retrying another model. @@ -284,7 +284,7 @@ Create `.env` from `.env.sample` or set the same values in your shell. A hosted ```dotenv DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2 DATAROBOT_API_TOKEN= -JOINTFM_SCHEMA_VERSION=v3 +JOINTFM_SCHEMA_VERSION=v4 JOINTFM_DEPLOYMENT_ID= # Or: JOINTFM_DEPLOYMENT_IDS=chevron-id,research-id # Optional drift-detection pin; the SDK discovers the model version from /healthz when unset: diff --git a/config.sample.yaml b/config.sample.yaml index 768f4c5..a49d28e 100644 --- a/config.sample.yaml +++ b/config.sample.yaml @@ -52,7 +52,7 @@ transport: - X-DataRobot-Execution-ID user_agent_header: User-Agent forecast: - schema_version: v3 + schema_version: v4 query_mode: forecast return_mode: mean time_index_mode: ordinal diff --git a/docs/api-reference.md b/docs/api-reference.md index e8de616..efc7f82 100644 --- a/docs/api-reference.md +++ b/docs/api-reference.md @@ -18,7 +18,7 @@ This reference covers the supported public Python surface exported by `jointfm_c | `DataFrameSchema` | Describes tabular history layout. Fields are `columns`, `time_index_mode`, `time_column`, `time_scale_seconds`, `use_local_normalized_time`, `calendar_id`, and `timezone`. | | `ForecastRequestMetadata` | Holds `schema_version`, `model_version`, `query_mode` (`forecast` or `condition`), and `return_mode` for one forecast request. | | `ForecastRequest` | Validated request object that combines metadata, schema, history rows, query times, requested columns, sample or quantile controls, `seed`, an optional `condition` block, and the optional `query_rows` a scored request supplies, then emits a JSON-compatible payload with `to_payload()`. `query_mode="condition"` and a `condition` block must appear together, as must `return_mode="log_prob"` and `query_rows`, and both are validated against the request's own schema and `query_times` before any round trip. | -| `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The pinned column leaves the read-out set, so it must not appear in `requested_columns`. | +| `EqualityCondition` | One column of the conditioned position pinned to a finite `value`. The column may be named in `requested_columns` like any other, and reads back the pinned value the request supplied. | | `IntervalCondition` | One column of the conditioned position confined to `[lower, upper]`; `None` leaves a side open, at least one side must be bounded, and `lower < upper`. The column stays readable and the response describes its distribution inside the range. | | `ConditionBlock` | Every condition of one request: `query_time_index` into the request's `query_times` and a sequence of `EqualityCondition` or `IntervalCondition` values, at most one per column. `kinds`, `pinned_columns`, and `conditioned_columns` expose what the capability gate and the request validation need. | | `ConditionPlausibility` | What the model thinks of the conditions it was given: `equality_log_density` of the pinned values and `region_log_probability` of the interval region, each `None` when the request carried no condition of that kind. Reported by the service and never refused on. | @@ -80,7 +80,7 @@ All SDK-specific exceptions inherit from `JointFMError`. | `JointFMHTTPStatusError` | The service returns an HTTP error status. | | `JointFMServiceError` | A response body contains non-empty JointFM `errors`, including the case where HTTP status unexpectedly succeeded. | | `JointFMCompatibilityError` | Base class for fail-fast service compatibility failures. | -| `UnsupportedSchemaVersionError` | The service or response advertises a schema version other than `v3`. | +| `UnsupportedSchemaVersionError` | The service or response advertises a schema version other than `v4`. | | `UnsupportedModelVersionError` | The service or response model version differs from the configured or requested version. | | `UnsupportedServiceContractError` | The service-health payload advertises mode capabilities or a `decoding_strategy` outside the recorded service contract, or a condition request targets a deployment that does not advertise the `condition` mode or one of the block's condition kinds. | @@ -126,7 +126,7 @@ All SDK-specific exceptions inherit from `JointFMError`. | --- | --- | --- | | `DATAROBOT_ENDPOINT` | Hosted calls | HTTPS DataRobot API v2 endpoint, normalized without a trailing slash and required to end in `/api/v2`. | | `DATAROBOT_API_TOKEN` | Hosted calls | Non-empty, whitespace-free API token used in the hosted bearer authorization header. | -| `JOINTFM_SCHEMA_VERSION` | Hosted calls | Request schema pin. The SDK supports only `v3`. | +| `JOINTFM_SCHEMA_VERSION` | Hosted calls | Request schema pin. The SDK supports only `v4`. | | `JOINTFM_MODEL_VERSION` | Hosted calls | Exact JointFM deployment model version expected from the service-health payload and prediction responses. | | `JOINTFM_DEPLOYMENT_ID` | One selector | Deployment ID used to build hosted health and prediction URLs. | | `JOINTFM_DEPLOYMENT_IDS` | One selector | Comma-separated hosted deployment IDs for round-robin load balancing (at least two unique IDs). Mutually exclusive with other selectors. Peers must share `model_version` and `checkpoint_version`. `health()` uses the minimum reachable `max_sample_count` as the sample-batch cap; `health_instances()` sums reachable caps for overall parallel capacity and reports topology. | @@ -157,7 +157,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Required | Description | | --- | --- | --- | | `request_type` | Optional | One of `"predict"` (default) or `"health"`. Forecast requests omit this field or set it to `"predict"`. | -| `schema_version` | Yes | Must be `"v3"`. | +| `schema_version` | Yes | Must be `"v4"`. | | `model_version` | Yes | Exact deployed model version expected by the caller. | | `query_mode` | Yes | `"forecast"` for the unconditional forecast or `"condition"` for a conditional query at one future position. | | `return_mode` | Yes | One of `"mean"`, `"samples"`, `"quantiles"`, or `"log_prob"`, covered by the `forecast_mean`, `forecast_samples`, `forecast_quantiles`, and `forecast_log_prob` helpers. | @@ -166,12 +166,12 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | `history_rows` | Yes | Non-empty array of history row objects in the declared schema. | | `query_times` | Yes | Non-empty future forecast horizon values. Absolute datetimes are encoded timezone-stably. | | `time_column` | For absolute datetime, optional otherwise | Name of the history time column. It must not duplicate a modeled column name. | -| `requested_columns` | Optional | Output column names or integer indices. Duplicates are rejected. Defaults to all modeled columns, minus the columns an equality condition pinned. Under `return_mode="log_prob"` an explicit list must name every readable column in schema order, because the score is joint over them. | +| `requested_columns` | Optional | Output column names or integer indices, answered in the order given. Duplicates are rejected. Any declared column may be named, whatever its `role` and whatever a condition says about it. Defaults to every declared column in declared order, so a condition response and a forecast response over the same schema line up column for column. Under `return_mode="log_prob"` an explicit list must name every declared column in declared order, because the score is joint over them. | | `query_rows` | `log_prob` mode | Observed rows to score, one per entry of `query_times`, each carrying a value for every declared column. Required by `return_mode="log_prob"` and rejected for every other mode. A row that contradicts a condition — a pinned column given another value, or a bounded column outside its range — is refused by the service instead of scored. | | `n_samples` | Samples and quantiles controls | Positive sample count when sampling controls are needed. Oversized sample forecasts are batched automatically against the cap advertised in health metadata. | | `quantiles` | Quantiles mode | Quantile levels in `(0, 1)`, required for `return_mode="quantiles"`. | | `seed` | Optional | Integer random seed for reproducible stochastic outputs. | -| `condition` | With `query_mode="condition"` | Object with `query_time_index` (index into `query_times`) and `conditions`, a list of `{"column", "kind": "equality", "value"}` or `{"column", "kind": "interval", "lower", "upper"}` entries with `null` for an open bound. At most one condition per column, at least one column left unconditioned, and no pinned column in `requested_columns`. Forbidden with any other query mode. | +| `condition` | With `query_mode="condition"` | Object with `query_time_index` (index into `query_times`) and `conditions`, a list of `{"column", "kind": "equality", "value"}` or `{"column", "kind": "interval", "lower", "upper"}` entries with `null` for an open bound. At most one condition per column and at least one column left unconditioned. Forbidden with any other query mode. | | `time_scale_seconds` | Optional | Positive scale for continuous time indexes. | | `use_local_normalized_time` | Optional | Whether the service should use local normalized time features. | | `calendar_id` | Optional | Calendar identifier, defaulting to `pandas-default`. | @@ -200,7 +200,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Description | | --- | --- | -| `schema_version` | Response schema, expected to be `"v3"`. | +| `schema_version` | Response schema, expected to be `"v4"`. | | `image_version` | Service image version that produced the response. | | `model_version` | Model version that produced the response. | | `checkpoint_version` | Checkpoint version that produced the response. | @@ -226,7 +226,7 @@ The string literals are exposed as `PREDICT_REQUEST_TYPE`, `HEALTH_REQUEST_TYPE` | Field | Description | | --- | --- | | `status` | Service status string. | -| `schema_version` | Advertised schema version. The SDK requires `v3`. | +| `schema_version` | Advertised schema version. The SDK requires `v4`. | | `image_version` | Running service image version. | | `model_version` | Running model version. | | `checkpoint_version` | Loaded checkpoint version. | diff --git a/notebooks/forecast_condition.ipynb b/notebooks/forecast_condition.ipynb index d3fc9f9..4fac89e 100644 --- a/notebooks/forecast_condition.ipynb +++ b/notebooks/forecast_condition.ipynb @@ -444,7 +444,7 @@ "## Scoring an outcome you already have\n", "Every read-out so far answers *what does the model expect*. `log_prob` asks the opposite — how plausible are values I already hold — and is the only return mode where the caller supplies the future instead of receiving it. Under a condition it scores those values against the conditional, which is how two candidate outcomes of the same what-if are compared.\n", "\n", - "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every *readable* column in schema order — a narrower projection would score a different distribution than the caller means to ask about — and the pinned column is the one exception, since conditioning fixed its value. The deployment refuses a row that contradicts the condition, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", + "Three rules follow from scoring a joint rather than a projection of one, and the client enforces all three before anything is sent. `query_rows` carries one observed row per entry of `query_times`, each with a value for every declared column. `requested_columns` must name every declared column in declared order — a narrower projection would score a different distribution than the caller means to ask about — and a condition excuses nothing, so omitting the field is the natural spelling. The deployment refuses a row that contradicts the condition, a pinned column given another value or a banded column outside its range, instead of scoring it.\n", "\n", "The scores below are log densities of whole rows, so the unit caveat from the plausibility section applies unchanged: read their difference, not their level." ] diff --git a/src/jointfm_client/client.py b/src/jointfm_client/client.py index c15b1f9..8a6a5bd 100644 --- a/src/jointfm_client/client.py +++ b/src/jointfm_client/client.py @@ -541,7 +541,8 @@ def forecast_log_prob( of ``query_times``, and the result is the model's log density of those values under its joint at that position. It therefore scores the whole joint, and ``requested_columns`` — when given at all — must name every - readable column in schema order. + declared column in declared order, which is what omitting it already + resolves to. With ``condition`` the score is taken under the conditional at the one future position the block names, and the service refuses rows that diff --git a/src/jointfm_client/contract.py b/src/jointfm_client/contract.py index 208a1d9..2b028b4 100644 --- a/src/jointfm_client/contract.py +++ b/src/jointfm_client/contract.py @@ -37,7 +37,7 @@ # pyproject.toml the way a hand-maintained literal here did. PACKAGE_VERSION: Final = importlib.metadata.version(DISTRIBUTION_NAME) -SCHEMA_VERSION: Final = "v3" +SCHEMA_VERSION: Final = "v4" # The service sums and averages its log densities in the model's own precision, # so its summaries differ from a recomputation in the last few digits. DERIVED_SCORE_TOLERANCE: Final = 1e-6 @@ -317,8 +317,9 @@ class EqualityCondition: An equality condition is an event of probability zero under a continuous head, so the deployment answers it analytically rather than by filtering - draws. It fixes the column's value, which is why that column leaves the - read-out projection. + draws. It fixes the column's value, so a projection that names the column + reads that value back rather than a model output — which is what lets a + scenario answer line up column for column with an unconditioned one. """ column: str @@ -423,7 +424,7 @@ def __post_init__(self) -> None: @property def pinned_columns(self) -> tuple[str, ...]: - """Columns an equality condition fixes, which leave the read-out set.""" + """Columns an equality condition fixes, whose answer is the request's own value.""" return tuple( condition.column for condition in self.conditions @@ -581,29 +582,17 @@ def _validate_condition(self) -> None: "reads out is the conditional distribution of the columns it does " "not condition" ) - requested = _resolve_requested_columns( - self.schema.columns, self.requested_columns - ) - if requested is None: - return - requested_names = {value for value in requested if isinstance(value, str)} - pinned = sorted(set(block.pinned_columns) & requested_names) - if pinned: - raise ValueError( - f"requested_columns lists pinned columns {pinned}; an equality " - "condition fixes the value, so reading it back returns only what " - "the request supplied" - ) def _validate_query_rows(self) -> None: """Check the observed rows a log-density request scores. The service scores the whole joint at each future position, so a row must carry every declared column and ``requested_columns`` must name - every readable one in schema order: a narrower projection would score a - different distribution than the caller believes it asked about. An - equality condition is the one exception, since it fixes its column's - value and the conditioning contract forbids reading it back. + every one of them in declared order: a narrower or reordered projection + would score a different distribution than the caller believes it asked + about. A ``condition`` block changes nothing here — its columns are + declared columns like any other — and omitting the field already + resolves to exactly this projection. Only two per-value failures are decidable without the deployment's own encoding and are checked here — a missing value and a non-finite number. @@ -643,16 +632,11 @@ def _validate_query_rows(self) -> None: ) if requested is None: return - pinned = set(() if self.condition is None else self.condition.pinned_columns) - readable = [name for name in declared if name not in pinned] - if requested != readable: - excused = ( - f", the pinned columns {sorted(pinned)} excepted" if pinned else "" - ) + if requested != declared: raise ValueError( "return_mode='log_prob' scores the whole joint, so " - "requested_columns must list every declared column in schema " - f"order{excused}: expected {readable}, got {requested}" + "requested_columns must list every declared column in declared " + f"order: expected {declared}, got {requested}" ) def to_payload(self) -> dict[str, Any]: diff --git a/tests/fixtures/condition_interval_response.json b/tests/fixtures/condition_interval_response.json index e9ab16e..f217332 100644 --- a/tests/fixtures/condition_interval_response.json +++ b/tests/fixtures/condition_interval_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/condition_log_prob_request.json b/tests/fixtures/condition_log_prob_request.json index 2407223..b6ef705 100644 --- a/tests/fixtures/condition_log_prob_request.json +++ b/tests/fixtures/condition_log_prob_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "condition", "return_mode": "log_prob", @@ -26,7 +26,7 @@ } ], "query_times": [2, 3], - "requested_columns": ["target"], + "requested_columns": ["driver", "target"], "condition": { "query_time_index": 1, "conditions": [ diff --git a/tests/fixtures/condition_log_prob_response.json b/tests/fixtures/condition_log_prob_response.json index 21d9f62..04513ea 100644 --- a/tests/fixtures/condition_log_prob_response.json +++ b/tests/fixtures/condition_log_prob_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -8,7 +8,7 @@ "return_mode": "log_prob", "outputs": { "query_times": [3], - "requested_columns": ["target"], + "requested_columns": ["driver", "target"], "mean": null, "samples": null, "quantiles": null, diff --git a/tests/fixtures/condition_mean_request.json b/tests/fixtures/condition_mean_request.json index ba7d9f4..f506b20 100644 --- a/tests/fixtures/condition_mean_request.json +++ b/tests/fixtures/condition_mean_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "condition", "return_mode": "mean", diff --git a/tests/fixtures/condition_mean_response.json b/tests/fixtures/condition_mean_response.json index 622745b..0c06ed6 100644 --- a/tests/fixtures/condition_mean_response.json +++ b/tests/fixtures/condition_mean_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_log_prob_request.json b/tests/fixtures/forecast_log_prob_request.json index f843ce4..6ba470f 100644 --- a/tests/fixtures/forecast_log_prob_request.json +++ b/tests/fixtures/forecast_log_prob_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "log_prob", diff --git a/tests/fixtures/forecast_log_prob_response.json b/tests/fixtures/forecast_log_prob_response.json index f6ad94a..ab9ee00 100644 --- a/tests/fixtures/forecast_log_prob_response.json +++ b/tests/fixtures/forecast_log_prob_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_mean_request.json b/tests/fixtures/forecast_mean_request.json index c4b7ab7..17c29ff 100644 --- a/tests/fixtures/forecast_mean_request.json +++ b/tests/fixtures/forecast_mean_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "mean", diff --git a/tests/fixtures/forecast_mean_response.json b/tests/fixtures/forecast_mean_response.json index d6aa094..d289237 100644 --- a/tests/fixtures/forecast_mean_response.json +++ b/tests/fixtures/forecast_mean_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_quantiles_request.json b/tests/fixtures/forecast_quantiles_request.json index 602ddc4..32442ba 100644 --- a/tests/fixtures/forecast_quantiles_request.json +++ b/tests/fixtures/forecast_quantiles_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "quantiles", diff --git a/tests/fixtures/forecast_quantiles_response.json b/tests/fixtures/forecast_quantiles_response.json index aa8e9bd..cd49f12 100644 --- a/tests/fixtures/forecast_quantiles_response.json +++ b/tests/fixtures/forecast_quantiles_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/forecast_samples_request.json b/tests/fixtures/forecast_samples_request.json index 0f161c2..9203269 100644 --- a/tests/fixtures/forecast_samples_request.json +++ b/tests/fixtures/forecast_samples_request.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "samples", diff --git a/tests/fixtures/forecast_samples_response.json b/tests/fixtures/forecast_samples_response.json index 60e60e1..29cff5c 100644 --- a/tests/fixtures/forecast_samples_response.json +++ b/tests/fixtures/forecast_samples_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/health_metadata.json b/tests/fixtures/health_metadata.json index 08d9aeb..804fdce 100644 --- a/tests/fixtures/health_metadata.json +++ b/tests/fixtures/health_metadata.json @@ -1,6 +1,6 @@ { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", diff --git a/tests/fixtures/input_size_exceeded_response.json b/tests/fixtures/input_size_exceeded_response.json index e244f5c..2e1c7c1 100644 --- a/tests/fixtures/input_size_exceeded_response.json +++ b/tests/fixtures/input_size_exceeded_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "INPUT_SIZE_EXCEEDED", diff --git a/tests/fixtures/model_version_mismatch_response.json b/tests/fixtures/model_version_mismatch_response.json index b5fa39e..07f5363 100644 --- a/tests/fixtures/model_version_mismatch_response.json +++ b/tests/fixtures/model_version_mismatch_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "MODEL_VERSION_MISMATCH", diff --git a/tests/fixtures/schema_version_mismatch_response.json b/tests/fixtures/schema_version_mismatch_response.json index ffb299b..d5cffa9 100644 --- a/tests/fixtures/schema_version_mismatch_response.json +++ b/tests/fixtures/schema_version_mismatch_response.json @@ -1,9 +1,9 @@ { - "schema_version": "v2", + "schema_version": "v3", "errors": [ { "code": "SCHEMA_VERSION_MISMATCH", - "message": "Unsupported schema_version: expected 'v3', got 'v2'", + "message": "Unsupported schema_version: expected 'v4', got 'v3'", "field": "schema_version", "detail": {} } diff --git a/tests/fixtures/validation_error_response.json b/tests/fixtures/validation_error_response.json index 7f12104..30d6160 100644 --- a/tests/fixtures/validation_error_response.json +++ b/tests/fixtures/validation_error_response.json @@ -1,5 +1,5 @@ { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", diff --git a/tests/test_cli.py b/tests/test_cli.py index 440447a..006392a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -48,7 +48,7 @@ class FakeHealthClient: "deployment-id/predictionsUnstructured" ), deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -58,7 +58,7 @@ def health(self, *, cache: bool = False, refresh: bool = False) -> HealthMetadat del cache, refresh return HealthMetadata( status="ok", - schema_version="v3", + schema_version="v4", image_version="0.3.0", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", checkpoint_version="sdk-test", @@ -172,7 +172,7 @@ def test_predict_command_writes_response_file(monkeypatch, tmp_path: Path) -> No request_file.write_text( json.dumps( { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } ), diff --git a/tests/test_condition_mode.py b/tests/test_condition_mode.py index 31d7deb..7a1e25c 100644 --- a/tests/test_condition_mode.py +++ b/tests/test_condition_mode.py @@ -97,7 +97,7 @@ def _health( """Build one advertisement without going through a transport.""" return HealthMetadata( status="ok", - schema_version="v3", + schema_version="v4", image_version="0.3.0", model_version=_MODEL_VERSION, checkpoint_version="sdk-test", @@ -195,14 +195,6 @@ def test_the_mode_and_the_block_must_agree() -> None: ("target",), "outside the 2 requested future positions", ), - ( - ConditionBlock( - query_time_index=0, - conditions=(EqualityCondition(column="driver", value=1.0),), - ), - ("driver", "target"), - "pinned columns", - ), ( ConditionBlock( query_time_index=0, @@ -249,6 +241,38 @@ def test_an_interval_column_may_still_be_read_out() -> None: assert request.to_payload()["requested_columns"] == ["driver", "target"] +def test_a_pinned_column_may_be_read_out_in_any_position() -> None: + """A projection is condition-agnostic, so a pin is nameable like any column. + + Its answer is the request's own value, which is what lets a scenario frame + be compared against an unconditioned one column for column. + """ + block = ConditionBlock( + query_time_index=0, + conditions=(EqualityCondition(column="driver", value=1.0),), + ) + + request = _request(block, requested_columns=("target", "driver")) + + assert request.to_payload()["requested_columns"] == ["target", "driver"] + + +def test_omitting_the_projection_states_nothing_on_the_wire() -> None: + """The default lives in the service, so the client must not invent one. + + A client-side default would be a second copy of the rule, free to drift from + the one the deployment actually applies. + """ + block = ConditionBlock( + query_time_index=0, + conditions=(EqualityCondition(column="driver", value=1.0),), + ) + + payload = _request(block, requested_columns=None).to_payload() + + assert "requested_columns" not in payload + + def test_the_payload_carries_the_block_the_service_parses() -> None: """The wire form names its position once and each condition names its kind.""" payload = build_forecast_payload( @@ -394,7 +418,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: assert isinstance(sample_count, int) start = sum(cast(int, earlier["n_samples"]) for earlier in self.payloads[:-1]) return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", @@ -432,7 +456,7 @@ def _health_payload( """Build one health advertisement as the service serializes it.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", diff --git a/tests/test_configuration.py b/tests/test_configuration.py index 5ee8cab..3a6fe7b 100644 --- a/tests/test_configuration.py +++ b/tests/test_configuration.py @@ -47,7 +47,7 @@ class _HealthTransport: _METADATA: dict[str, object] = { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", "checkpoint_version": "yaml", @@ -147,7 +147,7 @@ def test_load_settings_layers_config_below_dotenv_and_environment( "datarobot_endpoint": "https://app.datarobot.com/api/v2", "datarobot_api_token": "yaml-token", "deployment_id": "yaml-deployment-id", - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", } }, @@ -191,7 +191,7 @@ def test_client_from_env_uses_transport_defaults_from_config( "datarobot_endpoint": "https://app.datarobot.com/api/v2", "datarobot_api_token": "yaml-token", "deployment_id": "yaml-deployment-id", - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.yaml", }, "transport": { diff --git a/tests/test_contract.py b/tests/test_contract.py index 90af77b..9d80ad3 100644 --- a/tests/test_contract.py +++ b/tests/test_contract.py @@ -76,7 +76,7 @@ def _health_metadata() -> dict[str, object]: """Health metadata.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -111,7 +111,7 @@ def test_package_identity_contract() -> None: assert DISTRIBUTION_NAME == "jointfm-client" assert IMPORT_NAMESPACE == "jointfm_client" assert FIRST_SUPPORTED_PYTHON_VERSION == "3.11" - assert SCHEMA_VERSION == "v3" + assert SCHEMA_VERSION == "v4" def test_package_version_matches_installed_distribution() -> None: @@ -267,7 +267,7 @@ def test_forecast_payload_matches_service_contract_without_mutating_inputs() -> ) assert payload == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -362,7 +362,7 @@ def test_dataframe_payload_matches_service_forecast_request_shape() -> None: ) assert payload == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "mean", @@ -777,7 +777,7 @@ def test_health_and_response_models_parse_current_payloads() -> None: health = HealthMetadata.from_payload(_health_metadata()) response = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -840,7 +840,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - """Forecast result conversion helpers cover mean samples and quantiles.""" mean_result = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -860,7 +860,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - ) sample_result = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -883,7 +883,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - ) quantile_result = ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -957,7 +957,7 @@ def test_forecast_result_conversion_helpers_cover_mean_samples_and_quantiles() - def test_forecast_response_validates_request_scoped_shapes() -> None: """Forecast response validates request scoped shapes.""" request_payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "samples", @@ -966,7 +966,7 @@ def test_forecast_response_validates_request_scoped_shapes() -> None: "n_samples": 3, } response_payload = { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -996,7 +996,7 @@ def test_forecast_response_raises_typed_error_for_success_payload_errors() -> No with pytest.raises(JointFMServiceError) as exc_info: ForecastResponse.from_payload( { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -1162,7 +1162,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: def _forecast_response_payload(*, return_mode: str) -> dict[str, object]: """Forecast response payload.""" return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", diff --git a/tests/test_contract_models.py b/tests/test_contract_models.py index 9d5186b..af48ec8 100644 --- a/tests/test_contract_models.py +++ b/tests/test_contract_models.py @@ -76,7 +76,7 @@ def test_request_models_serialize_direct_payloads_without_mutating_inputs() -> N ) assert metadata.to_payload() == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -91,7 +91,7 @@ def test_request_models_serialize_direct_payloads_without_mutating_inputs() -> N "timezone": "UTC", } assert request.to_payload() == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "quantiles", @@ -303,7 +303,7 @@ def test_response_models_reject_direct_validation_edges() -> None: def test_forecast_response_rejects_request_scoped_metadata_mismatches() -> None: """Forecast response rejects request scoped metadata mismatches.""" request_payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "mean", @@ -354,7 +354,7 @@ def test_forecast_response_rejects_request_scoped_metadata_mismatches() -> None: def test_forecast_response_rejects_sample_bound_violations() -> None: """Forecast response rejects sample bound violations.""" request_payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "query_mode": "forecast", "return_mode": "samples", @@ -397,7 +397,7 @@ def test_forecast_response_rejects_sample_bound_violations() -> None: def _mean_response_payload() -> dict[str, Any]: """Mean response payload.""" return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", @@ -419,7 +419,7 @@ def _mean_response_payload() -> dict[str, Any]: def _sample_response_payload() -> dict[str, Any]: """Sample response payload.""" return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.smoke-1", "checkpoint_version": "smoke-1", diff --git a/tests/test_feature_importance.py b/tests/test_feature_importance.py index 3d6fe0a..c706175 100644 --- a/tests/test_feature_importance.py +++ b/tests/test_feature_importance.py @@ -47,7 +47,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: self.payloads.append(dict(payload)) samples = self.sample_batches[len(self.payloads) - 1] return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": _MODEL_VERSION, "checkpoint_version": "sdk-test", diff --git a/tests/test_log_prob_mode.py b/tests/test_log_prob_mode.py index cdb6e01..9c0cd34 100644 --- a/tests/test_log_prob_mode.py +++ b/tests/test_log_prob_mode.py @@ -20,9 +20,10 @@ rows nothing else carries, and a projection narrower than the scored joint asks a different question than the caller believes they asked. The response carries a block with no column axis whose summary fields only repeat what ``values`` -already says. Under a condition the two meet: a pinned column leaves the -read-out set but still has to be supplied, because the service refuses a row -contradicting the pin rather than scoring it. +already says. Under a condition the two meet: a pinned column is projected and +supplied like any other, because the service refuses a row contradicting the +pin rather than scoring it, while the score itself stays the conditional +density of the columns the request does not condition. """ from __future__ import annotations @@ -170,18 +171,22 @@ def test_a_row_that_cannot_be_scored_is_refused_before_any_round_trip( def test_a_partial_projection_is_refused_because_it_scores_something_else() -> None: """Dropping a column would score a narrower joint than the caller believes.""" - with pytest.raises(ValueError, match="every declared column in schema order"): + with pytest.raises(ValueError, match="every declared column in declared order"): _request(requested_columns=("target",)) -def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: - """Conditioning fixes the column, and the contract forbids reading it back.""" - pinned = _request( - requested_columns=("target",), +def test_a_score_may_leave_out_no_column_at_all_under_a_condition() -> None: + """A condition excuses nothing: the scorer covers the whole declared joint. + + Omitting the projection is the natural spelling, because the service's own + default already resolves to every declared column in declared order. + """ + stated = _request( + requested_columns=("driver", "target"), query_rows=_PINNED_QUERY_ROWS, condition=_PIN_DRIVER, ) - assert pinned.to_payload()["requested_columns"] == ["target"] + assert stated.to_payload()["requested_columns"] == ["driver", "target"] default_projection = _request( requested_columns=None, @@ -190,8 +195,12 @@ def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: ) assert "requested_columns" not in default_projection.to_payload() - with pytest.raises(ValueError, match="lists pinned columns"): - _request(query_rows=_PINNED_QUERY_ROWS, condition=_PIN_DRIVER) + with pytest.raises(ValueError, match="every declared column in declared order"): + _request( + requested_columns=("target",), + query_rows=_PINNED_QUERY_ROWS, + condition=_PIN_DRIVER, + ) @pytest.mark.parametrize( @@ -201,7 +210,7 @@ def test_a_pinned_column_is_the_one_column_a_score_may_leave_out() -> None: ( "condition_log_prob_request", _PIN_DRIVER, - ("target",), + ("driver", "target"), _PINNED_QUERY_ROWS, ), ], @@ -329,7 +338,7 @@ def test_the_client_scores_under_a_condition( schema=_schema(), query_times=[2, 3], query_rows=list(_PINNED_QUERY_ROWS), - requested_columns=["target"], + requested_columns=["driver", "target"], model_version=_MODEL_VERSION, seed=7, condition=_PIN_DRIVER, diff --git a/tests/test_pool.py b/tests/test_pool.py index 837f19e..9de17bb 100644 --- a/tests/test_pool.py +++ b/tests/test_pool.py @@ -56,7 +56,7 @@ def _health( """Health.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": model_version, "checkpoint_version": checkpoint_version, @@ -134,7 +134,7 @@ def test_pool_retries_next_instance_on_470() -> None: pool = _pool(transport=_Transport(fail_ids=frozenset({"a"}))) assert pool.next_instance().deployment_id == "a" assert pool.next_instance().deployment_id == "b" - assert pool.post_json({"schema_version": "v3"}) == { + assert pool.post_json({"schema_version": "v4"}) == { "ok": True, "deployment_id": "b", } @@ -144,7 +144,7 @@ def test_pool_raises_when_all_instances_unavailable() -> None: """Pool raises when all instances unavailable.""" pool = _pool(transport=_Transport(fail_ids=frozenset({"a", "b"}))) with pytest.raises(JointFMHTTPStatusError, match="unavailable"): - pool.post_json({"schema_version": "v3"}) + pool.post_json({"schema_version": "v4"}) def test_pool_health_rejects_mismatch_and_aligns_sample_cap() -> None: @@ -275,7 +275,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: executor.submit( pool.post_json_to, pool.instance_at(index), - {"schema_version": "v3"}, + {"schema_version": "v4"}, ) for index in range(2) ] @@ -299,7 +299,7 @@ def test_pool_health_routes_only_reachable_peers() -> None: assert pool.instance_at(0).deployment_id == "b" assert pool.instance_at(1).deployment_id == "b" assert pool.next_instance().deployment_id == "b" - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "b" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "b" def test_pool_failover_retries_health_excluded_peer() -> None: @@ -313,7 +313,7 @@ def test_pool_failover_retries_health_excluded_peer() -> None: assert pool.instance_at(0).deployment_id == "b" transport.fail_ids = frozenset({"b"}) - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "a" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "a" assert pool.instance_at(0).deployment_id == "a" @@ -333,7 +333,7 @@ def test_pool_health_skips_incompatible_peer_when_another_matches_pin() -> None: assert metadata.model_version == pinned assert pool.instance_at(0).deployment_id == "a" assert pool.instance_at(1).deployment_id == "a" - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "a" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "a" def test_pool_cooldown_restores_peer_after_transient_failure( @@ -345,7 +345,7 @@ def test_pool_cooldown_restores_peer_after_transient_failure( transport = _Transport(fail_ids=frozenset({"a"})) pool = _pool(transport=transport, peer_cooldown_seconds=10.0) - assert pool.post_json({"schema_version": "v3"})["deployment_id"] == "b" + assert pool.post_json({"schema_version": "v4"})["deployment_id"] == "b" assert pool.instance_at(0).deployment_id == "b" assert pool.instance_at(1).deployment_id == "b" @@ -355,4 +355,4 @@ def test_pool_cooldown_restores_peer_after_transient_failure( clock["now"] = 110.0 assert {pool.instance_at(i).deployment_id for i in range(2)} == {"a", "b"} - assert pool.post_json({"schema_version": "v3"})["deployment_id"] in {"a", "b"} + assert pool.post_json({"schema_version": "v4"})["deployment_id"] in {"a", "b"} diff --git a/tests/test_settings.py b/tests/test_settings.py index 22c74f3..13632c0 100644 --- a/tests/test_settings.py +++ b/tests/test_settings.py @@ -54,7 +54,7 @@ def _hosted_env(**overrides: str) -> dict[str, str]: DATAROBOT_ENDPOINT_ENV: "https://app.datarobot.com/api/v2/", DATAROBOT_API_TOKEN_ENV: "secret-token", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", } env.update(overrides) @@ -66,7 +66,7 @@ def test_load_settings_from_environment_with_deployment_id_builds_hosted_url() - settings = load_settings(env=_hosted_env(), dotenv_path=None) assert settings.datarobot_endpoint == "https://app.datarobot.com/api/v2" - assert settings.schema_version == "v3" + assert settings.schema_version == "v4" assert settings.model_version == "jointfm-inference:0.3.0+ckpt.sdk-test" assert settings.deployment_id == "deployment-id" assert settings.predict_url == ( @@ -82,7 +82,7 @@ def test_load_settings_with_local_service_base_url_builds_direct_urls() -> None: settings = load_settings( env={ JOINTFM_LOCAL_BASE_URL_ENV: "http://127.0.0.1:8080/", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.local-test", }, dotenv_path=None, @@ -158,7 +158,7 @@ def test_load_settings_reads_dotenv_without_overriding_environment(tmp_path) -> "DATAROBOT_ENDPOINT=https://app.datarobot.com/api/v2", "DATAROBOT_API_TOKEN=file-token", "JOINTFM_DEPLOYMENT_ID=file-deployment-id", - "JOINTFM_SCHEMA_VERSION=v3", + "JOINTFM_SCHEMA_VERSION=v4", "JOINTFM_MODEL_VERSION=jointfm-inference:0.3.0+ckpt.sdk-test", ] ), @@ -209,7 +209,7 @@ def test_load_settings_rejects_missing_credentials_without_defaults() -> None: env={ DATAROBOT_API_TOKEN_ENV: "secret-token", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, @@ -220,7 +220,7 @@ def test_load_settings_rejects_missing_credentials_without_defaults() -> None: env={ DATAROBOT_ENDPOINT_ENV: "https://app.datarobot.com/api/v2", JOINTFM_DEPLOYMENT_ID_ENV: "deployment-id", - JOINTFM_SCHEMA_VERSION_ENV: "v3", + JOINTFM_SCHEMA_VERSION_ENV: "v4", JOINTFM_MODEL_VERSION_ENV: "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, diff --git a/tests/test_surfaces.py b/tests/test_surfaces.py index 1d73b3d..bfd2a26 100644 --- a/tests/test_surfaces.py +++ b/tests/test_surfaces.py @@ -70,7 +70,7 @@ def test_hosted_surface_uses_datarobot_routes_and_auth_headers( health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version=request_payload["model_version"], deployment_id="deployment-id", ) @@ -209,7 +209,7 @@ def test_hosted_surface_auto_discovers_model_version_when_settings_unpinned( health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", deployment_id="deployment-id", ) assert settings.model_version is None diff --git a/tests/test_transport.py b/tests/test_transport.py index 441d7aa..c5fd1af 100644 --- a/tests/test_transport.py +++ b/tests/test_transport.py @@ -154,7 +154,7 @@ def _health_payload( """Health payload.""" return { "status": "ok", - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": model_version, "checkpoint_version": checkpoint_version, @@ -190,7 +190,7 @@ def _forecast_response_payload(*, return_mode: str = "mean") -> dict[str, object else None, } return { - "schema_version": "v3", + "schema_version": "v4", "image_version": "0.3.0", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "checkpoint_version": "sdk-test", @@ -234,7 +234,7 @@ def test_transport_posts_json_with_headers_timeout_and_user_agent() -> None: session.mount("https://", adapter) result = transport.post_json( - "https://example.com/predict", {"schema_version": "v3"} + "https://example.com/predict", {"schema_version": "v4"} ) assert result == {"ok": True} @@ -248,7 +248,7 @@ def test_transport_posts_json_with_headers_timeout_and_user_agent() -> None: assert adapter.kwargs[0]["timeout"] == (1.5, 2.5) request_body = request.body assert isinstance(request_body, bytes) - assert json.loads(request_body.decode("utf-8")) == {"schema_version": "v3"} + assert json.loads(request_body.decode("utf-8")) == {"schema_version": "v4"} def test_transport_from_settings_attaches_hosted_auth_headers_and_closes_session() -> ( @@ -265,7 +265,7 @@ def test_transport_from_settings_attaches_hosted_auth_headers_and_closes_session "deployment-id/predictionsUnstructured" ), deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -293,7 +293,7 @@ def test_transport_from_local_settings_omits_hosted_auth_headers() -> None: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.local-test", local_base_url="http://127.0.0.1:8080", ) @@ -319,7 +319,7 @@ def test_transport_retries_retryable_server_responses() -> None: retry_config=JointFMRetryConfig(max_attempts=2) ) - result = transport.post_json(_server_url(server), {"schema_version": "v3"}) + result = transport.post_json(_server_url(server), {"schema_version": "v4"}) assert result == {"ok": True} assert handler.request_count == 2 @@ -351,7 +351,7 @@ def test_transport_retries_html_bodied_gateway_errors() -> None: ), ) - result = transport.post_json(_server_url(server), {"schema_version": "v3"}) + result = transport.post_json(_server_url(server), {"schema_version": "v4"}) assert result == {"ok": True} assert handler.request_count == 2 @@ -380,7 +380,7 @@ def test_transport_raises_status_error_when_gateway_html_persists() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v3"}) + transport.post_json(_server_url(server), {"schema_version": "v4"}) assert exc_info.value.status_code == HTTPStatus.BAD_GATEWAY assert "502 Bad Gateway" in exc_info.value.response_body_excerpt @@ -444,7 +444,7 @@ def request(self, *args: Any, **kwargs: Any) -> requests.Response: ) result = transport.post_json( - "https://example.com/predict", {"schema_version": "v3"} + "https://example.com/predict", {"schema_version": "v4"} ) assert result == {"ok": True} @@ -463,7 +463,7 @@ def test_status_error_carries_parsed_retry_after_seconds() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v3"}) + transport.post_json(_server_url(server), {"schema_version": "v4"}) assert exc_info.value.retry_after_seconds == 0.5 assert handler.request_count == 1 @@ -481,7 +481,7 @@ def test_transport_does_not_retry_validation_errors() -> None: ) with pytest.raises(JointFMHTTPStatusError) as exc_info: - transport.post_json(_server_url(server), {"schema_version": "v3"}) + transport.post_json(_server_url(server), {"schema_version": "v4"}) assert exc_info.value.status_code == HTTPStatus.BAD_REQUEST assert exc_info.value.datarobot_request_id == "request-id-1" @@ -507,7 +507,7 @@ def test_transport_rejects_non_json_serializable_payloads() -> None: with pytest.raises(JointFMRequestEncodingError, match="JSON-serializable"): transport.post_json( "https://example.com/predict", - {"schema_version": "v3", "bad": object()}, + {"schema_version": "v4", "bad": object()}, ) @@ -572,14 +572,14 @@ def test_client_predict_uses_configured_transport_and_settings() -> None: health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/healthz", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) transport = RecordingTransport() client = JointFMClient(settings=settings, transport=transport) payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } @@ -598,7 +598,7 @@ def test_client_health_returns_typed_metadata_and_caches_only_when_requested() - health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -629,7 +629,7 @@ def test_client_health_instances_returns_one_entry_for_single_endpoint() -> None health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -691,7 +691,7 @@ def capture_transport( "DATAROBOT_ENDPOINT": "https://app.datarobot.com/api/v2", "DATAROBOT_API_TOKEN": "secret-token", "JOINTFM_DEPLOYMENT_ID": "deployment-id", - "JOINTFM_SCHEMA_VERSION": "v3", + "JOINTFM_SCHEMA_VERSION": "v4", "JOINTFM_MODEL_VERSION": "jointfm-inference:0.3.0+ckpt.sdk-test", }, dotenv_path=None, @@ -713,7 +713,7 @@ def test_client_health_rejects_cached_model_mismatch() -> None: health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -739,7 +739,7 @@ def test_client_hosted_health_posts_request_type_health_to_predict_url() -> None health_url=predict_url, predict_url=predict_url, deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) @@ -763,7 +763,7 @@ def test_client_local_health_keeps_get_request_to_healthz_route() -> None: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", local_base_url="http://127.0.0.1:8080", ) @@ -805,7 +805,7 @@ def test_client_forecast_builds_payload_from_rows_and_returns_typed_response() - assert result.outputs.mean == ((12.0,),) assert transport.predict_url == "http://localhost:8080/predict" assert transport.payload == { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", "query_mode": "forecast", "return_mode": "mean", @@ -980,7 +980,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: health_url="http://127.0.0.1:8080/healthz", predict_url="http://127.0.0.1:8080/predict", deployment_selector="local_service", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", local_base_url="http://127.0.0.1:8080", ) @@ -1032,13 +1032,13 @@ def test_client_predict_raises_typed_service_error_for_success_payload_errors() health_url="https://app.datarobot.com/api/v2/deployments/deployment-id/healthz", predict_url="https://app.datarobot.com/api/v2/deployments/deployment-id/predictionsUnstructured", deployment_selector="deployment_id", - schema_version="v3", + schema_version="v4", model_version="jointfm-inference:0.3.0+ckpt.sdk-test", deployment_id="deployment-id", ) transport = RecordingTransport() transport.predict_payload = { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", @@ -1052,7 +1052,7 @@ def test_client_predict_raises_typed_service_error_for_success_payload_errors() with pytest.raises(JointFMServiceError) as exc_info: client.predict( { - "schema_version": "v3", + "schema_version": "v4", "model_version": settings.model_version, } ) @@ -1141,7 +1141,7 @@ def do_POST(self) -> None: payload = {"ok": True} else: payload = { - "schema_version": "v3", + "schema_version": "v4", "errors": [ { "code": "VALIDATION_ERROR", @@ -1225,7 +1225,7 @@ def _pool_settings(primary: str, backup: str) -> JointFMSettings: health_url=primary, predict_url=primary, deployment_selector="deployment_ids", - schema_version="v3", + schema_version="v4", instances=( JointFMInstanceSettings(deployment_id="primary-id", predict_url=primary), JointFMInstanceSettings(deployment_id="backup-id", predict_url=backup), @@ -1267,7 +1267,7 @@ def post_json(self, url: str, payload: Mapping[str, Any]) -> Mapping[str, Any]: transport = PoolTransport() client = JointFMClient(settings=settings, transport=transport) payload = { - "schema_version": "v3", + "schema_version": "v4", "model_version": "jointfm-inference:0.3.0+ckpt.sdk-test", } client.predict(payload)