diff --git a/docs/STABILITY.md b/docs/STABILITY.md index a601e33..4b3df24 100644 --- a/docs/STABILITY.md +++ b/docs/STABILITY.md @@ -27,6 +27,7 @@ This page provides module-level guidance and deprecation notes. | `tributo.config` — `AlgorithmExecutionConfig` and nested algorithm execution models | `alpha` | Strict JSON envelope shared by local Ray and Kubernetes-hosted Ray execution | | `tributo.job` — `TributoClient` | `stable` | Primary Ray Jobs client | | `tributo.job` — `RayJob` | `stable` annotation with runtime deprecation warning | Use `TributoClient`; the annotation and warning conflict is documented without changing the public contract in this documentation update | +| `tributo.ray_jobs` | `alpha` | Workload-neutral submission identity, ambiguous-submit reconciliation, status, logs, and stop helpers | | `tributo.exceptions` — core exceptions | `stable` | ``TributoError`` and 16 common subtypes | | `tributo.exceptions` — `ResultMaterializationError` | `alpha` | Credential-safe lazy inference action failure | | `tributo.exceptions` — Bundle/Plugin exceptions | `beta` | ``BundleExportError``, ``BundleCommitBusyError``, ``AliasConflict``, ``UnsupportedArtifactFormat``, ``PostPublishCallbackError``, ``PluginLoadIssue`` | @@ -204,6 +205,8 @@ from the legacy setup-only propagation rule. | `tributo.integrations.model_importers.*` | `alpha` | Canonical ModelImporter protocol/registry plus explicit MLflow and typed artifact-to-Bundle implementations | | `tributo.integrations.sinks.parquet` | `alpha` | Parquet inference ResultSink adapter | | `tributo.integrations.sinks.lance` | `alpha` | Generic Lance inference ResultSink adapter | +| `tributo.integrations.broker` | `alpha` | Minimal transport-neutral Broker API v1; transport implementations and consume loops are external | +| `tributo.integrations.broker_registry` | `alpha` | Lazy broker discovery and explicit provider resolution | ### Inference (tributo.inference.*) diff --git a/docs/adr/002-broker-plugin-boundary.md b/docs/adr/002-broker-plugin-boundary.md new file mode 100644 index 0000000..f338d33 --- /dev/null +++ b/docs/adr/002-broker-plugin-boundary.md @@ -0,0 +1,66 @@ +# Broker provider boundary + +## Status + +Accepted for the Alpha API. + +## Context + +Tributo needs optional message-broker integrations without becoming a message +queue platform. A broker task is a control-plane admission request; bounded +data ingestion, unbounded inference streams, training, inference, Bundle +publication, and result sinks retain their existing Tributo contracts. + +Transport clients and external wire protocols must remain independently +installable. Core must be usable and testable without Redis, Kafka, RabbitMQ, +or another provider dependency. + +## Decision + +Core owns a deliberately small Broker API v1 with Alpha stability: + +- `BrokerPlugin` discovery, structural version checks, capability metadata, + stability metadata, explicit config validation, and runtime construction; +- opaque `Message` payloads with a delivery token and restricted string + metadata; +- `TaskConsumer`, `BrokerRuntime`, `TaskDisposition`, and a minimal + `TaskOutcome` with an optional credential-safe `BrokerError`; +- workload-neutral `RayJobSubmission` identity and deterministic submission + IDs derived from an operation namespace, `run_id`, and `attempt_id`; +- ambiguous Ray submission reconciliation plus status and stop operations + keyed by `submission_id`. + +`BROKER_API_VERSION = 1` checks structural compatibility; it does not imply a +Beta or long-term compatibility promise. Discovery is lazy and fail-open with +diagnostics. Explicit resolution and configuration validation fail closed. +Discovery never instantiates a provider or performs connectivity checks. + +The following concerns belong to provider packages: + +- broker connections, polling, acknowledgments, re-delivery, recovery, dead + letters, cancellation watchers, and the production consume CLI/runtime; +- external request and event schemas, operation mapping, capability profiles, + credential references, error mapping, redaction, and event durability; +- structured terminal-event publication from existing `TrainingResult`, + `InferenceResult`, Bundle, and result-sink receipts. + +Core does not define a workload registry or an external operation schema. +Providers submit one thin execution-driver Ray Job through the generic helper; +that driver calls existing in-process training or batch-inference APIs. Worker +side broker cancellation, arbitrary execution context, a generic Core consume +loop, and durable workflow semantics are outside the Alpha contract. + +`submission_id` is the primary Ray Jobs identity for admission, status, logs, +and stop. `ray_job_id` is optional execution metadata and is populated only +from a real Ray `JobDetails.job_id`; Core never substitutes `submission_id` for +it. An optional credential-free `request_digest` may be recorded as Ray +metadata, but Core does not persist it or promise cross-restart conflict +detection. + +## Consequences + +Normal Tributo installations remain free of broker dependencies, and a +provider can evolve transport and protocol behavior independently. Providers +must own their infrastructure and healthy-path tests. The first release does +not promise exactly-once execution, durable terminal events, high availability, +or complete pending-message recovery. diff --git a/docs/architecture/index.md b/docs/architecture/index.md index 6c5928e..3432271 100644 --- a/docs/architecture/index.md +++ b/docs/architecture/index.md @@ -18,4 +18,5 @@ benchmark-protocol version-policy decision-log ../adr/001-data-and-bundle-contracts +../adr/002-broker-plugin-boundary ``` diff --git a/docs/reference/api/algorithms-training.md b/docs/reference/api/algorithms-training.md index 74b53ef..157b8cd 100644 --- a/docs/reference/api/algorithms-training.md +++ b/docs/reference/api/algorithms-training.md @@ -589,6 +589,9 @@ documentation for every public stability tier. ```{autofunction} tributo.training.job_submitter.submit_training_job ``` +```{autofunction} tributo.training.job_submitter.submit_training_job_with_identity +``` + ```{autofunction} tributo.training.job_submitter.submit_training_job_with_retry ``` @@ -708,3 +711,6 @@ documentation for every public stability tier. ```{autofunction} tributo.training.xgboost_trainer.run_training_from_json ``` + +```{autofunction} tributo.training.xgboost_trainer.run_training_result_with_config +``` diff --git a/docs/reference/api/core.md b/docs/reference/api/core.md index e731bb0..80c98d2 100644 --- a/docs/reference/api/core.md +++ b/docs/reference/api/core.md @@ -26,8 +26,7 @@ documentation for every public stability tier. ```{autoexception} tributo._common.dependencies.DependencyUnavailableError ``` -```{autoclass} tributo._common.dependencies.MissingOptionalDependency -:no-members: +```{autoexception} tributo._common.dependencies.MissingOptionalDependency ``` ```{autofunction} tributo._common.dependencies.probe_dependency @@ -73,8 +72,7 @@ documentation for every public stability tier. ## `tributo.exceptions` -```{autoclass} tributo.exceptions.AliasConflict -:no-members: +```{autoexception} tributo.exceptions.AliasConflict ``` ```{autoexception} tributo.exceptions.ArtifactCorruptedError @@ -131,8 +129,7 @@ documentation for every public stability tier. ```{autoexception} tributo.exceptions.ModelSchemaMismatchError ``` -```{autoclass} tributo.exceptions.PluginLoadIssue -:no-members: +```{autoexception} tributo.exceptions.PluginLoadIssue ``` ```{autoexception} tributo.exceptions.PostPublishCallbackError @@ -159,8 +156,7 @@ documentation for every public stability tier. ```{autoexception} tributo.TributoError ``` -```{autoclass} tributo.exceptions.UnsupportedArtifactFormat -:no-members: +```{autoexception} tributo.exceptions.UnsupportedArtifactFormat ``` @@ -173,3 +169,22 @@ documentation for every public stability tier. ```{autoclass} tributo.TributoClient :no-members: ``` + + +## `tributo.ray_jobs` + +```{autoclass} tributo.ray_jobs.RayJobSubmission +:no-members: +``` + +```{autofunction} tributo.ray_jobs.get_ray_job_logs +``` + +```{autofunction} tributo.ray_jobs.get_ray_job_status +``` + +```{autofunction} tributo.ray_jobs.stop_ray_job +``` + +```{autofunction} tributo.ray_jobs.submit_ray_job +``` diff --git a/docs/reference/api/extensions.md b/docs/reference/api/extensions.md index 9cbf6ed..8d27de9 100644 --- a/docs/reference/api/extensions.md +++ b/docs/reference/api/extensions.md @@ -9,6 +9,48 @@ a public annotation or moving a public object. Stable, Beta, and Alpha objects appear because Ray-style API policy requires documentation for every public stability tier. +## `tributo.integrations.broker` + +```{autoclass} tributo.integrations.broker.BrokerError +:no-members: +``` + +```{autoclass} tributo.integrations.broker.BrokerPlugin +:no-members: +``` + +```{autoclass} tributo.integrations.broker.BrokerRuntime +:no-members: +``` + +```{autoclass} tributo.integrations.broker.Message +:no-members: +``` + +```{autoclass} tributo.integrations.broker.TaskConsumer +:no-members: +``` + +```{autoclass} tributo.integrations.broker.TaskDisposition +:no-members: +``` + +```{autoclass} tributo.integrations.broker.TaskOutcome +:no-members: +``` + + +## `tributo.integrations.broker_registry` + +```{autoclass} tributo.integrations.broker_registry.BrokerDescriptor +:no-members: +``` + +```{autoclass} tributo.integrations.broker_registry.BrokerRegistry +:no-members: +``` + + ## `tributo.pipeline.core` ```{autoclass} tributo.pipeline.core.ArtifactRef diff --git a/docs/reference/api/inference-serving.md b/docs/reference/api/inference-serving.md index e446157..c63b96a 100644 --- a/docs/reference/api/inference-serving.md +++ b/docs/reference/api/inference-serving.md @@ -299,12 +299,18 @@ documentation for every public stability tier. ```{autofunction} tributo.inference.job_runner.submit_inference_request ``` +```{autofunction} tributo.inference.job_runner.submit_inference_request_with_identity +``` + ```{autofunction} tributo.inference.job_runner.submit_inference_request_with_retry ``` ```{autofunction} tributo.inference.job_runner.submit_resolved_inference ``` +```{autofunction} tributo.inference.job_runner.submit_resolved_inference_with_identity +``` + ```{autofunction} tributo.inference.job_runner.wait_for_job ``` diff --git a/src/tributo/cli.py b/src/tributo/cli.py index 986c540..3ffb6b1 100644 --- a/src/tributo/cli.py +++ b/src/tributo/cli.py @@ -26,7 +26,49 @@ logger = logging.getLogger(__name__) -@click.group() +class _LazyTributoGroup(click.Group): + """Load the broker command module only when that command is selected.""" + + def get_command(self, ctx: click.Context, cmd_name: str) -> click.Command | None: + if cmd_name == "broker": + from tributo.cli_broker import broker + + return broker + return super().get_command(ctx, cmd_name) + + def list_commands(self, ctx: click.Context) -> list[str]: + commands = super().list_commands(ctx) + if "broker" not in commands: + commands.append("broker") + return sorted(commands) + + def format_commands( + self, + ctx: click.Context, + formatter: click.HelpFormatter, + ) -> None: + """Render root help without importing the broker command module.""" + commands: list[tuple[str, click.Command]] = [] + for command_name in self.list_commands(ctx): + command = ( + click.Command( + "broker", + help="Discover and validate explicitly selected broker providers.", + ) + if command_name == "broker" + else self.get_command(ctx, command_name) + ) + if command is not None and not command.hidden: + commands.append((command_name, command)) + if not commands: + return + limit = formatter.width - 6 - max(len(name) for name, _ in commands) + rows = [(name, command.get_short_help_str(limit)) for name, command in commands] + with formatter.section("Commands"): + formatter.write_dl(rows) + + +@click.group(cls=_LazyTributoGroup) @click.version_option(package_name="tributo") def main(): """Tributo: Unified framework for submitting Ray Jobs.""" diff --git a/src/tributo/cli_broker.py b/src/tributo/cli_broker.py new file mode 100644 index 0000000..b292093 --- /dev/null +++ b/src/tributo/cli_broker.py @@ -0,0 +1,72 @@ +"""Broker-specific CLI commands, mounted by :mod:`tributo.cli`.""" + +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +import click + +from tributo.exceptions import JobConfigurationError +from tributo.integrations.broker_registry import BrokerRegistry + + +def _load_config(path: str) -> dict[str, Any]: + config_path = Path(path) + if config_path.suffix.lower() in {".yaml", ".yml"}: + raise click.ClickException("YAML broker config is not supported; use JSON.") + try: + value = json.loads(config_path.read_text(encoding="utf-8")) + except (OSError, json.JSONDecodeError) as exc: + raise click.ClickException(f"Unable to read broker config: {exc}") from exc + if not isinstance(value, dict): + raise click.ClickException("Broker config root must be a JSON object.") + return value + + +@click.group() +def broker() -> None: + """Discover and validate explicitly selected broker providers.""" + + +@broker.command("list") +def broker_list() -> None: + """List installed broker plugins without connecting to a broker.""" + registry = BrokerRegistry() + descriptors = registry.list() + for descriptor in descriptors: + capabilities = ",".join(descriptor.capabilities) or "-" + click.echo( + f"{descriptor.broker_id}\tapi={descriptor.api_version}" + f"\tstability={descriptor.stability}\tcapabilities={capabilities}" + ) + for diagnostic in registry.diagnostics(): + click.echo( + f"diagnostic\t{diagnostic.entry_point_name}\t{diagnostic.reason}", + err=True, + ) + + +@broker.command("validate") +@click.option("--broker", "broker_id", required=True) +@click.option("--config", "config_path", required=True, type=click.Path(exists=True)) +@click.option( + "--check-connectivity", + is_flag=True, + help="Ask the provider to perform an explicit connectivity probe.", +) +def broker_validate(broker_id: str, config_path: str, check_connectivity: bool) -> None: + """Validate provider-owned JSON config.""" + try: + config = _load_config(config_path) + BrokerRegistry().validate( + broker_id, + config, + check_connectivity=check_connectivity, + ) + except (JobConfigurationError, click.ClickException) as exc: + if isinstance(exc, click.ClickException): + raise + raise click.ClickException(str(exc)) from exc + click.echo(f"Broker configuration is valid: {broker_id}") diff --git a/src/tributo/inference/__init__.py b/src/tributo/inference/__init__.py index 1ca2c9a..3fe4c0b 100644 --- a/src/tributo/inference/__init__.py +++ b/src/tributo/inference/__init__.py @@ -40,7 +40,9 @@ def __call__(self, batch): ... from tributo.inference.job_runner import ( submit_inference_job, submit_inference_request, + submit_inference_request_with_identity, submit_resolved_inference, + submit_resolved_inference_with_identity, ) from tributo.inference.pipeline import ( InferenceConfig, @@ -72,5 +74,7 @@ def __call__(self, batch): ... "XGBoostONNXPredictor", "submit_inference_job", "submit_inference_request", + "submit_inference_request_with_identity", "submit_resolved_inference", + "submit_resolved_inference_with_identity", ] diff --git a/src/tributo/inference/job_runner.py b/src/tributo/inference/job_runner.py index 47280f6..9b7fd47 100644 --- a/src/tributo/inference/job_runner.py +++ b/src/tributo/inference/job_runner.py @@ -18,6 +18,7 @@ from tributo._common.submission_id import generate_submission_id from tributo.inference.contracts import InferenceRequest, ResolvedInference from tributo.inference.resolver import InferenceResolver +from tributo.ray_jobs import RayJobSubmission, _submit_ray_job_with_client from tributo.util.annotations import PublicAPI logger = logging.getLogger(__name__) @@ -52,7 +53,7 @@ class InferenceJobAttempt(BaseModel): run_id: str = Field(min_length=1) attempt_id: str = Field(min_length=1) submission_id: str = Field(min_length=1) - job_id: str = Field(min_length=1) + ray_job_id: str | None = Field(default=None, min_length=1) attempt_number: int = Field(ge=1) status: InferenceJobStatus retryable: bool = False @@ -69,7 +70,8 @@ class InferenceJobResult(BaseModel): ) run_id: str = Field(min_length=1) - job_id: str = Field(min_length=1) + submission_id: str = Field(min_length=1) + ray_job_id: str | None = Field(default=None, min_length=1) status: TerminalInferenceJobStatus logs: str = "" attempts: tuple[InferenceJobAttempt, ...] = () @@ -105,9 +107,7 @@ def submit_inference_job( resolved_run_id = run_id or generate_submission_id( "infer-run", config_path, str(sorted((env_vars or {}).items())) ) - submission_id = generate_submission_id( - "infer", resolved_run_id, attempt_id, config_path - ) + submission_id = generate_submission_id("infer", resolved_run_id, attempt_id) job_env = dict(env_vars or {}) job_env.update( { @@ -125,14 +125,20 @@ def submit_inference_job( f"python -m tributo.inference.batch_job --config {shlex.quote(config_path)}" ) client = _get_submission_client(dashboard_url) - job_id = _submit_attempt( + submission = _submit_attempt( client, entrypoint=entrypoint, runtime_env=runtime_env, + run_id=resolved_run_id, + attempt_id=attempt_id, submission_id=submission_id, ) - logger.info("Submitted inference job %s: config=%s", job_id, config_path) - return job_id + logger.info( + "Submitted inference job %s: config=%s", + submission.submission_id, + config_path, + ) + return submission.submission_id @PublicAPI(stability="alpha") @@ -144,17 +150,35 @@ def submit_inference_request( project_root: Path | None = None, resolver: InferenceResolver | None = None, ) -> str: - """Resolve once, serialize the credential-free plan, and submit it.""" + """Resolve once and return the accepted Ray Jobs submission identity.""" + return submit_inference_request_with_identity( + request, + dashboard_url=dashboard_url, + env_vars=env_vars, + project_root=project_root, + resolver=resolver, + ).submission_id + + +@PublicAPI(stability="alpha") +def submit_inference_request_with_identity( + request: InferenceRequest, + *, + dashboard_url: str = DEFAULT_DASHBOARD_URL, + env_vars: dict[str, str] | None = None, + project_root: Path | None = None, + resolver: InferenceResolver | None = None, +) -> RayJobSubmission: + """Resolve once, freeze the plan, and return complete submission identity.""" _validate_env_vars(env_vars) plan = (resolver or InferenceResolver()).resolve(request) client = _get_submission_client(dashboard_url) - job_id, _ = _submit_resolved_plan( + return _submit_resolved_plan( client, plan=plan, env_vars=env_vars, project_root=project_root, ) - return job_id @PublicAPI(stability="alpha") @@ -165,16 +189,32 @@ def submit_resolved_inference( env_vars: dict[str, str] | None = None, project_root: Path | None = None, ) -> str: - """Submit an already-frozen plan without re-resolving external aliases.""" + """Submit an already-frozen plan and return its submission identity.""" + return submit_resolved_inference_with_identity( + plan, + dashboard_url=dashboard_url, + env_vars=env_vars, + project_root=project_root, + ).submission_id + + +@PublicAPI(stability="alpha") +def submit_resolved_inference_with_identity( + plan: ResolvedInference, + *, + dashboard_url: str = DEFAULT_DASHBOARD_URL, + env_vars: dict[str, str] | None = None, + project_root: Path | None = None, +) -> RayJobSubmission: + """Submit a frozen plan and return workload-neutral Ray Jobs identity.""" _validate_env_vars(env_vars) client = _get_submission_client(dashboard_url) - job_id, _ = _submit_resolved_plan( + return _submit_resolved_plan( client, plan=plan, env_vars=env_vars, project_root=project_root, ) - return job_id @PublicAPI(stability="alpha") @@ -201,7 +241,7 @@ def submit_inference_request_with_retry( for attempt_number in range(1, max_attempts + 1): plan = _plan_for_attempt(first_plan, attempt_number) - job_id, submission_id = _submit_resolved_plan( + submission = _submit_resolved_plan( client, plan=plan, env_vars=env_vars, @@ -209,7 +249,7 @@ def submit_inference_request_with_retry( ) last = wait_for_job( client, - job_id, + submission.submission_id, timeout=timeout, poll_interval=poll_interval, ) @@ -225,8 +265,8 @@ def submit_inference_request_with_retry( InferenceJobAttempt( run_id=plan.run_id, attempt_id=plan.attempt_id, - submission_id=submission_id, - job_id=job_id, + submission_id=submission.submission_id, + ray_job_id=submission.ray_job_id, attempt_number=attempt_number, status=status, retryable=retryable, @@ -240,7 +280,8 @@ def submit_inference_request_with_retry( final_status = cast(TerminalInferenceJobStatus, attempts[-1].status) return InferenceJobResult( run_id=first_plan.run_id, - job_id=attempts[-1].job_id, + submission_id=attempts[-1].submission_id, + ray_job_id=attempts[-1].ray_job_id, status=final_status, logs=str(last.get("logs", "")), attempts=tuple(attempts), @@ -254,7 +295,7 @@ def _submit_resolved_plan( plan: ResolvedInference, env_vars: dict[str, str] | None, project_root: Path | None, -) -> tuple[str, str]: +) -> RayJobSubmission: plan = ResolvedInference.model_validate(plan.model_dump(mode="python")) encoded_plan = base64.urlsafe_b64encode( plan.model_dump_json().encode("utf-8") @@ -281,22 +322,25 @@ def _submit_resolved_plan( project_root=project_root, env_vars=job_env, ) - job_id = _submit_attempt( + submission = _submit_attempt( client, entrypoint=( "python -m tributo.inference.batch_job " "--resolved-plan-env TRIBUTO_INFERENCE_PLAN_B64" ), runtime_env=runtime_env, + run_id=plan.run_id, + attempt_id=plan.attempt_id, submission_id=plan.submission_id, + request_digest=plan.plan_digest, ) logger.info( - "Submitted inference attempt %s for run %s as Ray job %s", + "Submitted inference attempt %s for run %s as submission %s", plan.attempt_id, plan.run_id, - job_id, + submission.submission_id, ) - return job_id, plan.submission_id + return submission def _submit_attempt( @@ -304,38 +348,27 @@ def _submit_attempt( *, entrypoint: str, runtime_env: dict[str, Any], + run_id: str, + attempt_id: str, submission_id: str, -) -> str: - try: - return str( - client.submit_job( - entrypoint=entrypoint, - runtime_env=runtime_env, - submission_id=submission_id, - ) - ) - except Exception as exc: - try: - status = client.get_job_status(submission_id) - except Exception as query_exc: - raise exc from query_exc - if status is None: - raise exc from None - logger.warning( - "Reconciled inference submission %s after ambiguous error (status=%s)", - submission_id, - _ray_status_name(status), - ) - return submission_id + request_digest: str | None = None, +) -> RayJobSubmission: + return _submit_ray_job_with_client( + client, + entrypoint=entrypoint, + runtime_env=runtime_env, + run_id=run_id, + attempt_id=attempt_id, + submission_id=submission_id, + request_digest=request_digest, + ) def _plan_for_attempt( first_plan: ResolvedInference, attempt_number: int ) -> ResolvedInference: attempt_id = f"attempt-{attempt_number}" - submission_id = generate_submission_id( - "infer", first_plan.run_id, attempt_id, first_plan.plan_digest - ) + submission_id = generate_submission_id("infer", first_plan.run_id, attempt_id) return first_plan.model_copy( update={"attempt_id": attempt_id, "submission_id": submission_id} ) @@ -408,7 +441,9 @@ def _get_submission_client(dashboard_url: str) -> JobSubmissionClient: "map_ray_job_status", "submit_inference_job", "submit_inference_request", + "submit_inference_request_with_identity", "submit_inference_request_with_retry", "submit_resolved_inference", + "submit_resolved_inference_with_identity", "wait_for_job", ] diff --git a/src/tributo/inference/resolver.py b/src/tributo/inference/resolver.py index f3887af..6e83fea 100644 --- a/src/tributo/inference/resolver.py +++ b/src/tributo/inference/resolver.py @@ -117,7 +117,7 @@ def resolve(self, request: InferenceRequest) -> ResolvedInference: if os.environ.get("TRIBUTO_JOB_KIND") == "inference" else None ) or "attempt-1" - submission_id = generate_submission_id("infer", run_id, attempt_id, plan_digest) + submission_id = generate_submission_id("infer", run_id, attempt_id) return ResolvedInference( plan_digest=plan_digest, diff --git a/src/tributo/integrations/broker.py b/src/tributo/integrations/broker.py index f953451..61498f7 100644 --- a/src/tributo/integrations/broker.py +++ b/src/tributo/integrations/broker.py @@ -1,143 +1,170 @@ -"""Message broker abstraction for ML job lifecycle integration. +"""Transport-neutral contracts for independently installed broker providers. -Provides abstract interfaces for task consumption, event reporting, -and cancellation checking. Third-party implementations (Redis, Kafka, -Pulsar, etc.) will register via the ``tributo.brokers`` entry point group. - -.. note:: - - Plugin discovery and the ``tributo.brokers`` entry point group are - planned for v1.1. The ABCs in this module currently have zero - concrete implementations — they define the contract that broker - packages (e.g. ``tributo-broker-redis``) will fulfill. - -See Also: - :mod:`tributo.plugin` — plugin discovery infrastructure. +Core intentionally knows nothing about Redis, Kafka, RabbitMQ, or an external +operation protocol. Providers own transport semantics, request mapping, event +publication, and their production consume loop. """ from __future__ import annotations from abc import ABC, abstractmethod from dataclasses import dataclass, field -from typing import Any +from enum import StrEnum +from types import MappingProxyType +from typing import Any, ClassVar, Mapping + +from tributo.util.annotations import PublicAPI + +BROKER_API_VERSION = 1 -# ── Data types ──────────────────────────────────────────────────────────────── +@PublicAPI(stability="alpha") +class TaskDisposition(StrEnum): + """Provider decision for one broker delivery.""" -@dataclass + ACK = "ack" + RETRY = "retry" + REJECT = "reject" + + +@PublicAPI(stability="alpha") +@dataclass(frozen=True) class Message: - """A job request consumed from a message queue. + """Opaque provider delivery passed across the Core Broker boundary. - Attributes: - job_id: Unique identifier for the job. - payload: The deserialized request body (e.g. training config dict). - metadata: Optional routing headers or trace context. + ``delivery_token`` identifies the transport delivery, not a business + operation or Ray Job. Providers parse operation identity from ``payload``. + Metadata is restricted to string keys and values so transport clients and + credentials cannot be smuggled through this convenience surface. """ - job_id: str - payload: dict[str, Any] - metadata: dict[str, Any] = field(default_factory=dict) + payload: Any + delivery_token: str + metadata: Mapping[str, str] = field(default_factory=dict) + def __post_init__(self) -> None: + if not isinstance(self.delivery_token, str) or not self.delivery_token.strip(): + raise ValueError("Message.delivery_token must not be empty") + if not all( + isinstance(key, str) and isinstance(value, str) + for key, value in self.metadata.items() + ): + raise ValueError("Message.metadata requires string keys and values") + object.__setattr__(self, "metadata", MappingProxyType(dict(self.metadata))) -@dataclass -class JobResult: - """Outcome of a completed ML training job. - Attributes: - job_id: Unique identifier for the job. - status: Terminal status — ``"success"``, ``"failed"``, or ``"cancelled"``. - metrics: Final evaluation metrics (e.g. accuracy, loss). - artifacts: Paths or URIs of produced artifacts (model files, reports). - error: Human-readable error message when status is ``"failed"``. - """ +@PublicAPI(stability="alpha") +@dataclass(frozen=True) +class BrokerError: + """Minimal credential-safe provider error attached to a delivery outcome.""" - job_id: str - status: str - metrics: dict[str, float] = field(default_factory=dict) - artifacts: list[str] = field(default_factory=list) - error: str | None = None + code: str + sanitized_message: str + def __post_init__(self) -> None: + if not isinstance(self.code, str) or not self.code.strip(): + raise ValueError("BrokerError.code must not be empty") + if not isinstance(self.sanitized_message, str): + raise TypeError("BrokerError.sanitized_message must be a string") -# ── Abstract interfaces ─────────────────────────────────────────────────────── +@PublicAPI(stability="alpha") +@dataclass(frozen=True) +class TaskOutcome: + """Transport-neutral disposition returned by a provider runtime.""" -class TaskConsumer(ABC): - """Consume ML job requests from a message queue. + disposition: TaskDisposition + error: BrokerError | None = None - Implementations wrap a specific broker client (Redis Streams, - Kafka consumer group, etc.) and yield :class:`Message` objects. - """ + +@PublicAPI(stability="alpha") +class TaskConsumer(ABC): + """Consume opaque deliveries from a provider-owned transport.""" @abstractmethod def poll(self, timeout_ms: int = 5000) -> Message | None: - """Block until a job request arrives or *timeout_ms* expires. - - Args: - timeout_ms: Maximum time to wait in milliseconds. - - Returns: - A :class:`Message` if one is available, or ``None`` on timeout. - """ + """Block until a delivery arrives or ``timeout_ms`` expires.""" ... @abstractmethod def ack(self, message: Message) -> None: - """Acknowledge successful processing of *message*. - - After ``ack`` the broker guarantees the message will not be - redelivered. - """ + """Acknowledge a delivery according to provider semantics.""" ... + def retry(self, message: Message, error: BrokerError | None = None) -> None: + """Apply provider-defined retry semantics. -class EventReporter(ABC): - """Publish ML job lifecycle events. - - Every event is keyed by *job_id* so downstream systems can - reconstruct the full timeline of a job. - """ + The default deliberately does nothing. Providers must override this + hook unless leaving the delivery pending is their explicit retry + policy. + """ + del message, error - @abstractmethod - def report_phase(self, job_id: str, phase: str) -> None: - """Report a lifecycle phase transition. + def reject(self, message: Message, error: BrokerError | None = None) -> None: + """Apply a provider-defined permanent rejection policy. - Typical phases: ``"initializing"``, ``"training"``, - ``"exporting"``, ``"completed"``. + The default deliberately does nothing. A provider must override this + hook before returning ``REJECT`` for a delivery. """ - ... + del message, error - @abstractmethod - def report_metrics( - self, job_id: str, metrics: dict[str, float], progress: float - ) -> None: - """Report intermediate training metrics. + def recover_pending(self) -> int: + """Best-effort provider hook; Core makes no recovery guarantee.""" + return 0 + + def close(self) -> None: + """Close transport resources; the default is a no-op.""" + return None - Args: - job_id: The job identifier. - metrics: Current metric values (e.g. ``{"loss": 0.35}``). - progress: Progress fraction in ``[0.0, 1.0]``. - """ - ... +@PublicAPI(stability="alpha") +class BrokerRuntime(ABC): + """Provider runtime for mapping one delivery into a delivery outcome.""" + + @property @abstractmethod - def report_completed(self, job_id: str, result: JobResult) -> None: - """Report job completion with final results.""" + def consumer(self) -> TaskConsumer: + """Return the provider-owned consumer.""" ... @abstractmethod - def report_failed(self, job_id: str, error: str) -> None: - """Report job failure with error details.""" + def handle(self, message: Message) -> TaskOutcome: + """Handle one message without applying transport ACK side effects.""" ... + def close(self) -> None: + """Close provider resources.""" + self.consumer.close() -class CancellationChecker(ABC): - """Check whether a running job has been requested to cancel. - The training loop calls :meth:`is_cancelled` periodically and - stops early when it returns ``True``. - """ +@PublicAPI(stability="alpha") +class BrokerPlugin(ABC): + """Structural base for an independently installed broker provider.""" + + api_version: ClassVar[int] = BROKER_API_VERSION + broker_id: ClassVar[str] + capabilities: ClassVar[frozenset[str]] = frozenset() + stability: ClassVar[str] = "alpha" @abstractmethod - def is_cancelled(self, job_id: str) -> bool: - """Return ``True`` if *job_id* should stop early.""" + def validate_config( + self, config: Mapping[str, Any], *, check_connectivity: bool = False + ) -> None: + """Validate provider config and optionally probe connectivity.""" ... + + @abstractmethod + def create_runtime(self, config: Mapping[str, Any]) -> BrokerRuntime: + """Create a provider runtime; discovery itself must remain side-effect free.""" + ... + + +__all__ = [ + "BrokerError", + "BrokerPlugin", + "BrokerRuntime", + "Message", + "TaskConsumer", + "TaskDisposition", + "TaskOutcome", +] diff --git a/src/tributo/integrations/broker_registry.py b/src/tributo/integrations/broker_registry.py new file mode 100644 index 0000000..723ede9 --- /dev/null +++ b/src/tributo/integrations/broker_registry.py @@ -0,0 +1,99 @@ +"""Lazy discovery and explicit resolution for broker providers.""" + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any, cast + +from tributo.exceptions import JobConfigurationError +from tributo.exporting.models import PluginLoadDiagnostic +from tributo.integrations.broker import BrokerPlugin +from tributo.plugin import discover_broker_plugins, resolve_broker_plugin +from tributo.util.annotations import PublicAPI + + +@PublicAPI(stability="alpha") +@dataclass(frozen=True) +class BrokerDescriptor: + """Side-effect-free metadata reported for one installed provider.""" + + broker_id: str + api_version: int + capabilities: tuple[str, ...] + stability: str + + +@PublicAPI(stability="alpha") +class BrokerRegistry: + """Resolve only explicitly selected providers, failing closed on errors.""" + + def __init__(self) -> None: + self._diagnostics: list[PluginLoadDiagnostic] = [] + + def diagnostics(self) -> tuple[PluginLoadDiagnostic, ...]: + """Return non-fatal diagnostics from the most recent listing.""" + return tuple(self._diagnostics) + + def list(self) -> tuple[BrokerDescriptor, ...]: + """List providers without constructing them or connecting to a broker.""" + self._diagnostics.clear() + descriptors: list[BrokerDescriptor] = [] + seen: set[str] = set() + for cls in discover_broker_plugins(self._diagnostics): + broker_id = cls.broker_id + if broker_id in seen: + self._diagnostics.append( + PluginLoadDiagnostic( + group="tributo.brokers", + entry_point_name=broker_id, + reason="Duplicate broker_id discovered", + ) + ) + continue + seen.add(broker_id) + descriptors.append( + BrokerDescriptor( + broker_id=broker_id, + api_version=cls.api_version, + capabilities=tuple(sorted(cls.capabilities)), + stability=cls.stability, + ) + ) + return tuple(descriptors) + + def resolve(self, broker_id: str) -> BrokerPlugin: + """Load and instantiate one explicitly selected provider.""" + cls = resolve_broker_plugin(broker_id) + try: + plugin = cls() + except Exception as exc: + raise JobConfigurationError( + f"Failed to initialize broker {broker_id!r} ({type(exc).__name__})" + ) from exc + return cast(BrokerPlugin, plugin) + + def validate( + self, + broker_id: str, + config: Mapping[str, Any], + *, + check_connectivity: bool = False, + ) -> BrokerPlugin: + """Resolve a provider and delegate its config validation.""" + plugin = self.resolve(broker_id) + try: + plugin.validate_config( + config, + check_connectivity=check_connectivity, + ) + except JobConfigurationError: + raise + except Exception as exc: + raise JobConfigurationError( + f"Broker {broker_id!r} rejected configuration ({type(exc).__name__})" + ) from exc + return plugin + + +__all__ = ["BrokerDescriptor", "BrokerRegistry"] diff --git a/src/tributo/job.py b/src/tributo/job.py index 5070e62..703d1e6 100644 --- a/src/tributo/job.py +++ b/src/tributo/job.py @@ -36,8 +36,8 @@ class TributoClient: Example: >>> client = TributoClient("http://127.0.0.1:8265") - >>> job_id = client.submit(entrypoint="python script.py") - >>> status = client.get_status(job_id) + >>> submission_id = client.submit(entrypoint="python script.py") + >>> status = client.get_status(submission_id) """ def __init__(self, address: str): @@ -88,7 +88,7 @@ def submit( ``working_dir`` upload. Returns: - Job ID string. + Ray Jobs submission identity. Raises: JobSubmissionError: If submission fails. @@ -147,7 +147,8 @@ def get_status(self, job_id: str) -> str: """Get the status of a submitted job. Args: - job_id: The job ID to query. + job_id: Ray Jobs submission identity to query. The parameter name + is retained for compatibility. Returns: Job status string (e.g. ``"RUNNING"``, ``"SUCCEEDED"``). @@ -166,7 +167,8 @@ def get_logs(self, job_id: str) -> str: """Get logs for a submitted job. Args: - job_id: The job ID to query. + job_id: Ray Jobs submission identity to query. The parameter name + is retained for compatibility. Returns: Job logs as a string. @@ -184,7 +186,8 @@ def stop_job(self, job_id: str) -> bool: """Stop a running job. Args: - job_id: The job ID to stop. + job_id: Ray Jobs submission identity to stop. The parameter name + is retained for compatibility. Returns: True if the job was stopped successfully. diff --git a/src/tributo/plugin.py b/src/tributo/plugin.py index c7be8cf..b48a064 100644 --- a/src/tributo/plugin.py +++ b/src/tributo/plugin.py @@ -895,3 +895,153 @@ def _discover_storage_adapter_plugins( classes.append(cls) logger.info("Discovered storage adapter %r (%s)", ep.name, ep.value) return classes + + +# ═══════════════════════════════════════════════════════════════════════════════ +# Broker plugins +# ═══════════════════════════════════════════════════════════════════════════════ + + +def _broker_contract_issues(cls: Any) -> tuple[str, ...]: + """Return structural API issues without instantiating a provider.""" + issues: list[str] = [] + if not isinstance(cls, type): + return ("provider class",) + if type(getattr(cls, "api_version", None)) is not int: + issues.append("api_version") + broker_id = getattr(cls, "broker_id", None) + if not isinstance(broker_id, str) or not broker_id.strip(): + issues.append("broker_id") + capabilities = getattr(cls, "capabilities", None) + if not isinstance(capabilities, frozenset) or not all( + isinstance(value, str) and bool(value.strip()) for value in capabilities + ): + issues.append("capabilities") + if getattr(cls, "stability", None) not in {"alpha", "beta", "stable"}: + issues.append("stability") + for method in ( + "validate_config", + "create_runtime", + ): + if not callable(getattr(cls, method, None)): + issues.append(method) + return tuple(issues) + + +def discover_broker_plugins( + diagnostics: list[PluginLoadDiagnostic] | None = None, +) -> list[type[Any]]: + """Discover broker provider classes from ``tributo.brokers``. + + Discovery is fail-open and never instantiates a provider. In particular, + it cannot create a Redis client or perform a network probe. Explicit + resolution is provided by :func:`resolve_broker_plugin` and is + fail-closed. + """ + from tributo.integrations.broker import BROKER_API_VERSION + + enabled = _get_enabled_plugins() + classes: list[type[Any]] = [] + for ep in _iter_entry_points("tributo.brokers"): + if enabled is not None and ep.name not in enabled: + logger.debug("Skipping broker plugin %r", ep.name) + continue + try: + cls = ep.load() + except Exception as exc: + logger.warning( + "Failed to load broker plugin %r (%s; %s)", + ep.name, + ep.value, + type(exc).__name__, + ) + _record_diagnostic( + diagnostics, + "tributo.brokers", + ep.name, + f"Failed to load entry point ({type(exc).__name__})", + error_type=type(exc).__name__, + ) + continue + + issues = _broker_contract_issues(cls) + if issues: + logger.warning( + "Broker plugin %r does not satisfy API v%d: %s", + ep.name, + BROKER_API_VERSION, + ", ".join(issues), + ) + _record_diagnostic( + diagnostics, + "tributo.brokers", + ep.name, + "Missing or invalid BrokerPlugin members: " + ", ".join(issues), + ) + continue + if cls.api_version != BROKER_API_VERSION: + reason = ( + f"Unsupported BrokerPlugin api_version {cls.api_version!r}; " + f"expected {BROKER_API_VERSION}" + ) + _record_diagnostic(diagnostics, "tributo.brokers", ep.name, reason) + logger.warning("Broker plugin %r: %s", ep.name, reason) + continue + if ep.name != cls.broker_id: + reason = ( + f"Entry-point name {ep.name!r} does not match broker_id " + f"{cls.broker_id!r}" + ) + _record_diagnostic(diagnostics, "tributo.brokers", ep.name, reason) + logger.warning("Broker plugin %r: %s", ep.name, reason) + continue + classes.append(cls) + logger.info("Discovered broker plugin %r (%s)", ep.name, ep.value) + return classes + + +def resolve_broker_plugin(broker_id: str) -> type[Any]: + """Resolve one explicitly selected broker, using fail-closed semantics.""" + enabled = _get_enabled_plugins() + if enabled is not None and broker_id not in enabled: + raise JobConfigurationError( + f"Broker {broker_id!r} is disabled by TRIBUTO_PLUGINS" + ) + + matches = [ + ep for ep in _iter_entry_points("tributo.brokers") if ep.name == broker_id + ] + if not matches: + raise JobConfigurationError(f"Unknown broker {broker_id!r}") + if len(matches) > 1: + raise JobConfigurationError( + f"Multiple entry points are registered for broker {broker_id!r}" + ) + + ep = matches[0] + try: + cls = cast(type[Any], ep.load()) + except Exception as exc: + raise JobConfigurationError( + f"Failed to load broker {broker_id!r} ({type(exc).__name__})" + ) from exc + + issues = _broker_contract_issues(cls) + from tributo.integrations.broker import BROKER_API_VERSION + + if issues: + raise JobConfigurationError( + f"Broker {broker_id!r} does not implement the BrokerPlugin v" + f"{BROKER_API_VERSION} contract: {', '.join(issues)}" + ) + if cls.api_version != BROKER_API_VERSION: + raise JobConfigurationError( + f"Broker {broker_id!r} has unsupported api_version " + f"{cls.api_version!r}; expected {BROKER_API_VERSION}" + ) + if cls.broker_id != ep.name: + raise JobConfigurationError( + f"Broker entry-point name {ep.name!r} does not match broker_id " + f"{cls.broker_id!r}" + ) + return cast(type[Any], cls) diff --git a/src/tributo/ray_jobs.py b/src/tributo/ray_jobs.py new file mode 100644 index 0000000..d61d82a --- /dev/null +++ b/src/tributo/ray_jobs.py @@ -0,0 +1,266 @@ +"""Workload-neutral Ray Jobs admission and control helpers.""" + +from __future__ import annotations + +import logging +from pathlib import Path +from typing import Any + +from pydantic import BaseModel, ConfigDict, Field +from ray.job_submission import JobSubmissionClient + +from tributo._common import DEFAULT_DASHBOARD_URL, build_runtime_env +from tributo._common.retry import retry_with_exponential_backoff +from tributo._common.submission_id import generate_submission_id +from tributo.util.annotations import PublicAPI + +logger = logging.getLogger(__name__) + +_RESERVED_ENV_KEYS = frozenset( + { + "TRIBUTO_RUN_ID", + "TRIBUTO_ATTEMPT_ID", + "TRIBUTO_SUBMISSION_ID", + } +) +_REQUEST_DIGEST_METADATA_KEY = "tributo.request_digest" + + +def _require_submission_id(submission_id: str) -> None: + if not isinstance(submission_id, str) or not submission_id.strip(): + raise ValueError("submission_id must not be empty") + + +@PublicAPI(stability="alpha") +class RayJobSubmission(BaseModel): + """Identity returned after one Ray Jobs attempt is accepted or reconciled.""" + + model_config = ConfigDict(frozen=True, extra="forbid") + + run_id: str = Field(min_length=1) + attempt_id: str = Field(min_length=1) + submission_id: str = Field(min_length=1) + ray_job_id: str | None = Field(default=None, min_length=1) + request_digest: str | None = Field(default=None, min_length=1) + + +def _validate_inputs( + operation_namespace: str, + run_id: str, + attempt_id: str, + env_vars: dict[str, str] | None, + metadata: dict[str, str] | None, + request_digest: str | None, +) -> None: + for name, value in ( + ("operation_namespace", operation_namespace), + ("run_id", run_id), + ("attempt_id", attempt_id), + ): + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"{name} must not be empty") + conflicts = _RESERVED_ENV_KEYS.intersection(env_vars or {}) + if conflicts: + raise ValueError( + "env_vars must not override Ray Job identity: " + + ", ".join(sorted(conflicts)) + ) + if metadata is not None and _REQUEST_DIGEST_METADATA_KEY in metadata: + raise ValueError( + f"metadata must not define reserved key {_REQUEST_DIGEST_METADATA_KEY!r}" + ) + if request_digest is not None and not request_digest.strip(): + raise ValueError("request_digest must not be empty") + + +def _ray_job_id(client: JobSubmissionClient, submission_id: str) -> str | None: + get_job_info = getattr(client, "get_job_info", None) + if not callable(get_job_info): + return None + try: + info = get_job_info(submission_id) + except Exception as exc: + logger.debug( + "Ray JobDetails unavailable for submission %s (%s)", + submission_id, + type(exc).__name__, + ) + return None + value = getattr(info, "job_id", None) + return value if isinstance(value, str) and value else None + + +def _submit_ray_job_with_client( + client: JobSubmissionClient, + *, + entrypoint: str, + run_id: str, + attempt_id: str, + submission_id: str, + runtime_env: dict[str, Any] | None, + metadata: dict[str, str] | None = None, + request_digest: str | None = None, + entrypoint_num_cpus: float | None = None, + entrypoint_num_gpus: float | None = None, + entrypoint_memory: int | None = None, +) -> RayJobSubmission: + """Submit through an existing client; shared by Core workload adapters.""" + job_metadata = dict(metadata or {}) + if _REQUEST_DIGEST_METADATA_KEY in job_metadata: + raise ValueError( + f"metadata must not define reserved key {_REQUEST_DIGEST_METADATA_KEY!r}" + ) + if request_digest is not None: + if not request_digest.strip(): + raise ValueError("request_digest must not be empty") + job_metadata[_REQUEST_DIGEST_METADATA_KEY] = request_digest + + try: + client.submit_job( + entrypoint=entrypoint, + runtime_env=runtime_env, + metadata=job_metadata or None, + submission_id=submission_id, + entrypoint_num_cpus=entrypoint_num_cpus, + entrypoint_num_gpus=entrypoint_num_gpus, + entrypoint_memory=entrypoint_memory, + ) + except Exception as submit_error: + try: + status = client.get_job_status(submission_id) + except Exception as query_error: + raise submit_error from query_error + if status is None: + raise submit_error from None + logger.warning( + "Reconciled Ray submission %s after an ambiguous submit response", + submission_id, + ) + + return RayJobSubmission( + run_id=run_id, + attempt_id=attempt_id, + submission_id=submission_id, + ray_job_id=_ray_job_id(client, submission_id), + request_digest=request_digest, + ) + + +@PublicAPI(stability="alpha") +def submit_ray_job( + entrypoint: str, + *, + operation_namespace: str, + run_id: str, + attempt_id: str = "attempt-1", + dashboard_url: str = DEFAULT_DASHBOARD_URL, + env_vars: dict[str, str] | None = None, + project_root: Path | None = None, + extra_excludes: list[str] | None = None, + metadata: dict[str, str] | None = None, + request_digest: str | None = None, + entrypoint_num_cpus: float | None = None, + entrypoint_num_gpus: float | None = None, + entrypoint_memory: int | None = None, +) -> RayJobSubmission: + """Submit one deterministic Ray Job and reconcile an ambiguous response.""" + + if not entrypoint.strip(): + raise ValueError("entrypoint must not be empty") + _validate_inputs( + operation_namespace, + run_id, + attempt_id, + env_vars, + metadata, + request_digest, + ) + submission_id = generate_submission_id( + operation_namespace, + run_id, + attempt_id, + ) + job_env = dict(env_vars or {}) + job_env.update( + { + "TRIBUTO_RUN_ID": run_id, + "TRIBUTO_ATTEMPT_ID": attempt_id, + "TRIBUTO_SUBMISSION_ID": submission_id, + } + ) + runtime_env = build_runtime_env( + project_root=project_root, + env_vars=job_env, + extra_excludes=extra_excludes, + ) + client = _get_submission_client(dashboard_url) + return _submit_ray_job_with_client( + client, + entrypoint=entrypoint, + run_id=run_id, + attempt_id=attempt_id, + submission_id=submission_id, + runtime_env=runtime_env, + metadata=metadata, + request_digest=request_digest, + entrypoint_num_cpus=entrypoint_num_cpus, + entrypoint_num_gpus=entrypoint_num_gpus, + entrypoint_memory=entrypoint_memory, + ) + + +@PublicAPI(stability="alpha") +def get_ray_job_status( + submission_id: str, + *, + dashboard_url: str = DEFAULT_DASHBOARD_URL, +) -> str: + """Return the normalized Ray Jobs status for a submission identity.""" + + _require_submission_id(submission_id) + status = _get_submission_client(dashboard_url).get_job_status(submission_id) + if status is None: + raise LookupError(f"Unknown Ray submission {submission_id!r}") + return str(getattr(status, "value", status)).upper() + + +@PublicAPI(stability="alpha") +def get_ray_job_logs( + submission_id: str, + *, + dashboard_url: str = DEFAULT_DASHBOARD_URL, +) -> str: + """Return logs for the Ray Job identified by ``submission_id``.""" + + _require_submission_id(submission_id) + return _get_submission_client(dashboard_url).get_job_logs(submission_id) + + +@PublicAPI(stability="alpha") +def stop_ray_job( + submission_id: str, + *, + dashboard_url: str = DEFAULT_DASHBOARD_URL, +) -> bool: + """Request that Ray stop the job identified by ``submission_id``.""" + + _require_submission_id(submission_id) + return bool(_get_submission_client(dashboard_url).stop_job(submission_id)) + + +@retry_with_exponential_backoff( + max_retries=3, + base_delay=1.0, + exceptions=(ConnectionError, TimeoutError, OSError), +) +def _get_submission_client(dashboard_url: str) -> JobSubmissionClient: + return JobSubmissionClient(dashboard_url) + + +__all__ = [ + "RayJobSubmission", + "get_ray_job_logs", + "get_ray_job_status", + "stop_ray_job", + "submit_ray_job", +] diff --git a/src/tributo/training/__init__.py b/src/tributo/training/__init__.py index b9ca95a..73339bd 100644 --- a/src/tributo/training/__init__.py +++ b/src/tributo/training/__init__.py @@ -57,6 +57,7 @@ JobAttempt, TrainingJobResult, submit_training_job, + submit_training_job_with_identity, submit_training_job_with_retry, wait_for_job, ) @@ -67,7 +68,11 @@ ) from tributo.training.registry import get_trainer, list_trainers, register from tributo.training.tune_runner import TuneRunner, extract_best_params - from tributo.training.xgboost_trainer import build_trainer, run_training_from_json + from tributo.training.xgboost_trainer import ( + build_trainer, + run_training_from_json, + run_training_result_with_config, + ) _LAZY_EXPORTS = { "get_trainer": ("tributo.training.registry", "get_trainer"), @@ -86,6 +91,10 @@ "tributo.training.job_submitter", "submit_training_job_with_retry", ), + "submit_training_job_with_identity": ( + "tributo.training.job_submitter", + "submit_training_job_with_identity", + ), "wait_for_job": ("tributo.training.job_submitter", "wait_for_job"), "TuneRunner": ("tributo.training.tune_runner", "TuneRunner"), "extract_best_params": ( @@ -97,6 +106,10 @@ "tributo.training.xgboost_trainer", "run_training_from_json", ), + "run_training_result_with_config": ( + "tributo.training.xgboost_trainer", + "run_training_result_with_config", + ), "DNNTrainerImpl": ("tributo.training.dnn_trainer", "DNNTrainerImpl"), "run_dnn_training_from_json": ( "tributo.training.dnn_trainer", @@ -162,6 +175,7 @@ def __getattr__(name: str) -> Any: "JobAttempt", "TrainingJobResult", "submit_training_job", + "submit_training_job_with_identity", "submit_training_job_with_retry", "wait_for_job", "export_to_onnx", @@ -177,6 +191,7 @@ def __getattr__(name: str) -> Any: "warn_search_space_conflicts", "build_trainer", "run_training_from_json", + "run_training_result_with_config", ] if importlib.util.find_spec("torch") is not None: diff --git a/src/tributo/training/job_submitter.py b/src/tributo/training/job_submitter.py index 0b1317b..3a07f81 100644 --- a/src/tributo/training/job_submitter.py +++ b/src/tributo/training/job_submitter.py @@ -6,7 +6,6 @@ from __future__ import annotations -import logging import time from collections.abc import Callable from pathlib import Path @@ -20,13 +19,19 @@ from tributo._common.submission_id import generate_submission_id from tributo.algorithms.api import EnvironmentSpec from tributo.algorithms.api.artifacts import AlgorithmArtifact, ImageProfile +from tributo.ray_jobs import RayJobSubmission, _submit_ray_job_with_client from tributo.util.annotations import PublicAPI -logger = logging.getLogger(__name__) DEFAULT_TIMEOUT = 180 JobAttemptStatus = Literal["PENDING", "RUNNING", "SUCCEEDED", "FAILED", "STOPPED"] TerminalJobStatus = Literal["SUCCEEDED", "FAILED", "STOPPED"] -_RESERVED_ENV_KEYS = frozenset({"TRIBUTO_RUN_ID", "TRIBUTO_ATTEMPT_ID"}) +_RESERVED_ENV_KEYS = frozenset( + { + "TRIBUTO_RUN_ID", + "TRIBUTO_ATTEMPT_ID", + "TRIBUTO_SUBMISSION_ID", + } +) def _resolve_algorithm_dependencies( @@ -49,7 +54,7 @@ class JobAttempt(BaseModel): run_id: str = Field(..., min_length=1) attempt_id: str = Field(..., min_length=1) submission_id: str = Field(..., min_length=1) - job_id: str = Field(..., min_length=1) + ray_job_id: str | None = Field(default=None, min_length=1) attempt_number: int = Field(..., ge=1) status: JobAttemptStatus retryable: bool = False @@ -63,7 +68,8 @@ class TrainingJobResult(BaseModel): run_id: str = Field(..., min_length=1) bundle_id: str = Field(..., min_length=1) - job_id: str = Field(..., min_length=1) + submission_id: str = Field(..., min_length=1) + ray_job_id: str | None = Field(default=None, min_length=1) status: TerminalJobStatus logs: str = "" attempts: tuple[JobAttempt, ...] = () @@ -106,35 +112,21 @@ def _submit_training_job_attempt( runtime_env: dict[str, Any], run_id: str, attempt_id: str, + submission_id: str, metadata: dict[str, str] | None = None, -) -> tuple[str, str]: - """Submit one stable attempt and reconcile ambiguous server responses.""" - submission_id = generate_submission_id("train", run_id, attempt_id) - try: - job_id = client.submit_job( - entrypoint=entrypoint, - runtime_env=runtime_env, - metadata=metadata, - submission_id=submission_id, - ) - except Exception as exc: - # The request may have reached Ray before the client observed an - # error. Query the deterministic submission ID before considering a - # retry; inventing another ID here could run the same attempt twice. - try: - status = client.get_job_status(submission_id) - except Exception as query_exc: - raise exc from query_exc - if status is None: - raise exc from None - logger.warning( - "Reconciled submission %s after ambiguous error (status=%s)", - submission_id, - _status_name(status), - ) - return submission_id, submission_id - logger.info("Submitted training job %s: %s", job_id, entrypoint) - return str(job_id), submission_id + request_digest: str | None = None, +) -> RayJobSubmission: + """Submit one stable attempt through the workload-neutral Core helper.""" + return _submit_ray_job_with_client( + client, + entrypoint=entrypoint, + runtime_env=runtime_env, + run_id=run_id, + attempt_id=attempt_id, + submission_id=submission_id, + metadata=metadata, + request_digest=request_digest, + ) @PublicAPI(stability="beta") @@ -152,6 +144,7 @@ def submit_training_job( image_profile: ImageProfile | None = None, declared_dependencies: tuple[str, ...] = (), environment: EnvironmentSpec | None = None, + request_digest: str | None = None, ) -> str: """Submit a training job via the Ray Jobs API. @@ -178,21 +171,62 @@ def submit_training_job( declared_dependencies: Additional PEP 508 constraints to preflight. environment: Optional formal or ``from_sklearn()`` EnvironmentSpec; its dependencies are merged into the same preflight. + request_digest: Optional credential-free request digest stored only as + Ray submission metadata. Returns: - Submitted job ID on success. + Deterministic Ray Jobs submission identity on success. Raises: RuntimeError: Submission failed. """ + return submit_training_job_with_identity( + entrypoint, + dashboard_url=dashboard_url, + env_vars=env_vars, + project_root=project_root, + extra_excludes=extra_excludes, + run_id=run_id, + attempt_id=attempt_id, + metadata=metadata, + algorithm_artifact=algorithm_artifact, + image_profile=image_profile, + declared_dependencies=declared_dependencies, + environment=environment, + request_digest=request_digest, + ).submission_id + + +@PublicAPI(stability="alpha") +def submit_training_job_with_identity( + entrypoint: str, + *, + dashboard_url: str = DEFAULT_DASHBOARD_URL, + env_vars: dict[str, str] | None = None, + project_root: Path | None = None, + extra_excludes: list[str] | None = None, + run_id: str | None = None, + attempt_id: str | None = None, + metadata: dict[str, str] | None = None, + algorithm_artifact: AlgorithmArtifact | None = None, + image_profile: ImageProfile | None = None, + declared_dependencies: tuple[str, ...] = (), + environment: EnvironmentSpec | None = None, + request_digest: str | None = None, +) -> RayJobSubmission: + """Submit one attempt and return workload-neutral Ray Jobs identity.""" resolved_run_id = _resolve_run_id(entrypoint, env_vars, run_id) resolved_attempt_id = attempt_id or "attempt-1" + submission_id = generate_submission_id( + "train", resolved_run_id, resolved_attempt_id + ) _validate_metadata(metadata) job_env_vars = dict(env_vars or {}) job_env_vars.update( { "TRIBUTO_RUN_ID": resolved_run_id, "TRIBUTO_ATTEMPT_ID": resolved_attempt_id, + "TRIBUTO_SUBMISSION_ID": submission_id, } ) runtime_env = build_runtime_env( @@ -208,15 +242,16 @@ def submit_training_job( ) client = _get_submission_client(dashboard_url) - job_id, _submission_id = _submit_training_job_attempt( + return _submit_training_job_attempt( client, entrypoint=entrypoint, runtime_env=runtime_env, run_id=resolved_run_id, attempt_id=resolved_attempt_id, + submission_id=submission_id, metadata=metadata, + request_digest=request_digest, ) - return job_id @PublicAPI(stability="beta") @@ -237,6 +272,7 @@ def submit_training_job_with_retry( image_profile: ImageProfile | None = None, declared_dependencies: tuple[str, ...] = (), environment: EnvironmentSpec | None = None, + request_digest: str | None = None, ) -> TrainingJobResult: """Submit, reconcile and optionally retry a training run. @@ -268,8 +304,10 @@ def submit_training_job_with_retry( for attempt_number in range(1, max_attempts + 1): attempt_id = f"attempt-{attempt_number}" + submission_id = generate_submission_id("train", resolved_run_id, attempt_id) attempt_env_vars = dict(job_env_vars) attempt_env_vars["TRIBUTO_ATTEMPT_ID"] = attempt_id + attempt_env_vars["TRIBUTO_SUBMISSION_ID"] = submission_id runtime_env = build_runtime_env( project_root=project_root, env_vars=attempt_env_vars, @@ -281,17 +319,19 @@ def submit_training_job_with_retry( environment, ), ) - job_id, submission_id = _submit_training_job_attempt( + submission = _submit_training_job_attempt( client, entrypoint=entrypoint, runtime_env=runtime_env, run_id=resolved_run_id, attempt_id=attempt_id, + submission_id=submission_id, metadata=metadata, + request_digest=request_digest, ) last_result = wait_for_job( client, - job_id, + submission.submission_id, timeout=timeout, poll_interval=poll_interval, ) @@ -307,8 +347,8 @@ def submit_training_job_with_retry( JobAttempt( run_id=resolved_run_id, attempt_id=attempt_id, - submission_id=submission_id, - job_id=job_id, + submission_id=submission.submission_id, + ray_job_id=submission.ray_job_id, attempt_number=attempt_number, status=status_name, retryable=retryable, @@ -325,7 +365,8 @@ def submit_training_job_with_retry( return TrainingJobResult( run_id=resolved_run_id, bundle_id=bundle_id_for_request(resolved_run_id), - job_id=attempts[-1].job_id, + submission_id=attempts[-1].submission_id, + ray_job_id=attempts[-1].ray_job_id, status=cast(TerminalJobStatus, final_status), logs=str(last_result.get("logs", "")), attempts=tuple(attempts), @@ -360,7 +401,8 @@ def wait_for_job( Args: client: Ray Jobs API client. - job_id: Job ID to wait for. + job_id: Ray Jobs submission identity to wait for. The parameter name + is retained for compatibility. timeout: Maximum wait time in seconds. poll_interval: Polling interval in seconds. diff --git a/src/tributo/training/xgboost_trainer.py b/src/tributo/training/xgboost_trainer.py index d3ae389..0eab8f7 100644 --- a/src/tributo/training/xgboost_trainer.py +++ b/src/tributo/training/xgboost_trainer.py @@ -21,7 +21,6 @@ XGBOOST_DESCRIPTOR, build_legacy_spec, ) -from tributo.integrations.broker import CancellationChecker from tributo.training.base import BaseTrainer from tributo.training.checkpoint import ResumeConfig from tributo.training.resource import ( @@ -37,6 +36,8 @@ import ray.data from ray.train.xgboost import XGBoostTrainer + from tributo.training.results import TrainingResult + # XGBoost params reserved by Tributo: silently passing these as native # training parameters would change the execution path (e.g. external-memory # / data_iter) without going through the materialization-budget contract. @@ -528,30 +529,6 @@ def train_loop_per_worker(config: dict[str, Any]) -> None: # pass would double-count bytes and rows against the shared budget. test_labels: list[Any] = [] - # Cancel signal — inject CancellationChecker via config when using a broker. - # TODO(v1.1): _tributo_cancel_key and _tributo_cancel_checker are dead code - # until a broker implementation (e.g. tributo-broker-redis) populates them. - _cancel_key: str | None = config.get("_tributo_cancel_key") - _cancel_checker: CancellationChecker | None = config.get("_tributo_cancel_checker") - - class _CancelCallback(xgboost.callback.TrainingCallback): - """Check cancellation signal after each iteration (broker protocol).""" - - def after_iteration( - self, model: xgboost.Booster, epoch: int, evts_log: dict - ) -> bool: - if _cancel_key is None or _cancel_checker is None: - return False - try: - return _cancel_checker.is_cancelled(_cancel_key) - except Exception: - logger.warning( - "Cancellation check failed for job %s", - _cancel_key, - exc_info=True, - ) - return False # transient error → don't cancel - def _make_quantile_dmatrix( dataset_key: str, ref: xgboost.QuantileDMatrix | None = None, @@ -777,7 +754,7 @@ def after_iteration( evals=evals, evals_result=current_evals_result, early_stopping_rounds=config.get("early_stopping_rounds"), - callbacks=[_CancelCallback(), _ResumeCheckpointCallback()], + callbacks=[_ResumeCheckpointCallback()], xgb_model=initial_booster, ) evals_result = _merge_xgb_eval_results( @@ -1240,6 +1217,17 @@ def run_training_with_config(config: dict[str, Any]) -> dict[str, Any]: return trainer.run() +@PublicAPI(stability="alpha") +def run_training_result_with_config(config: dict[str, Any]) -> TrainingResult: + """Run training in-process and return the structured terminal result.""" + from tributo.training.results import TrainingResult + + summary = run_training_with_config(config) + return TrainingResult.model_validate( + {key: summary.get(key) for key in TrainingResult.model_fields if key in summary} + ) + + # Built-in registration _trainer_spec = build_legacy_spec( diff --git a/tests/docs/test_docs_tooling.py b/tests/docs/test_docs_tooling.py index 04bffd6..5e9d783 100644 --- a/tests/docs/test_docs_tooling.py +++ b/tests/docs/test_docs_tooling.py @@ -2,6 +2,7 @@ from __future__ import annotations +import ast import importlib import json import runpy @@ -128,6 +129,23 @@ def test_generated_public_api_reference_covers_source_inventory() -> None: assert check_pages(inventory) == [] +def test_public_api_inventory_classifies_exceptions_by_base_class() -> None: + tree = ast.parse( + """class BrokerError: + pass + +class ActualError(Exception): + pass +""" + ) + classes = [node for node in tree.body if isinstance(node, ast.ClassDef)] + + assert [public_api_generator._is_exception_class(node) for node in classes] == [ + False, + True, + ] + + def test_api_reference_validation_reuses_source_inventory( monkeypatch: pytest.MonkeyPatch, ) -> None: diff --git a/tests/inference/test_job_runner.py b/tests/inference/test_job_runner.py index 5630e79..2b654e7 100644 --- a/tests/inference/test_job_runner.py +++ b/tests/inference/test_job_runner.py @@ -13,6 +13,7 @@ from tributo.inference.contracts import ResolvedInference from tributo.inference.job_runner import ( map_ray_job_status, + submit_inference_job, submit_inference_request, submit_inference_request_with_retry, submit_resolved_inference, @@ -32,8 +33,6 @@ def test_passes_submission_id(self): "tributo.inference.job_runner.JobSubmissionClient", return_value=mock_client, ): - from tributo.inference.job_runner import submit_inference_job - submit_inference_job( config_path="jobs/inference.yaml", dashboard_url="http://127.0.0.1:8265", @@ -54,8 +53,6 @@ def test_submission_id_is_deterministic(self): "tributo.inference.job_runner.JobSubmissionClient", return_value=mock_client, ): - from tributo.inference.job_runner import submit_inference_job - submit_inference_job( config_path="jobs/inference.yaml", dashboard_url="http://127.0.0.1:8265", @@ -75,14 +72,12 @@ def test_client_creation_is_retried_on_connection_error(self): "tributo.inference.job_runner.JobSubmissionClient", side_effect=side_effects, ) as mock_constructor: - from tributo.inference.job_runner import submit_inference_job - job_id = submit_inference_job( config_path="jobs/inference.yaml", dashboard_url="http://127.0.0.1:8265", ) - assert job_id == "job-123" + assert job_id.startswith("tributo-infer-") assert mock_constructor.call_count == 2 @@ -113,7 +108,7 @@ def test_already_resolved_plan_is_submitted_without_a_resolver(self) -> None: ): job_id = submit_resolved_inference(plan) - assert job_id == "job-frozen" + assert job_id == plan.submission_id assert client.submit_job.call_args.kwargs["submission_id"] == plan.submission_id def test_frozen_plan_is_transported_without_re_resolution_in_job(self) -> None: @@ -135,7 +130,7 @@ def test_frozen_plan_is_transported_without_re_resolution_in_job(self) -> None: ): job_id = submit_inference_request(object(), resolver=resolver) - assert job_id == "job-1" + assert job_id == plan.submission_id resolver.resolve.assert_called_once() call = client.submit_job.call_args.kwargs assert call["submission_id"] == plan.submission_id diff --git a/tests/test_broker.py b/tests/test_broker.py new file mode 100644 index 0000000..b2099ba --- /dev/null +++ b/tests/test_broker.py @@ -0,0 +1,201 @@ +"""Core Broker API v1 and lazy provider-discovery tests.""" + +from __future__ import annotations + +import operator +from typing import Any, ClassVar, cast + +import pytest + +import tributo.integrations.broker as broker_contract +import tributo.plugin as plugin +from tributo.exceptions import JobConfigurationError +from tributo.integrations.broker import ( + BROKER_API_VERSION, + BrokerError, + BrokerPlugin, + BrokerRuntime, + Message, + TaskConsumer, + TaskDisposition, + TaskOutcome, +) +from tributo.integrations.broker_registry import BrokerRegistry + + +class _EntryPoint: + def __init__(self, name: str, loaded: Any) -> None: + self.name = name + self.value = "tests:_Plugin" + self._loaded = loaded + + def load(self) -> Any: + if isinstance(self._loaded, Exception): + raise self._loaded + return self._loaded + + +class _Consumer(TaskConsumer): + def poll(self, timeout_ms: int = 5000) -> Message | None: + del timeout_ms + return None + + def ack(self, message: Message) -> None: + del message + + +class _Runtime(BrokerRuntime): + consumer = _Consumer() + + def handle(self, message: Message) -> TaskOutcome: + del message + return TaskOutcome(TaskDisposition.ACK) + + +class _Plugin(BrokerPlugin): + api_version: ClassVar[int] = BROKER_API_VERSION + broker_id: ClassVar[str] = "fake" + capabilities: ClassVar[frozenset[str]] = frozenset({"task-consumer"}) + stability: ClassVar[str] = "alpha" + + def validate_config(self, config, *, check_connectivity=False) -> None: + del config, check_connectivity + + def create_runtime(self, config) -> _Runtime: + del config + return _Runtime() + + +def test_message_keeps_payload_opaque_and_metadata_restricted() -> None: + payload = object() + message = Message( + payload, + "delivery-1", + metadata={"attempt": "1"}, + ) + + assert message.payload is payload + assert message.metadata == {"attempt": "1"} + with pytest.raises(TypeError): + operator.setitem(message.metadata, "attempt", "2") + with pytest.raises(ValueError, match="string keys and values"): + Message( + {}, + "delivery-2", + metadata=cast(Any, {"attempt": 1}), + ) + + +def test_task_outcome_is_not_a_workload_result_contract() -> None: + outcome = TaskOutcome( + TaskDisposition.RETRY, + BrokerError(code="RAY_UNAVAILABLE", sanitized_message="retry later"), + ) + + assert outcome.error is not None + assert outcome.error.code == "RAY_UNAVAILABLE" + assert not hasattr(outcome, "result") + assert not hasattr(broker_contract, "JobResult") + assert not hasattr(broker_contract, "EventReporter") + + +def test_discovery_is_lazy_and_records_import_diagnostics(monkeypatch) -> None: + monkeypatch.setattr( + plugin, + "_iter_entry_points", + lambda group: iter( + [ + _EntryPoint("broken", ImportError("optional dependency unavailable")), + _EntryPoint("fake", _Plugin), + ] + if group == "tributo.brokers" + else [] + ), + ) + diagnostics = [] + + assert plugin.discover_broker_plugins(diagnostics) == [_Plugin] + assert diagnostics[0].entry_point_name == "broken" + assert diagnostics[0].error_type == "ImportError" + assert "optional dependency unavailable" not in diagnostics[0].reason + + +@pytest.mark.parametrize( + ("attribute", "value"), + [ + ("capabilities", ("task-consumer",)), + ("stability", "prototype"), + ], +) +def test_discovery_rejects_invalid_provider_metadata( + monkeypatch, attribute: str, value: object +) -> None: + invalid = type("InvalidPlugin", (_Plugin,), {attribute: value}) + monkeypatch.setattr( + plugin, + "_iter_entry_points", + lambda group: ( + iter([_EntryPoint("fake", invalid)]) + if group == "tributo.brokers" + else iter(()) + ), + ) + diagnostics = [] + + assert plugin.discover_broker_plugins(diagnostics) == [] + assert attribute in diagnostics[0].reason + + +def test_discovery_rejects_version_and_entrypoint_identity_mismatch( + monkeypatch, +) -> None: + class _WrongVersion(_Plugin): + api_version = BROKER_API_VERSION + 1 + + class _WrongIdentity(_Plugin): + broker_id = "other" + + monkeypatch.setattr( + plugin, + "_iter_entry_points", + lambda group: iter( + [ + _EntryPoint("wrong-version", _WrongVersion), + _EntryPoint("fake", _WrongIdentity), + ] + if group == "tributo.brokers" + else [] + ), + ) + diagnostics = [] + + assert plugin.discover_broker_plugins(diagnostics) == [] + assert "api_version" in diagnostics[0].reason + assert "does not match broker_id" in diagnostics[1].reason + + +def test_explicit_disabled_provider_fails_closed(monkeypatch) -> None: + monkeypatch.setenv("TRIBUTO_PLUGINS", "another") + monkeypatch.setattr( + plugin, + "_iter_entry_points", + lambda _group: iter([_EntryPoint("fake", _Plugin)]), + ) + + with pytest.raises(JobConfigurationError, match="disabled"): + plugin.resolve_broker_plugin("fake") + + +def test_registry_reports_metadata_and_duplicate_ids(monkeypatch) -> None: + monkeypatch.setattr( + "tributo.integrations.broker_registry.discover_broker_plugins", + lambda _diagnostics: [_Plugin, _Plugin], + ) + registry = BrokerRegistry() + + descriptors = registry.list() + + assert descriptors[0].broker_id == "fake" + assert descriptors[0].stability == "alpha" + assert descriptors[0].capabilities == ("task-consumer",) + assert registry.diagnostics()[0].reason == "Duplicate broker_id discovered" diff --git a/tests/test_broker_cli.py b/tests/test_broker_cli.py new file mode 100644 index 0000000..c72624d --- /dev/null +++ b/tests/test_broker_cli.py @@ -0,0 +1,74 @@ +"""Core broker CLI isolation tests.""" + +from __future__ import annotations + +import json +import os +import subprocess +import sys +from pathlib import Path + +from click.testing import CliRunner + +from tributo.cli import main + + +def test_broker_config_is_json_only_and_provider_owned(tmp_path) -> None: + path = tmp_path / "broker.json" + path.write_text(json.dumps({"opaque_provider_field": True}), encoding="utf-8") + result = CliRunner().invoke( + main, ["broker", "validate", "--broker", "missing", "--config", str(path)] + ) + + assert result.exit_code != 0 + assert "Unknown broker" in result.output + + +def test_normal_cli_does_not_require_a_broker_provider() -> None: + result = CliRunner().invoke(main, ["--help"]) + + assert result.exit_code == 0 + assert "broker" in result.output + + +def test_broker_list_without_provider_is_empty(monkeypatch) -> None: + monkeypatch.setattr("tributo.plugin._iter_entry_points", lambda _group: iter(())) + + result = CliRunner().invoke(main, ["broker", "list"]) + + assert result.exit_code == 0 + assert result.output == "" + + +def test_core_cli_has_no_provider_consume_loop() -> None: + result = CliRunner().invoke(main, ["broker", "--help"]) + + assert result.exit_code == 0 + assert "validate" in result.output + assert "consume" not in result.output + + +def test_import_and_root_help_do_not_import_broker_module() -> None: + env = dict(os.environ) + env["PYTHONPATH"] = os.pathsep.join( + [str(Path(__file__).parents[1] / "src"), env.get("PYTHONPATH", "")] + ) + script = """ +import sys +from click.testing import CliRunner +import tributo.cli as cli +assert 'tributo.cli_broker' not in sys.modules +result = CliRunner().invoke(cli.main, ['--help']) +assert result.exit_code == 0, result.output +assert 'broker' in result.output +assert 'tributo.cli_broker' not in sys.modules +""" + result = subprocess.run( + [sys.executable, "-c", script], + env=env, + check=False, + capture_output=True, + text=True, + ) + + assert result.returncode == 0, result.stderr or result.stdout diff --git a/tests/test_ray_jobs.py b/tests/test_ray_jobs.py new file mode 100644 index 0000000..7c180bb --- /dev/null +++ b/tests/test_ray_jobs.py @@ -0,0 +1,124 @@ +"""Workload-neutral Ray Jobs submission identity tests.""" + +from __future__ import annotations + +from typing import Any +from unittest.mock import MagicMock, patch + +import pytest +from ray.job_submission import JobStatus + +from tributo.ray_jobs import ( + RayJobSubmission, + get_ray_job_logs, + get_ray_job_status, + stop_ray_job, + submit_ray_job, +) + + +def _runtime_env(*args: Any, **kwargs: Any) -> dict[str, Any]: + del args + return {"env_vars": kwargs.get("env_vars", {})} + + +def test_submission_identity_is_workload_neutral_and_ray_job_id_is_real() -> None: + client = MagicMock() + client.submit_job.return_value = "ray-api-return-value" + client.get_job_info.return_value = type( + "JobInfo", (), {"job_id": "ray-core-job-1"} + )() + + with ( + patch("tributo.ray_jobs._get_submission_client", return_value=client), + patch("tributo.ray_jobs.build_runtime_env", side_effect=_runtime_env), + ): + result = submit_ray_job( + "python -m provider.driver", + operation_namespace="broker", + run_id="run-1", + attempt_id="attempt-2", + ) + + assert isinstance(result, RayJobSubmission) + assert result.run_id == "run-1" + assert result.attempt_id == "attempt-2" + assert result.submission_id.startswith("tributo-broker-") + assert result.ray_job_id == "ray-core-job-1" + assert client.submit_job.call_args.kwargs["submission_id"] == result.submission_id + + +def test_ambiguous_submission_reconciles_by_submission_id() -> None: + client = MagicMock() + client.submit_job.side_effect = TimeoutError("response lost") + client.get_job_status.return_value = JobStatus.RUNNING + client.get_job_info.side_effect = LookupError("driver not created") + + with ( + patch("tributo.ray_jobs._get_submission_client", return_value=client), + patch("tributo.ray_jobs.build_runtime_env", side_effect=_runtime_env), + ): + result = submit_ray_job( + "python -m provider.driver", + operation_namespace="broker", + run_id="run-1", + ) + + client.get_job_status.assert_called_once_with(result.submission_id) + assert result.ray_job_id is None + + +def test_request_digest_is_optional_metadata_not_submission_identity() -> None: + client = MagicMock() + client.get_job_info.return_value = type("JobInfo", (), {"job_id": None})() + + with ( + patch("tributo.ray_jobs._get_submission_client", return_value=client), + patch("tributo.ray_jobs.build_runtime_env", side_effect=_runtime_env), + ): + first = submit_ray_job( + "python -m provider.driver", + operation_namespace="broker", + run_id="run-1", + request_digest="digest-a", + ) + second = submit_ray_job( + "python -m provider.driver", + operation_namespace="broker", + run_id="run-1", + request_digest="digest-b", + ) + + assert first.submission_id == second.submission_id + assert client.submit_job.call_args_list[0].kwargs["metadata"] == { + "tributo.request_digest": "digest-a" + } + assert client.submit_job.call_args_list[1].kwargs["metadata"] == { + "tributo.request_digest": "digest-b" + } + + +def test_reserved_identity_environment_is_rejected_before_submission() -> None: + with pytest.raises(ValueError, match="must not override Ray Job identity"): + submit_ray_job( + "python -m provider.driver", + operation_namespace="broker", + run_id="run-1", + env_vars={"TRIBUTO_SUBMISSION_ID": "external"}, + ) + + +def test_status_and_stop_use_submission_identity() -> None: + client = MagicMock() + client.get_job_status.return_value = JobStatus.RUNNING + client.get_job_logs.return_value = "driver logs" + client.stop_job.return_value = True + + with patch("tributo.ray_jobs._get_submission_client", return_value=client): + assert get_ray_job_status("submission-1") == "RUNNING" + assert get_ray_job_logs("submission-1") == "driver logs" + assert stop_ray_job("submission-1") is True + + client.get_job_status.assert_called_once_with("submission-1") + client.get_job_logs.assert_called_once_with("submission-1") + client.stop_job.assert_called_once_with("submission-1") diff --git a/tests/test_retry.py b/tests/test_retry.py index d36bbbf..81b066d 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -7,6 +7,7 @@ import pytest +import tributo.training.job_submitter as training_job_submitter from tributo._common.retry import retry_with_exponential_backoff @@ -89,19 +90,17 @@ def test_training_job_submitter_retries_connection(self): client_mock = MagicMock() client_mock.submit_job.return_value = "job-456" - import tributo.training.job_submitter as tjs - with patch.object( - tjs, + training_job_submitter, "JobSubmissionClient", side_effect=[ ConnectionError("timeout"), client_mock, ], ) as mock_jsc: - job_id = tjs.submit_training_job("python train.py") + job_id = training_job_submitter.submit_training_job("python train.py") - assert job_id == "job-456" + assert job_id.startswith("tributo-train-") assert mock_jsc.call_count == 2 diff --git a/tests/test_runtime_env.py b/tests/test_runtime_env.py index 5b47fd8..61426be 100644 --- a/tests/test_runtime_env.py +++ b/tests/test_runtime_env.py @@ -48,3 +48,13 @@ def test_runtime_env_debug_log_never_exposes_environment_values( assert "PYTHONPATH" not in runtime_env["env_vars"] assert profile_payload not in caplog.text assert "TRIBUTO_STORAGE_PROFILE_MODEL" in caplog.text + + +def test_default_runtime_env_does_not_add_extension_dependencies(tmp_path) -> None: + (tmp_path / "pyproject.toml").write_text( + "[project]\nname='test'\n", encoding="utf-8" + ) + (tmp_path / "tributo").mkdir() + runtime_env = build_runtime_env(project_root=tmp_path) + assert runtime_env["py_modules"] == [str(tmp_path / "tributo")] + assert "pip" not in runtime_env diff --git a/tests/test_stability_inventory.py b/tests/test_stability_inventory.py index 19da5ef..979dfd8 100644 --- a/tests/test_stability_inventory.py +++ b/tests/test_stability_inventory.py @@ -30,6 +30,8 @@ "tributo.config": "stable", "tributo.job": "stable", "tributo.exceptions": "stable", + # Core — alpha + "tributo.ray_jobs": "alpha", # Core — beta "tributo.cli": "beta", # Portable algorithm execution — alpha @@ -155,6 +157,8 @@ "tributo.integrations.hooks": "beta", "tributo.integrations.sinks.parquet": "alpha", "tributo.integrations.sinks.lance": "alpha", + "tributo.integrations.broker": "alpha", + "tributo.integrations.broker_registry": "alpha", # Inference — beta "tributo.inference.base": "beta", "tributo.inference.batch_predictor": "beta", @@ -225,9 +229,13 @@ "tributo.inference.job_runner.InferenceJobResult": "alpha", "tributo.inference.job_runner.map_ray_job_status": "alpha", "tributo.inference.job_runner.submit_inference_request": "alpha", + "tributo.inference.job_runner.submit_inference_request_with_identity": "alpha", "tributo.inference.job_runner.submit_resolved_inference": "alpha", + "tributo.inference.job_runner.submit_resolved_inference_with_identity": "alpha", "tributo.inference.job_runner.submit_inference_request_with_retry": "alpha", "tributo.inference.job_runner.wait_for_job": "alpha", + "tributo.training.job_submitter.submit_training_job_with_identity": "alpha", + "tributo.training.xgboost_trainer.run_training_result_with_config": "alpha", } diff --git a/tests/training/test_job_submitter.py b/tests/training/test_job_submitter.py index 8de3f3e..e2f52ec 100644 --- a/tests/training/test_job_submitter.py +++ b/tests/training/test_job_submitter.py @@ -10,6 +10,7 @@ from tributo.algorithms import AlgorithmArtifact, EnvironmentSpec, ImageProfile from tributo.training.job_submitter import ( submit_training_job, + submit_training_job_with_identity, submit_training_job_with_retry, ) @@ -69,15 +70,15 @@ def test_same_run_and_attempt_reuse_submission_id(self) -> None: attempt_id="attempt-1", ) - assert first == second == "job-1" + assert first == second + assert first.startswith("tributo-train-") ids = [ call.kwargs["submission_id"] for call in client.submit_job.call_args_list ] assert ids[0] == ids[1] - assert ( - "TRIBUTO_RUN_ID" - in client.submit_job.call_args.kwargs["runtime_env"]["env_vars"] - ) + worker_env = client.submit_job.call_args.kwargs["runtime_env"]["env_vars"] + assert "TRIBUTO_RUN_ID" in worker_env + assert worker_env["TRIBUTO_SUBMISSION_ID"] == ids[-1] def test_existing_failed_attempt_is_reconciled_without_timestamp_retry( self, @@ -85,6 +86,9 @@ def test_existing_failed_attempt_is_reconciled_without_timestamp_retry( client = MagicMock() client.submit_job.side_effect = RuntimeError("submission already exists") client.get_job_status.return_value = JobStatus.FAILED + client.get_job_info.return_value = type( + "JobInfo", (), {"job_id": "ray-job-1"} + )() with ( patch( @@ -96,13 +100,14 @@ def test_existing_failed_attempt_is_reconciled_without_timestamp_retry( side_effect=_runtime_env, ), ): - job_id = submit_training_job( + submission = submit_training_job_with_identity( "python train.py", run_id="run-1", attempt_id="attempt-1", ) - assert job_id.startswith("tributo-train-") + assert submission.submission_id.startswith("tributo-train-") + assert submission.ray_job_id == "ray-job-1" assert client.submit_job.call_count == 1 @@ -138,6 +143,13 @@ def test_failed_job_uses_next_attempt_and_succeeded_stops(self) -> None: assert result.attempts[1].attempt_id == "attempt-2" assert result.attempts[1].status == "SUCCEEDED" assert result.attempts[0].submission_id != result.attempts[1].submission_id + for call, attempt in zip( + client.submit_job.call_args_list, result.attempts, strict=True + ): + assert ( + call.kwargs["runtime_env"]["env_vars"]["TRIBUTO_SUBMISSION_ID"] + == attempt.submission_id + ) def test_stopped_job_is_never_retried(self) -> None: client = MagicMock() @@ -259,9 +271,51 @@ def test_algorithm_artifact_and_environment_dependencies_reach_preflight( ), ) - assert job_id == "job-artifact" + assert job_id.startswith("tributo-train-") assert build_runtime_env.call_args.kwargs["algorithm_artifact"] is artifact assert build_runtime_env.call_args.kwargs["image_profile"] is profile assert build_runtime_env.call_args.kwargs["declared_dependencies"] == ( "scikit-learn<2,>=1.6", ) + + def test_metadata_cannot_override_submission_identity(self) -> None: + with pytest.raises(ValueError, match="TRIBUTO_SUBMISSION_ID"): + submit_training_job( + "python train.py", + run_id="run-1", + metadata={"TRIBUTO_SUBMISSION_ID": "other-submission"}, + ) + + def test_submission_result_preserves_identity_and_optional_digest(self) -> None: + client = MagicMock() + client.submit_job.return_value = "submission-return" + client.get_job_info.return_value = type( + "JobInfo", (), {"job_id": "ray-job-1"} + )() + with ( + patch( + "tributo.training.job_submitter._get_submission_client", + return_value=client, + ), + patch( + "tributo.training.job_submitter.build_runtime_env", + side_effect=_runtime_env, + ), + ): + result = submit_training_job_with_identity( + "python -m worker", + run_id="business-job-1", + attempt_id="attempt-2", + request_digest="request-digest", + ) + assert result.run_id == "business-job-1" + assert result.attempt_id == "attempt-2" + assert result.submission_id.startswith("tributo-train-") + assert result.ray_job_id == "ray-job-1" + assert result.request_digest == "request-digest" + worker_env = client.submit_job.call_args.kwargs["runtime_env"]["env_vars"] + assert worker_env["TRIBUTO_SUBMISSION_ID"] == result.submission_id + assert "TRIBUTO_EXECUTION_CONTEXT" not in worker_env + assert client.submit_job.call_args.kwargs["metadata"] == { + "tributo.request_digest": "request-digest" + } diff --git a/tests/training/test_xgboost_trainer_unit.py b/tests/training/test_xgboost_trainer_unit.py index 6ba7987..e497c19 100644 --- a/tests/training/test_xgboost_trainer_unit.py +++ b/tests/training/test_xgboost_trainer_unit.py @@ -23,6 +23,7 @@ _managed_resume_checkpoint, _merge_xgb_eval_results, _populate_xgb_eval_metrics, + run_training_result_with_config, ) @@ -69,6 +70,29 @@ def test_managed_resume_checkpoint_cleans_directory_on_failure(tmp_path: Path) - assert not checkpoint_dir.exists() +def test_in_process_training_entrypoint_returns_training_result(monkeypatch) -> None: + monkeypatch.setattr( + "tributo.training.xgboost_trainer.run_training_with_config", + lambda _config: { + "model_uri": "file:///tmp/bundle", + "bundle_uri": "file:///tmp/bundle", + "metrics": {"accuracy": 0.9}, + "legacy_artifact_uri": None, + "training_status": "succeeded", + "bundle_status": "succeeded", + "hook_status": "not_configured", + "execution_id": "execution-1", + "status": "succeeded", + }, + ) + + result = run_training_result_with_config({}) + + assert result.training_status == "succeeded" + assert result.bundle_uri == "file:///tmp/bundle" + assert result.execution_id == "execution-1" + + class TestS3Config: """S3Config Pydantic 模型测试。""" diff --git a/tools/generate_public_api_reference.py b/tools/generate_public_api_reference.py index b7cb193..d0341ca 100644 --- a/tools/generate_public_api_reference.py +++ b/tools/generate_public_api_reference.py @@ -107,6 +107,19 @@ def _decorator_stability(decorator: ast.expr) -> str | None: ) +def _is_exception_class(node: ast.ClassDef) -> bool: + """Classify exception types from their bases, not their domain name.""" + base_names = [ + base.id + if isinstance(base, ast.Name) + else base.attr + if isinstance(base, ast.Attribute) + else "" + for base in node.bases + ] + return any(name.endswith(("Error", "Exception")) for name in base_names) + + def build_inventory(source_root: Path = SOURCE_ROOT) -> tuple[PublicSymbol, ...]: """Return every top-level source object annotated with ``@PublicAPI``.""" symbols: list[PublicSymbol] = [] @@ -135,11 +148,7 @@ def build_inventory(source_root: Path = SOURCE_ROOT) -> tuple[PublicSymbol, ...] f"{path}:{node.lineno}: unsupported stability {stability!r}" ) if isinstance(node, ast.ClassDef): - kind = ( - "exception" - if node.name.endswith(("Error", "Exception")) - else "class" - ) + kind = "exception" if _is_exception_class(node) else "class" else: kind = "function" symbols.append( @@ -164,7 +173,7 @@ def component_for(symbol: PublicSymbol) -> str: """Route a public symbol to one user-facing component page.""" parts = symbol.module.split(".") package = parts[1] if len(parts) > 1 else "core" - if package in {"config", "exceptions", "job", "_common"}: + if package in {"config", "exceptions", "job", "ray_jobs", "_common"}: return "core" if package in {"data", "streaming"}: return "data"