diff --git a/modelexpress_client/python/generate_proto.sh b/modelexpress_client/python/generate_proto.sh index d01257821..30c032ab7 100755 --- a/modelexpress_client/python/generate_proto.sh +++ b/modelexpress_client/python/generate_proto.sh @@ -13,27 +13,36 @@ SPDX_HEADER="# SPDX-FileCopyrightText: Copyright (c) 2025-${YEAR} NVIDIA CORPORA # SPDX-License-Identifier: Apache-2.0 #" +PROTO_NAMES=(p2p revision) +PROTO_FILES=() +for name in "${PROTO_NAMES[@]}"; do + PROTO_FILES+=("${PROTO_DIR}/${name}.proto") +done + # Generate protobuf files -echo "Generating protobuf files from ${PROTO_DIR}/p2p.proto..." +echo "Generating protobuf files from ${PROTO_NAMES[*]}..." python -m grpc_tools.protoc \ "-I${PROTO_DIR}" \ "--python_out=${OUT_DIR}" \ "--grpc_python_out=${OUT_DIR}" \ - "${PROTO_DIR}/p2p.proto" - -# Fix relative import in grpc file -echo "Fixing imports in p2p_pb2_grpc.py..." -tmp_file="$(mktemp)" -sed 's/^import p2p_pb2 as/from . import p2p_pb2 as/' "${OUT_DIR}/p2p_pb2_grpc.py" > "${tmp_file}" -mv "${tmp_file}" "${OUT_DIR}/p2p_pb2_grpc.py" - -# Add SPDX header to generated files -for file in "${OUT_DIR}/p2p_pb2.py" "${OUT_DIR}/p2p_pb2_grpc.py"; do - echo "Adding SPDX header to ${file}..." - tmp_file=$(mktemp) - echo "${SPDX_HEADER}" > "${tmp_file}" - cat "${file}" >> "${tmp_file}" - mv "${tmp_file}" "${file}" + "${PROTO_FILES[@]}" + +for name in "${PROTO_NAMES[@]}"; do + grpc_file="${OUT_DIR}/${name}_pb2_grpc.py" + echo "Fixing imports in ${name}_pb2_grpc.py..." + tmp_file="$(mktemp)" + sed -E 's/^import ([a-zA-Z0-9_]+_pb2) as/from . import \1 as/' "${grpc_file}" > "${tmp_file}" + mv "${tmp_file}" "${grpc_file}" + sed -i "s/+ f' but the generated code/+ ' but the generated code/" "${grpc_file}" + sed -i '/^import warnings$/d' "${grpc_file}" + + for file in "${OUT_DIR}/${name}_pb2.py" "${grpc_file}"; do + echo "Adding SPDX header to ${file}..." + tmp_file="$(mktemp)" + printf '%s\n' "${SPDX_HEADER}" > "${tmp_file}" + cat "${file}" >> "${tmp_file}" + mv "${tmp_file}" "${file}" + done done echo "Done." diff --git a/modelexpress_client/python/modelexpress/p2p_pb2_grpc.py b/modelexpress_client/python/modelexpress/p2p_pb2_grpc.py index 5f8ee047d..65cfc697c 100644 --- a/modelexpress_client/python/modelexpress/p2p_pb2_grpc.py +++ b/modelexpress_client/python/modelexpress/p2p_pb2_grpc.py @@ -4,7 +4,6 @@ # Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! """Client and server classes corresponding to protobuf-defined services.""" import grpc -import warnings from . import p2p_pb2 as p2p__pb2 @@ -21,7 +20,7 @@ if _version_not_supported: raise RuntimeError( f'The grpc package installed is at version {GRPC_VERSION},' - + f' but the generated code in p2p_pb2_grpc.py depends on' + + ' but the generated code in p2p_pb2_grpc.py depends on' + f' grpcio>={GRPC_GENERATED_VERSION}.' + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' diff --git a/modelexpress_client/python/modelexpress/refit/catalog.py b/modelexpress_client/python/modelexpress/refit/catalog.py new file mode 100644 index 000000000..03d44cb01 --- /dev/null +++ b/modelexpress_client/python/modelexpress/refit/catalog.py @@ -0,0 +1,99 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Typed boundary over the three minimal revision-catalog RPCs.""" + +from __future__ import annotations + +import math +from typing import Protocol, runtime_checkable + +import grpc + +from modelexpress import revision_pb2, revision_pb2_grpc + +from .manifest import RevisionManifest, RevisionRecord + + +@runtime_checkable +class RevisionCatalog(Protocol): + """Exact metadata operations available to the publisher and orchestrator.""" + + def publish_revision(self, manifest: RevisionManifest) -> RevisionRecord: ... + + def get_revision( + self, model_id: str, target_version: str + ) -> RevisionRecord: ... + + def commit_revision( + self, model_id: str, target_version: str + ) -> RevisionRecord: ... + + +class GrpcRevisionCatalog: + """Concrete :class:`RevisionCatalog` over the generated gRPC service.""" + + def __init__( + self, + endpoint: str | None = None, + stub=None, + timeout: float = 10.0, + ) -> None: + if not math.isfinite(timeout) or timeout <= 0: + raise ValueError("catalog RPC timeout must be finite and positive") + if (endpoint is None) == (stub is None): + raise ValueError( + "GrpcRevisionCatalog needs exactly one of endpoint or stub" + ) + self._channel = None + if stub is None: + assert endpoint is not None + if endpoint.startswith("https://"): + target = endpoint.removeprefix("https://") + self._channel = grpc.secure_channel( + target, grpc.ssl_channel_credentials() + ) + else: + target = endpoint.removeprefix("http://") + self._channel = grpc.insecure_channel(target) + stub = revision_pb2_grpc.RevisionCatalogServiceStub(self._channel) + self._stub = stub + self._timeout = timeout + + def __enter__(self) -> GrpcRevisionCatalog: + return self + + def __exit__(self, *_exc_info) -> None: + self.close() + + def close(self) -> None: + if self._channel is not None: + self._channel.close() + self._channel = None + + def publish_revision(self, manifest: RevisionManifest) -> RevisionRecord: + response = self._stub.PublishRevision( + revision_pb2.PublishRevisionRequest(manifest=manifest.to_proto()), + timeout=self._timeout, + ) + return RevisionRecord.from_proto(response) + + def get_revision(self, model_id: str, target_version: str) -> RevisionRecord: + response = self._stub.GetRevision( + revision_pb2.GetRevisionRequest( + model_id=model_id, + target_version=target_version, + ), + timeout=self._timeout, + ) + return RevisionRecord.from_proto(response) + + def commit_revision(self, model_id: str, target_version: str) -> RevisionRecord: + response = self._stub.CommitRevision( + revision_pb2.CommitRevisionRequest( + model_id=model_id, + target_version=target_version, + ), + timeout=self._timeout, + ) + return RevisionRecord.from_proto(response) diff --git a/modelexpress_client/python/modelexpress/refit/manifest.py b/modelexpress_client/python/modelexpress/refit/manifest.py new file mode 100644 index 000000000..7bcd0cef3 --- /dev/null +++ b/modelexpress_client/python/modelexpress/refit/manifest.py @@ -0,0 +1,105 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Minimal revision-catalog DTOs and their exact protobuf mapping.""" + +from __future__ import annotations + +from dataclasses import dataclass +from enum import IntEnum + +from modelexpress import revision_pb2 + + +class RevisionState(IntEnum): + UNSPECIFIED = revision_pb2.REVISION_STATE_UNSPECIFIED + READY = revision_pb2.REVISION_STATE_READY + COMMITTED = revision_pb2.REVISION_STATE_COMMITTED + + +@dataclass(frozen=True) +class S3Object: + bucket: str + key: str + checksum: str + object_version: str | None = None + + def to_proto(self) -> revision_pb2.S3Object: + fields = { + "bucket": self.bucket, + "key": self.key, + "checksum": self.checksum, + } + if self.object_version is not None: + fields["object_version"] = self.object_version + return revision_pb2.S3Object(**fields) + + @classmethod + def from_proto(cls, proto: revision_pb2.S3Object) -> S3Object: + return cls( + bucket=proto.bucket, + key=proto.key, + checksum=proto.checksum, + object_version=( + proto.object_version if proto.HasField("object_version") else None + ), + ) + + +@dataclass(frozen=True) +class RevisionManifest: + model_id: str + target_version: str + target_digest: str + format_digest: str + base_version: str | None = None + base_digest: str | None = None + payload: S3Object | None = None + + def to_proto(self) -> revision_pb2.RevisionManifest: + fields = { + "model_id": self.model_id, + "target_version": self.target_version, + "target_digest": self.target_digest, + "format_digest": self.format_digest, + } + if self.base_version is not None: + fields["base_version"] = self.base_version + if self.base_digest is not None: + fields["base_digest"] = self.base_digest + if self.payload is not None: + fields["payload"] = self.payload.to_proto() + return revision_pb2.RevisionManifest(**fields) + + @classmethod + def from_proto(cls, proto: revision_pb2.RevisionManifest) -> RevisionManifest: + return cls( + model_id=proto.model_id, + target_version=proto.target_version, + target_digest=proto.target_digest, + format_digest=proto.format_digest, + base_version=(proto.base_version if proto.HasField("base_version") else None), + base_digest=(proto.base_digest if proto.HasField("base_digest") else None), + payload=(S3Object.from_proto(proto.payload) if proto.HasField("payload") else None), + ) + + +@dataclass(frozen=True) +class RevisionRecord: + manifest: RevisionManifest + state: RevisionState + + def to_proto(self) -> revision_pb2.RevisionRecord: + return revision_pb2.RevisionRecord( + manifest=self.manifest.to_proto(), + state=int(self.state), + ) + + @classmethod + def from_proto(cls, proto: revision_pb2.RevisionRecord) -> RevisionRecord: + if not proto.HasField("manifest"): + raise ValueError("revision record is missing manifest") + return cls( + manifest=RevisionManifest.from_proto(proto.manifest), + state=RevisionState(proto.state), + ) diff --git a/modelexpress_client/python/modelexpress/revision_pb2.py b/modelexpress_client/python/modelexpress/revision_pb2.py new file mode 100644 index 000000000..9a4721228 --- /dev/null +++ b/modelexpress_client/python/modelexpress/revision_pb2.py @@ -0,0 +1,53 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# -*- coding: utf-8 -*- +# Generated by the protocol buffer compiler. DO NOT EDIT! +# NO CHECKED-IN PROTOBUF GENCODE +# source: revision.proto +# Protobuf Python Version: 5.27.2 +"""Generated protocol buffer code.""" +from google.protobuf import descriptor as _descriptor +from google.protobuf import descriptor_pool as _descriptor_pool +from google.protobuf import runtime_version as _runtime_version +from google.protobuf import symbol_database as _symbol_database +from google.protobuf.internal import builder as _builder +_runtime_version.ValidateProtobufRuntimeVersion( + _runtime_version.Domain.PUBLIC, + 5, + 27, + 2, + '', + 'revision.proto' +) +# @@protoc_insertion_point(imports) + +_sym_db = _symbol_database.Default() + + + + +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n\x0erevision.proto\x12\x16model_express.revision\"i\n\x08S3Object\x12\x0e\n\x06\x62ucket\x18\x01 \x01(\t\x12\x0b\n\x03key\x18\x02 \x01(\t\x12\x1b\n\x0eobject_version\x18\x03 \x01(\tH\x00\x88\x01\x01\x12\x10\n\x08\x63hecksum\x18\x04 \x01(\tB\x11\n\x0f_object_version\"\xf3\x01\n\x10RevisionManifest\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x16\n\x0etarget_version\x18\x02 \x01(\t\x12\x19\n\x0c\x62\x61se_version\x18\x03 \x01(\tH\x00\x88\x01\x01\x12\x18\n\x0b\x62\x61se_digest\x18\x04 \x01(\tH\x01\x88\x01\x01\x12\x15\n\rtarget_digest\x18\x05 \x01(\t\x12\x15\n\rformat_digest\x18\x06 \x01(\t\x12\x31\n\x07payload\x18\x07 \x01(\x0b\x32 .model_express.revision.S3ObjectB\x0f\n\r_base_versionB\x0e\n\x0c_base_digest\"\x82\x01\n\x0eRevisionRecord\x12:\n\x08manifest\x18\x01 \x01(\x0b\x32(.model_express.revision.RevisionManifest\x12\x34\n\x05state\x18\x02 \x01(\x0e\x32%.model_express.revision.RevisionState\"T\n\x16PublishRevisionRequest\x12:\n\x08manifest\x18\x01 \x01(\x0b\x32(.model_express.revision.RevisionManifest\">\n\x12GetRevisionRequest\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x16\n\x0etarget_version\x18\x02 \x01(\t\"A\n\x15\x43ommitRevisionRequest\x12\x10\n\x08model_id\x18\x01 \x01(\t\x12\x16\n\x0etarget_version\x18\x02 \x01(\t*g\n\rRevisionState\x12\x1e\n\x1aREVISION_STATE_UNSPECIFIED\x10\x00\x12\x18\n\x14REVISION_STATE_READY\x10\x01\x12\x1c\n\x18REVISION_STATE_COMMITTED\x10\x02\x32\xcf\x02\n\x16RevisionCatalogService\x12i\n\x0fPublishRevision\x12..model_express.revision.PublishRevisionRequest\x1a&.model_express.revision.RevisionRecord\x12\x61\n\x0bGetRevision\x12*.model_express.revision.GetRevisionRequest\x1a&.model_express.revision.RevisionRecord\x12g\n\x0e\x43ommitRevision\x12-.model_express.revision.CommitRevisionRequest\x1a&.model_express.revision.RevisionRecordb\x06proto3') + +_globals = globals() +_builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) +_builder.BuildTopDescriptorsAndMessages(DESCRIPTOR, 'revision_pb2', _globals) +if not _descriptor._USE_C_DESCRIPTORS: + DESCRIPTOR._loaded_options = None + _globals['_REVISIONSTATE']._serialized_start=745 + _globals['_REVISIONSTATE']._serialized_end=848 + _globals['_S3OBJECT']._serialized_start=42 + _globals['_S3OBJECT']._serialized_end=147 + _globals['_REVISIONMANIFEST']._serialized_start=150 + _globals['_REVISIONMANIFEST']._serialized_end=393 + _globals['_REVISIONRECORD']._serialized_start=396 + _globals['_REVISIONRECORD']._serialized_end=526 + _globals['_PUBLISHREVISIONREQUEST']._serialized_start=528 + _globals['_PUBLISHREVISIONREQUEST']._serialized_end=612 + _globals['_GETREVISIONREQUEST']._serialized_start=614 + _globals['_GETREVISIONREQUEST']._serialized_end=676 + _globals['_COMMITREVISIONREQUEST']._serialized_start=678 + _globals['_COMMITREVISIONREQUEST']._serialized_end=743 + _globals['_REVISIONCATALOGSERVICE']._serialized_start=851 + _globals['_REVISIONCATALOGSERVICE']._serialized_end=1186 +# @@protoc_insertion_point(module_scope) diff --git a/modelexpress_client/python/modelexpress/revision_pb2_grpc.py b/modelexpress_client/python/modelexpress/revision_pb2_grpc.py new file mode 100644 index 000000000..ee63a4071 --- /dev/null +++ b/modelexpress_client/python/modelexpress/revision_pb2_grpc.py @@ -0,0 +1,194 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Generated by the gRPC Python protocol compiler plugin. DO NOT EDIT! +"""Client and server classes corresponding to protobuf-defined services.""" +import grpc + +from . import revision_pb2 as revision__pb2 + +GRPC_GENERATED_VERSION = '1.66.2' +GRPC_VERSION = grpc.__version__ +_version_not_supported = False + +try: + from grpc._utilities import first_version_is_lower + _version_not_supported = first_version_is_lower(GRPC_VERSION, GRPC_GENERATED_VERSION) +except ImportError: + _version_not_supported = True + +if _version_not_supported: + raise RuntimeError( + f'The grpc package installed is at version {GRPC_VERSION},' + + ' but the generated code in revision_pb2_grpc.py depends on' + + f' grpcio>={GRPC_GENERATED_VERSION}.' + + f' Please upgrade your grpc module to grpcio>={GRPC_GENERATED_VERSION}' + + f' or downgrade your generated code using grpcio-tools<={GRPC_VERSION}.' + ) + + +class RevisionCatalogServiceStub(object): + """Stores immutable model-weight revision metadata. Payload bytes never pass + through this service. + """ + + def __init__(self, channel): + """Constructor. + + Args: + channel: A grpc.Channel. + """ + self.PublishRevision = channel.unary_unary( + '/model_express.revision.RevisionCatalogService/PublishRevision', + request_serializer=revision__pb2.PublishRevisionRequest.SerializeToString, + response_deserializer=revision__pb2.RevisionRecord.FromString, + _registered_method=True) + self.GetRevision = channel.unary_unary( + '/model_express.revision.RevisionCatalogService/GetRevision', + request_serializer=revision__pb2.GetRevisionRequest.SerializeToString, + response_deserializer=revision__pb2.RevisionRecord.FromString, + _registered_method=True) + self.CommitRevision = channel.unary_unary( + '/model_express.revision.RevisionCatalogService/CommitRevision', + request_serializer=revision__pb2.CommitRevisionRequest.SerializeToString, + response_deserializer=revision__pb2.RevisionRecord.FromString, + _registered_method=True) + + +class RevisionCatalogServiceServicer(object): + """Stores immutable model-weight revision metadata. Payload bytes never pass + through this service. + """ + + def PublishRevision(self, request, context): + """Atomically publishes one immutable revision as READY. + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + + def GetRevision(self, request, context): + """Returns one exact revision by model and target version. + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + + def CommitRevision(self, request, context): + """Records orchestrator acceptance after the required rollout cohort verifies. + """ + context.set_code(grpc.StatusCode.UNIMPLEMENTED) + context.set_details('Method not implemented!') + raise NotImplementedError('Method not implemented!') + + +def add_RevisionCatalogServiceServicer_to_server(servicer, server): + rpc_method_handlers = { + 'PublishRevision': grpc.unary_unary_rpc_method_handler( + servicer.PublishRevision, + request_deserializer=revision__pb2.PublishRevisionRequest.FromString, + response_serializer=revision__pb2.RevisionRecord.SerializeToString, + ), + 'GetRevision': grpc.unary_unary_rpc_method_handler( + servicer.GetRevision, + request_deserializer=revision__pb2.GetRevisionRequest.FromString, + response_serializer=revision__pb2.RevisionRecord.SerializeToString, + ), + 'CommitRevision': grpc.unary_unary_rpc_method_handler( + servicer.CommitRevision, + request_deserializer=revision__pb2.CommitRevisionRequest.FromString, + response_serializer=revision__pb2.RevisionRecord.SerializeToString, + ), + } + generic_handler = grpc.method_handlers_generic_handler( + 'model_express.revision.RevisionCatalogService', rpc_method_handlers) + server.add_generic_rpc_handlers((generic_handler,)) + server.add_registered_method_handlers('model_express.revision.RevisionCatalogService', rpc_method_handlers) + + + # This class is part of an EXPERIMENTAL API. +class RevisionCatalogService(object): + """Stores immutable model-weight revision metadata. Payload bytes never pass + through this service. + """ + + @staticmethod + def PublishRevision(request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None): + return grpc.experimental.unary_unary( + request, + target, + '/model_express.revision.RevisionCatalogService/PublishRevision', + revision__pb2.PublishRevisionRequest.SerializeToString, + revision__pb2.RevisionRecord.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @staticmethod + def GetRevision(request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None): + return grpc.experimental.unary_unary( + request, + target, + '/model_express.revision.RevisionCatalogService/GetRevision', + revision__pb2.GetRevisionRequest.SerializeToString, + revision__pb2.RevisionRecord.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) + + @staticmethod + def CommitRevision(request, + target, + options=(), + channel_credentials=None, + call_credentials=None, + insecure=False, + compression=None, + wait_for_ready=None, + timeout=None, + metadata=None): + return grpc.experimental.unary_unary( + request, + target, + '/model_express.revision.RevisionCatalogService/CommitRevision', + revision__pb2.CommitRevisionRequest.SerializeToString, + revision__pb2.RevisionRecord.FromString, + options, + channel_credentials, + insecure, + call_credentials, + compression, + wait_for_ready, + timeout, + metadata, + _registered_method=True) diff --git a/modelexpress_client/python/tests/test_revision_catalog.py b/modelexpress_client/python/tests/test_revision_catalog.py new file mode 100644 index 000000000..999c7dd89 --- /dev/null +++ b/modelexpress_client/python/tests/test_revision_catalog.py @@ -0,0 +1,216 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from concurrent.futures import ThreadPoolExecutor + +import grpc + +from modelexpress import revision_pb2, revision_pb2_grpc +from modelexpress.refit.catalog import GrpcRevisionCatalog, RevisionCatalog +from modelexpress.refit.manifest import ( + RevisionManifest, + RevisionRecord, + RevisionState, +) + + +class FakeStub: + def __init__(self, record: revision_pb2.RevisionRecord) -> None: + self.record = record + self.calls: list[tuple[str, object, dict[str, object]]] = [] + + def PublishRevision(self, request, **kwargs): + self.calls.append(("publish", request, kwargs)) + return self.record + + def GetRevision(self, request, **kwargs): + self.calls.append(("get", request, kwargs)) + return self.record + + def CommitRevision(self, request, **kwargs): + self.calls.append(("commit", request, kwargs)) + committed = revision_pb2.RevisionRecord() + committed.CopyFrom(self.record) + committed.state = revision_pb2.REVISION_STATE_COMMITTED + return committed + + +def launch_manifest() -> RevisionManifest: + return RevisionManifest( + model_id="model", + target_version="0", + target_digest="sha256:target-0", + format_digest="sha256:format", + ) + + +def catalog(timeout: float = 10.0) -> tuple[GrpcRevisionCatalog, FakeStub]: + record = RevisionRecord( + manifest=launch_manifest(), + state=RevisionState.READY, + ).to_proto() + stub = FakeStub(record) + return GrpcRevisionCatalog(stub=stub, timeout=timeout), stub + + +def test_protocol_has_only_three_revision_operations(): + assert { + name + for name, value in vars(RevisionCatalog).items() + if callable(value) and not name.startswith("_") + } == { + "publish_revision", + "get_revision", + "commit_revision", + } + + +def test_publish_sends_only_the_manifest_and_returns_direct_record(): + client, stub = catalog() + + result = client.publish_revision(launch_manifest()) + + assert result.state is RevisionState.READY + operation, request, kwargs = stub.calls[-1] + assert operation == "publish" + assert request == revision_pb2.PublishRevisionRequest( + manifest=launch_manifest().to_proto() + ) + assert kwargs == {"timeout": 10.0} + + +def test_get_and_commit_use_exact_model_and_target_version(): + client, stub = catalog() + + fetched = client.get_revision("model", "0") + committed = client.commit_revision("model", "0") + + assert fetched.state is RevisionState.READY + assert committed.state is RevisionState.COMMITTED + assert stub.calls[-2] == ( + "get", + revision_pb2.GetRevisionRequest(model_id="model", target_version="0"), + {"timeout": 10.0}, + ) + assert stub.calls[-1] == ( + "commit", + revision_pb2.CommitRevisionRequest(model_id="model", target_version="0"), + {"timeout": 10.0}, + ) + + +def test_configured_timeout_is_sent_to_every_rpc(): + client, stub = catalog(timeout=2.5) + + client.publish_revision(launch_manifest()) + client.get_revision("model", "0") + client.commit_revision("model", "0") + + assert [kwargs for _, _, kwargs in stub.calls] == [ + {"timeout": 2.5}, + {"timeout": 2.5}, + {"timeout": 2.5}, + ] + + +def test_timeout_must_be_finite_and_positive(): + for timeout in (0.0, -1.0, float("inf"), float("nan")): + try: + catalog(timeout=timeout) + except ValueError as error: + assert str(error) == "catalog RPC timeout must be finite and positive" + else: + raise AssertionError(f"accepted invalid timeout {timeout}") + + +def test_endpoint_scheme_selects_matching_grpc_transport(monkeypatch): + insecure_targets = [] + secure_targets = [] + channel = object() + monkeypatch.setattr( + grpc, + "insecure_channel", + lambda target: insecure_targets.append(target) or channel, + ) + monkeypatch.setattr( + grpc, + "secure_channel", + lambda target, credentials: secure_targets.append((target, credentials)) + or channel, + ) + credentials = object() + monkeypatch.setattr(grpc, "ssl_channel_credentials", lambda: credentials) + monkeypatch.setattr( + revision_pb2_grpc, + "RevisionCatalogServiceStub", + lambda selected_channel: selected_channel, + ) + + GrpcRevisionCatalog(endpoint="http://catalog:8001") + GrpcRevisionCatalog(endpoint="catalog:8001") + GrpcRevisionCatalog(endpoint="https://catalog:8443") + + assert insecure_targets == ["catalog:8001", "catalog:8001"] + assert secure_targets == [("catalog:8443", credentials)] + + +def test_client_has_no_deferred_catalog_operations(): + client, _ = catalog() + + assert not hasattr(client, "list_ready_revisions") + assert not hasattr(client, "get_recovery_candidates") + assert not hasattr(client, "update_receiver_state") + assert not hasattr(client, "commit_version") + + +class LoopbackCatalog(revision_pb2_grpc.RevisionCatalogServiceServicer): + def __init__(self) -> None: + self.records: dict[tuple[str, str], revision_pb2.RevisionRecord] = {} + + def PublishRevision(self, request, context): + manifest = request.manifest + key = (manifest.model_id, manifest.target_version) + existing = self.records.get(key) + if existing is not None: + if existing.manifest != manifest: + context.abort(grpc.StatusCode.ALREADY_EXISTS, "manifest conflict") + return existing + record = revision_pb2.RevisionRecord( + manifest=manifest, + state=revision_pb2.REVISION_STATE_READY, + ) + self.records[key] = record + return record + + def GetRevision(self, request, context): + record = self.records.get((request.model_id, request.target_version)) + if record is None: + context.abort(grpc.StatusCode.NOT_FOUND, "revision not found") + return record + + def CommitRevision(self, request, context): + key = (request.model_id, request.target_version) + record = self.records.get(key) + if record is None: + context.abort(grpc.StatusCode.NOT_FOUND, "revision not found") + record.state = revision_pb2.REVISION_STATE_COMMITTED + return record + + +def test_real_grpc_loopback_publish_get_and_commit(): + server = grpc.server(ThreadPoolExecutor(max_workers=1)) + revision_pb2_grpc.add_RevisionCatalogServiceServicer_to_server( + LoopbackCatalog(), server + ) + port = server.add_insecure_port("127.0.0.1:0") + server.start() + try: + with GrpcRevisionCatalog(f"127.0.0.1:{port}") as client: + published = client.publish_revision(launch_manifest()) + fetched = client.get_revision("model", "0") + committed = client.commit_revision("model", "0") + assert published.state is RevisionState.READY + assert fetched == published + assert committed.state is RevisionState.COMMITTED + finally: + server.stop(grace=None).wait() diff --git a/modelexpress_client/python/tests/test_revision_manifest.py b/modelexpress_client/python/tests/test_revision_manifest.py new file mode 100644 index 000000000..34142821b --- /dev/null +++ b/modelexpress_client/python/tests/test_revision_manifest.py @@ -0,0 +1,76 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import dataclasses + +from modelexpress.refit.manifest import ( + RevisionManifest, + RevisionRecord, + RevisionState, + S3Object, +) + + +def target_manifest() -> RevisionManifest: + return RevisionManifest( + model_id="model", + target_version="1", + base_version="0", + base_digest="sha256:target-0", + target_digest="sha256:target-1", + format_digest="sha256:format", + payload=S3Object( + bucket="bucket", + key="model/1/index.json", + object_version="object-version", + checksum="crc32c:01020304", + ), + ) + + +def test_minimal_manifest_round_trips_through_protobuf(): + manifest = target_manifest() + proto = manifest.to_proto() + + assert proto.model_id == "model" + assert proto.target_version == "1" + assert proto.HasField("base_version") + assert proto.HasField("base_digest") + assert proto.HasField("payload") + assert proto.payload.HasField("object_version") + assert RevisionManifest.from_proto(proto) == manifest + + +def test_launch_manifest_preserves_absence(): + manifest = RevisionManifest( + model_id="model", + target_version="0", + target_digest="sha256:target-0", + format_digest="sha256:format", + ) + proto = manifest.to_proto() + + assert not proto.HasField("base_version") + assert not proto.HasField("base_digest") + assert not proto.HasField("payload") + assert RevisionManifest.from_proto(proto) == manifest + + +def test_revision_record_round_trips_and_is_frozen(): + record = RevisionRecord(manifest=target_manifest(), state=RevisionState.READY) + + assert RevisionRecord.from_proto(record.to_proto()) == record + try: + record.state = RevisionState.COMMITTED + except dataclasses.FrozenInstanceError: + pass + else: + raise AssertionError("RevisionRecord must be frozen") + + +def test_revision_state_matches_wire_values(): + assert [(state.name, state.value) for state in RevisionState] == [ + ("UNSPECIFIED", 0), + ("READY", 1), + ("COMMITTED", 2), + ] diff --git a/modelexpress_client/python/tests/test_revision_proto.py b/modelexpress_client/python/tests/test_revision_proto.py new file mode 100644 index 000000000..2bb6a7cfe --- /dev/null +++ b/modelexpress_client/python/tests/test_revision_proto.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +import grpc +from google.protobuf import descriptor_pb2 + +from modelexpress import p2p_pb2, revision_pb2, revision_pb2_grpc + + +def _file_descriptor_proto() -> descriptor_pb2.FileDescriptorProto: + descriptor = descriptor_pb2.FileDescriptorProto() + descriptor.ParseFromString(revision_pb2.DESCRIPTOR.serialized_pb) + return descriptor + + +def _message(name: str) -> descriptor_pb2.DescriptorProto: + return next( + message + for message in _file_descriptor_proto().message_type + if message.name == name + ) + + +def test_revision_catalog_service_has_exact_public_methods(): + service = revision_pb2.DESCRIPTOR.services_by_name["RevisionCatalogService"] + + assert [method.name for method in service.methods] == [ + "PublishRevision", + "GetRevision", + "CommitRevision", + ] + assert not any(method.server_streaming for method in service.methods) + assert not any(method.client_streaming for method in service.methods) + + +def test_revision_catalog_rpc_types_are_minimal(): + service = revision_pb2.DESCRIPTOR.services_by_name["RevisionCatalogService"] + + assert { + method.name: (method.input_type.name, method.output_type.name) + for method in service.methods + } == { + "PublishRevision": ("PublishRevisionRequest", "RevisionRecord"), + "GetRevision": ("GetRevisionRequest", "RevisionRecord"), + "CommitRevision": ("CommitRevisionRequest", "RevisionRecord"), + } + + +def test_revision_proto_has_only_minimal_messages(): + assert [ + message.name for message in _file_descriptor_proto().message_type + ] == [ + "S3Object", + "RevisionManifest", + "RevisionRecord", + "PublishRevisionRequest", + "GetRevisionRequest", + "CommitRevisionRequest", + ] + + +def test_manifest_has_exact_minimal_structure(): + expected_fields = { + "S3Object": ["bucket", "key", "object_version", "checksum"], + "RevisionManifest": [ + "model_id", + "target_version", + "base_version", + "base_digest", + "target_digest", + "format_digest", + "payload", + ], + "RevisionRecord": ["manifest", "state"], + "PublishRevisionRequest": ["manifest"], + "GetRevisionRequest": ["model_id", "target_version"], + "CommitRevisionRequest": ["model_id", "target_version"], + } + + for message_name, field_names in expected_fields.items(): + message = _message(message_name) + assert [field.name for field in message.field] == field_names + assert [field.number for field in message.field] == list( + range(1, len(field_names) + 1) + ) + assert not message.reserved_range + + +def test_revision_state_has_only_ready_and_committed(): + enum = revision_pb2.DESCRIPTOR.enum_types_by_name["RevisionState"] + + assert [(value.name, value.number) for value in enum.values] == [ + ("REVISION_STATE_UNSPECIFIED", 0), + ("REVISION_STATE_READY", 1), + ("REVISION_STATE_COMMITTED", 2), + ] + + +def test_exact_base_fields_and_s3_object_version_preserve_presence(): + manifest = revision_pb2.RevisionManifest( + model_id="model", + target_version="1", + base_version="0", + base_digest="sha256:base", + target_digest="sha256:target", + format_digest="sha256:format", + payload=revision_pb2.S3Object( + bucket="bucket", + key="model/1/index.json", + object_version="object-version", + checksum="crc32c:01020304", + ), + ) + + assert manifest.HasField("base_version") + assert manifest.HasField("base_digest") + assert manifest.HasField("payload") + assert manifest.payload.HasField("object_version") + + launch = revision_pb2.RevisionManifest( + model_id="model", + target_version="0", + target_digest="sha256:target", + format_digest="sha256:format", + ) + assert not launch.HasField("base_version") + assert not launch.HasField("base_digest") + assert not launch.HasField("payload") + + +def test_generated_client_exposes_only_minimal_revision_methods(): + with grpc.insecure_channel("localhost:1") as channel: + stub = revision_pb2_grpc.RevisionCatalogServiceStub(channel) + + assert callable(stub.PublishRevision) + assert callable(stub.GetRevision) + assert callable(stub.CommitRevision) + assert not hasattr(stub, "ListReadyRevisions") + assert not hasattr(stub, "GetRecoveryCandidates") + assert not hasattr(stub, "UpdateReceiverState") + assert not hasattr(stub, "CommitVersion") + + +def test_revision_contract_is_independent_from_p2p_service(): + p2p_service = p2p_pb2.DESCRIPTOR.services_by_name["P2pService"] + revision_methods = {"PublishRevision", "GetRevision", "CommitRevision"} + + assert revision_methods.isdisjoint(method.name for method in p2p_service.methods) + assert "RevisionRecord" not in p2p_pb2.DESCRIPTOR.message_types_by_name + assert "RevisionState" not in p2p_pb2.DESCRIPTOR.enum_types_by_name diff --git a/modelexpress_common/build.rs b/modelexpress_common/build.rs index f3efa9cd9..536099843 100644 --- a/modelexpress_common/build.rs +++ b/modelexpress_common/build.rs @@ -13,6 +13,7 @@ fn main() -> Result<(), Box> { "proto/model.proto", "proto/p2p.proto", "proto/weight_sync.proto", + "proto/revision.proto", ], &["proto"], )?; diff --git a/modelexpress_common/proto/revision.proto b/modelexpress_common/proto/revision.proto new file mode 100644 index 000000000..8bf1113e2 --- /dev/null +++ b/modelexpress_common/proto/revision.proto @@ -0,0 +1,60 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +syntax = "proto3"; + +package model_express.revision; + +// Stores immutable model-weight revision metadata. Payload bytes never pass +// through this service. +service RevisionCatalogService { + // Atomically publishes one immutable revision as READY. + rpc PublishRevision(PublishRevisionRequest) returns (RevisionRecord); + // Returns one exact revision by model and target version. + rpc GetRevision(GetRevisionRequest) returns (RevisionRecord); + // Records orchestrator acceptance after the required rollout cohort verifies. + rpc CommitRevision(CommitRevisionRequest) returns (RevisionRecord); +} + +message S3Object { + string bucket = 1; + string key = 2; + optional string object_version = 3; + string checksum = 4; +} + +message RevisionManifest { + string model_id = 1; + string target_version = 2; + optional string base_version = 3; + optional string base_digest = 4; + string target_digest = 5; + string format_digest = 6; + // Absent only for the metadata-only launch anchor. + S3Object payload = 7; +} + +enum RevisionState { + REVISION_STATE_UNSPECIFIED = 0; + REVISION_STATE_READY = 1; + REVISION_STATE_COMMITTED = 2; +} + +message RevisionRecord { + RevisionManifest manifest = 1; + RevisionState state = 2; +} + +message PublishRevisionRequest { + RevisionManifest manifest = 1; +} + +message GetRevisionRequest { + string model_id = 1; + string target_version = 2; +} + +message CommitRevisionRequest { + string model_id = 1; + string target_version = 2; +} diff --git a/modelexpress_common/src/lib.rs b/modelexpress_common/src/lib.rs index bad2e35c6..7406702ec 100644 --- a/modelexpress_common/src/lib.rs +++ b/modelexpress_common/src/lib.rs @@ -12,6 +12,7 @@ pub mod download; pub mod envs; pub mod models; pub mod providers; +pub mod revision; #[cfg(any(test, feature = "test-support"))] #[doc(hidden)] pub mod test_support; @@ -37,6 +38,9 @@ pub mod grpc { pub mod weight_sync { tonic::include_proto!("weight_sync"); } + pub mod revision { + tonic::include_proto!("model_express.revision"); + } } /// Defines the shared response format between server and client (legacy HTTP) diff --git a/modelexpress_common/src/revision.rs b/modelexpress_common/src/revision.rs new file mode 100644 index 000000000..be7b4885c --- /dev/null +++ b/modelexpress_common/src/revision.rs @@ -0,0 +1,222 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use thiserror::Error; + +use crate::grpc::revision::RevisionManifest; + +#[derive(Debug, Error, PartialEq, Eq)] +pub enum RevisionManifestValidationError { + #[error("model_id is required")] + MissingModelId, + #[error("target_version is required")] + MissingTargetVersion, + #[error("format_digest is required")] + MissingFormatDigest, + #[error("target_digest is required")] + MissingTargetDigest, + #[error("launch revision must not have an exact base")] + UnexpectedLaunchBase, + #[error("launch revision must not have a payload")] + UnexpectedLaunchPayload, + #[error("base_version is required")] + MissingBaseVersion, + #[error("base_digest is required")] + MissingBaseDigest, + #[error("S3 payload is required")] + MissingPayload, + #[error("S3 payload bucket is required")] + MissingPayloadBucket, + #[error("S3 payload key is required")] + MissingPayloadKey, + #[error("S3 payload checksum is required")] + MissingPayloadChecksum, + #[error("S3 payload object_version must be non-empty when present")] + EmptyObjectVersion, +} + +pub fn validate_revision_manifest( + manifest: &RevisionManifest, +) -> Result<(), RevisionManifestValidationError> { + require_text( + &manifest.model_id, + RevisionManifestValidationError::MissingModelId, + )?; + require_text( + &manifest.target_version, + RevisionManifestValidationError::MissingTargetVersion, + )?; + require_text( + &manifest.format_digest, + RevisionManifestValidationError::MissingFormatDigest, + )?; + require_text( + &manifest.target_digest, + RevisionManifestValidationError::MissingTargetDigest, + )?; + + if manifest.target_version == "0" { + if manifest.base_version.is_some() || manifest.base_digest.is_some() { + return Err(RevisionManifestValidationError::UnexpectedLaunchBase); + } + if manifest.payload.is_some() { + return Err(RevisionManifestValidationError::UnexpectedLaunchPayload); + } + return Ok(()); + } + + require_optional_text( + manifest.base_version.as_deref(), + RevisionManifestValidationError::MissingBaseVersion, + )?; + require_optional_text( + manifest.base_digest.as_deref(), + RevisionManifestValidationError::MissingBaseDigest, + )?; + let payload = manifest + .payload + .as_ref() + .ok_or(RevisionManifestValidationError::MissingPayload)?; + require_text( + &payload.bucket, + RevisionManifestValidationError::MissingPayloadBucket, + )?; + require_text( + &payload.key, + RevisionManifestValidationError::MissingPayloadKey, + )?; + require_text( + &payload.checksum, + RevisionManifestValidationError::MissingPayloadChecksum, + )?; + if payload + .object_version + .as_deref() + .is_some_and(|value| value.trim().is_empty()) + { + return Err(RevisionManifestValidationError::EmptyObjectVersion); + } + Ok(()) +} + +fn require_optional_text( + value: Option<&str>, + error: RevisionManifestValidationError, +) -> Result<(), RevisionManifestValidationError> { + match value { + Some(value) => require_text(value, error), + None => Err(error), + } +} + +fn require_text( + value: &str, + error: RevisionManifestValidationError, +) -> Result<(), RevisionManifestValidationError> { + if value.trim().is_empty() { + Err(error) + } else { + Ok(()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::grpc::revision::S3Object; + + fn launch_manifest() -> RevisionManifest { + RevisionManifest { + model_id: "model".to_string(), + target_version: "0".to_string(), + target_digest: "sha256:target-0".to_string(), + format_digest: "sha256:format".to_string(), + ..Default::default() + } + } + + fn target_manifest() -> RevisionManifest { + RevisionManifest { + model_id: "model".to_string(), + target_version: "1".to_string(), + base_version: Some("0".to_string()), + base_digest: Some("sha256:target-0".to_string()), + target_digest: "sha256:target-1".to_string(), + format_digest: "sha256:format".to_string(), + payload: Some(S3Object { + bucket: "bucket".to_string(), + key: "model/1/index.json".to_string(), + object_version: None, + checksum: "crc32c:01020304".to_string(), + }), + } + } + + #[test] + fn launch_anchor_has_no_base_or_payload() { + assert_eq!(validate_revision_manifest(&launch_manifest()), Ok(())); + + let mut with_base = launch_manifest(); + with_base.base_version = Some("previous".to_string()); + assert_eq!( + validate_revision_manifest(&with_base), + Err(RevisionManifestValidationError::UnexpectedLaunchBase) + ); + + let mut with_payload = launch_manifest(); + with_payload.payload = target_manifest().payload; + assert_eq!( + validate_revision_manifest(&with_payload), + Err(RevisionManifestValidationError::UnexpectedLaunchPayload) + ); + } + + #[test] + fn required_text_rejects_whitespace_only_values() { + let mut manifest = target_manifest(); + manifest.model_id = " ".to_string(); + assert_eq!( + validate_revision_manifest(&manifest), + Err(RevisionManifestValidationError::MissingModelId) + ); + + let mut manifest = target_manifest(); + match manifest.payload.as_mut() { + Some(payload) => payload.checksum = "\t".to_string(), + None => panic!("target manifest must have a payload"), + } + assert_eq!( + validate_revision_manifest(&manifest), + Err(RevisionManifestValidationError::MissingPayloadChecksum) + ); + + let mut manifest = target_manifest(); + match manifest.payload.as_mut() { + Some(payload) => payload.object_version = Some(" ".to_string()), + None => panic!("target manifest must have a payload"), + } + assert_eq!( + validate_revision_manifest(&manifest), + Err(RevisionManifestValidationError::EmptyObjectVersion) + ); + } + + #[test] + fn later_revision_requires_exact_base_and_s3_payload() { + assert_eq!(validate_revision_manifest(&target_manifest()), Ok(())); + + let mut missing_base = target_manifest(); + missing_base.base_version = None; + assert_eq!( + validate_revision_manifest(&missing_base), + Err(RevisionManifestValidationError::MissingBaseVersion) + ); + + let mut missing_payload = target_manifest(); + missing_payload.payload = None; + assert_eq!( + validate_revision_manifest(&missing_payload), + Err(RevisionManifestValidationError::MissingPayload) + ); + } +} diff --git a/modelexpress_server/src/lib.rs b/modelexpress_server/src/lib.rs index f0aec8ade..f0cecea68 100644 --- a/modelexpress_server/src/lib.rs +++ b/modelexpress_server/src/lib.rs @@ -7,6 +7,7 @@ pub mod cache; pub mod config; pub mod p2p; pub mod registry; +pub mod revision; pub mod server; pub mod services; pub mod weight_sync; diff --git a/modelexpress_server/src/revision.rs b/modelexpress_server/src/revision.rs new file mode 100644 index 000000000..150308763 --- /dev/null +++ b/modelexpress_server/src/revision.rs @@ -0,0 +1,6 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +pub mod backend; +pub mod service; +pub mod state; diff --git a/modelexpress_server/src/revision/backend.rs b/modelexpress_server/src/revision/backend.rs new file mode 100644 index 000000000..c80a7b92c --- /dev/null +++ b/modelexpress_server/src/revision/backend.rs @@ -0,0 +1,94 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::sync::Arc; + +use async_trait::async_trait; +use modelexpress_common::grpc::revision::RevisionRecord; + +use crate::backend_config::BackendConfig; + +pub mod redis; +#[cfg(any(test, feature = "integration-tests"))] +pub(crate) mod testing; + +pub type CatalogResult = Result>; + +#[derive(Debug, Clone, PartialEq)] +pub enum PublishReadyOutcome { + Created(RevisionRecord), + Existing(RevisionRecord), + Conflict, +} + +#[derive(Debug, Clone, PartialEq)] +pub enum CommitOutcome { + Committed(RevisionRecord), + AlreadyCommitted(RevisionRecord), + NotFound, + InvalidState(RevisionRecord), +} + +#[async_trait] +pub trait RevisionCatalogBackend: Send + Sync { + async fn connect(&self) -> CatalogResult<()>; + + async fn publish_ready(&self, record: RevisionRecord) -> CatalogResult; + + async fn get_revision( + &self, + model_id: &str, + target_version: &str, + ) -> CatalogResult>; + + async fn commit_revision( + &self, + model_id: &str, + target_version: &str, + ) -> CatalogResult; +} + +pub type DynRevisionCatalogBackend = Arc; + +pub async fn create_revision_catalog_backend( + config: BackendConfig, +) -> CatalogResult> { + match config { + BackendConfig::Redis { url } => { + let backend = redis::RedisRevisionCatalogBackend::new(&url); + backend.connect().await?; + Ok(Some(Arc::new(backend))) + } + BackendConfig::Kubernetes { .. } => Ok(None), + #[cfg(feature = "memory-backend")] + BackendConfig::Memory => { + #[cfg(feature = "integration-tests")] + { + let backend = testing::TestRevisionCatalogBackend::new(); + backend.connect().await?; + Ok(Some(Arc::new(backend))) + } + #[cfg(not(feature = "integration-tests"))] + { + Err("in-memory revision catalog requires the 'integration-tests' feature".into()) + } + } + } +} + +#[cfg(test)] +#[allow(clippy::expect_used)] +mod tests { + use super::*; + + #[tokio::test] + async fn revision_catalog_is_disabled_for_kubernetes_metadata() { + let backend = create_revision_catalog_backend(BackendConfig::Kubernetes { + namespace: "test".to_string(), + }) + .await + .expect("Kubernetes metadata must not prevent server startup"); + + assert!(backend.is_none()); + } +} diff --git a/modelexpress_server/src/revision/backend/redis.rs b/modelexpress_server/src/revision/backend/redis.rs new file mode 100644 index 000000000..1fe01aa28 --- /dev/null +++ b/modelexpress_server/src/revision/backend/redis.rs @@ -0,0 +1,411 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::sync::Arc; + +use async_trait::async_trait; +use modelexpress_common::grpc::revision::{RevisionRecord, RevisionState}; +use prost::Message; +use redis::AsyncCommands; +use redis::aio::ConnectionManager; +use sha2::{Digest, Sha256}; +use tokio::sync::RwLock; + +use super::{CatalogResult, CommitOutcome, PublishReadyOutcome, RevisionCatalogBackend}; + +const PUBLISH_LUA: &str = r#" +local existing_manifest = redis.call('HGET', KEYS[1], 'manifest') +if existing_manifest then + local existing_record = redis.call('HGET', KEYS[1], 'record') + local existing_state = redis.call('HGET', KEYS[1], 'state') or '' + if existing_manifest == ARGV[1] then + return {2, existing_record, existing_state} + end + return {0, existing_record, existing_state} +end +redis.call('HSET', KEYS[1], 'manifest', ARGV[1], 'record', ARGV[2], 'state', ARGV[3]) +return {1, ARGV[2], ARGV[3]} +"#; + +const COMMIT_LUA: &str = r#" +local current = redis.call('HGET', KEYS[1], 'record') +if not current then return {0, '', ''} end +local state = redis.call('HGET', KEYS[1], 'state') +if not state then + state = ARGV[3] + redis.call('HSET', KEYS[1], 'state', state) +end +if state == ARGV[2] then return {2, current, state} end +if state ~= ARGV[1] then return {3, current, state or ''} end +redis.call('HSET', KEYS[1], 'state', ARGV[2]) +return {1, current, ARGV[2]} +"#; + +fn digest_hex(value: &str) -> String { + let digest = Sha256::digest(value.as_bytes()); + digest.iter().map(|byte| format!("{byte:02x}")).collect() +} + +fn revision_key(model_id: &str, target_version: &str) -> String { + format!( + "mx:revision-v0:{}", + digest_hex(&format!("{model_id}\0{target_version}")) + ) +} + +fn decode_record(bytes: &[u8]) -> CatalogResult { + RevisionRecord::decode(bytes).map_err(Into::into) +} + +fn verify_record_identity( + record: RevisionRecord, + model_id: &str, + target_version: &str, +) -> CatalogResult { + if record.manifest.as_ref().is_some_and(|manifest| { + manifest.model_id == model_id && manifest.target_version == target_version + }) { + Ok(record) + } else { + Err(format!("corrupt revision identity at '{model_id}/{target_version}'").into()) + } +} + +fn decode_stored_record( + bytes: &[u8], + state: Option, + model_id: &str, + target_version: &str, +) -> CatalogResult { + let mut record = verify_record_identity(decode_record(bytes)?, model_id, target_version)?; + if let Some(state) = state { + record.state = state; + } + Ok(record) +} + +pub struct RedisRevisionCatalogBackend { + redis: Arc>>, + redis_url: String, +} + +impl RedisRevisionCatalogBackend { + #[must_use] + pub fn new(redis_url: &str) -> Self { + Self { + redis: Arc::new(RwLock::new(None)), + redis_url: redis_url.to_string(), + } + } + + async fn connection(&self) -> CatalogResult { + { + let guard = self.redis.read().await; + if let Some(connection) = guard.as_ref() { + return Ok(connection.clone()); + } + } + let mut guard = self.redis.write().await; + if let Some(connection) = guard.as_ref() { + return Ok(connection.clone()); + } + let client = redis::Client::open(self.redis_url.as_str())?; + let connection = ConnectionManager::new(client).await?; + *guard = Some(connection.clone()); + Ok(connection) + } +} + +#[async_trait] +impl RevisionCatalogBackend for RedisRevisionCatalogBackend { + async fn connect(&self) -> CatalogResult<()> { + let mut connection = self.connection().await?; + let response: String = redis::cmd("PING").query_async(&mut connection).await?; + if response != "PONG" { + return Err(format!("unexpected Redis PING response: {response}").into()); + } + Ok(()) + } + + async fn publish_ready(&self, record: RevisionRecord) -> CatalogResult { + let submitted_manifest = record + .manifest + .as_ref() + .ok_or_else(|| "revision record is missing manifest".to_string())? + .clone(); + let key = revision_key( + &submitted_manifest.model_id, + &submitted_manifest.target_version, + ); + let manifest_bytes = submitted_manifest.encode_to_vec(); + let record_bytes = record.encode_to_vec(); + let mut connection = self.connection().await?; + let (code, existing, existing_state): (i32, Vec, String) = + redis::Script::new(PUBLISH_LUA) + .key(key) + .arg(manifest_bytes) + .arg(record_bytes) + .arg(RevisionState::Ready as i32) + .invoke_async(&mut connection) + .await?; + match code { + 1 => Ok(PublishReadyOutcome::Created(record)), + 2 => { + let existing = decode_stored_record( + &existing, + if existing_state.is_empty() { + None + } else { + Some(existing_state.parse()?) + }, + &submitted_manifest.model_id, + &submitted_manifest.target_version, + )?; + if existing.manifest.as_ref() != Some(&submitted_manifest) { + return Err("corrupt Redis revision record/manifest pair".into()); + } + Ok(PublishReadyOutcome::Existing(existing)) + } + 0 => Ok(PublishReadyOutcome::Conflict), + other => Err(format!("unexpected Redis publish result: {other}").into()), + } + } + + async fn get_revision( + &self, + model_id: &str, + target_version: &str, + ) -> CatalogResult> { + let mut connection = self.connection().await?; + let (bytes, state): (Option>, Option) = connection + .hget(revision_key(model_id, target_version), ("record", "state")) + .await?; + match bytes { + Some(bytes) => Ok(Some(decode_stored_record( + &bytes, + state, + model_id, + target_version, + )?)), + None => Ok(None), + } + } + + async fn commit_revision( + &self, + model_id: &str, + target_version: &str, + ) -> CatalogResult { + let Some(current) = self.get_revision(model_id, target_version).await? else { + return Ok(CommitOutcome::NotFound); + }; + let current_state = current.state; + let mut connection = self.connection().await?; + let (code, stored, stored_state): (i32, Vec, String) = redis::Script::new(COMMIT_LUA) + .key(revision_key(model_id, target_version)) + .arg(RevisionState::Ready as i32) + .arg(RevisionState::Committed as i32) + .arg(current_state) + .invoke_async(&mut connection) + .await?; + let state = || -> CatalogResult { Ok(stored_state.parse()?) }; + match code { + 1 => Ok(CommitOutcome::Committed(decode_stored_record( + &stored, + Some(state()?), + model_id, + target_version, + )?)), + 0 => Ok(CommitOutcome::NotFound), + 2 => Ok(CommitOutcome::AlreadyCommitted(decode_stored_record( + &stored, + Some(state()?), + model_id, + target_version, + )?)), + 3 => Ok(CommitOutcome::InvalidState(decode_stored_record( + &stored, + Some(state()?), + model_id, + target_version, + )?)), + other => Err(format!("unexpected Redis commit result: {other}").into()), + } + } +} + +#[cfg(test)] +#[allow(clippy::expect_used)] +mod tests { + use super::*; + use modelexpress_common::grpc::revision::RevisionManifest; + + #[test] + fn revision_keys_are_namespaced_and_bind_model_and_version() { + let first = revision_key("model-a", "1"); + assert!(first.starts_with("mx:revision-v0:")); + assert_ne!(first, revision_key("model-b", "1")); + assert_ne!(first, revision_key("model-a", "2")); + } + + #[test] + fn stored_record_identity_must_match_the_lookup_key() { + let record = RevisionRecord { + manifest: Some(RevisionManifest { + model_id: "model".to_string(), + target_version: "1".to_string(), + ..Default::default() + }), + state: RevisionState::Ready as i32, + }; + assert!(verify_record_identity(record.clone(), "model", "1").is_ok()); + assert!(verify_record_identity(record, "model", "2").is_err()); + } + + #[tokio::test] + #[ignore = "requires a live Redis at REDIS_URL"] + async fn commit_does_not_depend_on_byte_exact_protobuf_reencoding() { + let redis_url = + std::env::var("REDIS_URL").unwrap_or_else(|_| "redis://127.0.0.1:6379".to_string()); + let backend = RedisRevisionCatalogBackend::new(&redis_url); + backend.connect().await.expect("connect to Redis"); + let model_id = format!("revision-cas-{}", std::process::id()); + let target_version = "1"; + let key = revision_key(&model_id, target_version); + let manifest = RevisionManifest { + model_id: model_id.clone(), + target_version: target_version.to_string(), + target_digest: "sha256:target".to_string(), + format_digest: "sha256:format".to_string(), + ..Default::default() + }; + let record = RevisionRecord { + manifest: Some(manifest.clone()), + state: RevisionState::Ready as i32, + }; + let mut stored_record = record.encode_to_vec(); + stored_record.extend_from_slice(&[0x78, 0x01]); + let mut connection = backend.connection().await.expect("Redis connection"); + let _: () = redis::cmd("HSET") + .arg(&key) + .arg("manifest") + .arg(manifest.encode_to_vec()) + .arg("record") + .arg(stored_record.clone()) + .query_async(&mut connection) + .await + .expect("seed forward-compatible record"); + + let outcome = backend + .commit_revision(&model_id, target_version) + .await + .expect("commit"); + + assert!(matches!(outcome, CommitOutcome::Committed(_))); + let record_after_commit: Vec = redis::cmd("HGET") + .arg(&key) + .arg("record") + .query_async(&mut connection) + .await + .expect("read immutable record bytes"); + assert_eq!(record_after_commit, stored_record); + let fetched = backend + .get_revision(&model_id, target_version) + .await + .expect("read committed revision") + .expect("revision exists"); + assert_eq!(fetched.state, RevisionState::Committed as i32); + let _: () = redis::cmd("DEL") + .arg(key) + .query_async(&mut connection) + .await + .expect("cleanup revision"); + } + + #[tokio::test] + #[ignore = "requires a live Redis at REDIS_URL"] + async fn redis_commit_preserves_all_lifecycle_outcomes() { + let redis_url = + std::env::var("REDIS_URL").unwrap_or_else(|_| "redis://127.0.0.1:6379".to_string()); + let backend = RedisRevisionCatalogBackend::new(&redis_url); + backend.connect().await.expect("connect to Redis"); + let model_id = format!("revision-lifecycle-{}", std::process::id()); + let target_version = "1"; + let key = revision_key(&model_id, target_version); + let record = RevisionRecord { + manifest: Some(RevisionManifest { + model_id: model_id.clone(), + target_version: target_version.to_string(), + target_digest: "sha256:target".to_string(), + format_digest: "sha256:format".to_string(), + ..Default::default() + }), + state: RevisionState::Ready as i32, + }; + + assert!(matches!( + backend + .publish_ready(record.clone()) + .await + .expect("publish"), + PublishReadyOutcome::Created(_) + )); + assert!(matches!( + backend + .publish_ready(record.clone()) + .await + .expect("idempotent publish"), + PublishReadyOutcome::Existing(_) + )); + assert!(matches!( + backend + .commit_revision(&model_id, target_version) + .await + .expect("commit"), + CommitOutcome::Committed(_) + )); + assert!(matches!( + backend + .commit_revision(&model_id, target_version) + .await + .expect("idempotent commit"), + CommitOutcome::AlreadyCommitted(_) + )); + let PublishReadyOutcome::Existing(existing) = backend + .publish_ready(record.clone()) + .await + .expect("publish replay after commit") + else { + panic!("published revision should already exist"); + }; + assert_eq!(existing.state, RevisionState::Committed as i32); + assert!(matches!( + backend + .commit_revision(&model_id, "missing") + .await + .expect("not found"), + CommitOutcome::NotFound + )); + + let mut connection = backend.connection().await.expect("Redis connection"); + let _: () = redis::cmd("HSET") + .arg(&key) + .arg("state") + .arg(99) + .query_async(&mut connection) + .await + .expect("seed invalid state"); + assert!(matches!( + backend + .commit_revision(&model_id, target_version) + .await + .expect("invalid state"), + CommitOutcome::InvalidState(_) + )); + let _: () = redis::cmd("DEL") + .arg(key) + .query_async(&mut connection) + .await + .expect("cleanup revision"); + } +} diff --git a/modelexpress_server/src/revision/backend/testing.rs b/modelexpress_server/src/revision/backend/testing.rs new file mode 100644 index 000000000..4fb8f56cb --- /dev/null +++ b/modelexpress_server/src/revision/backend/testing.rs @@ -0,0 +1,103 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +//! Deterministic test double selected for `BackendConfig::Memory` when the +//! `integration-tests` feature is enabled. + +use std::collections::HashMap; +use std::sync::{Mutex, PoisonError}; + +use async_trait::async_trait; +use modelexpress_common::grpc::revision::{RevisionRecord, RevisionState}; + +use super::{CatalogResult, CommitOutcome, PublishReadyOutcome, RevisionCatalogBackend}; + +#[derive(Default)] +struct TestCatalog { + revisions: HashMap<(String, String), RevisionRecord>, +} + +#[derive(Default)] +pub struct TestRevisionCatalogBackend { + state: Mutex, +} + +impl TestRevisionCatalogBackend { + #[must_use] + pub fn new() -> Self { + Self::default() + } + + fn lock(&self) -> std::sync::MutexGuard<'_, TestCatalog> { + self.state.lock().unwrap_or_else(PoisonError::into_inner) + } + + #[cfg(test)] + pub fn insert(&self, record: RevisionRecord) { + let Some(manifest) = record.manifest.as_ref() else { + panic!("test record must include a manifest"); + }; + self.lock().revisions.insert( + (manifest.model_id.clone(), manifest.target_version.clone()), + record, + ); + } +} + +#[async_trait] +impl RevisionCatalogBackend for TestRevisionCatalogBackend { + async fn connect(&self) -> CatalogResult<()> { + Ok(()) + } + + async fn publish_ready(&self, record: RevisionRecord) -> CatalogResult { + let manifest = record + .manifest + .as_ref() + .ok_or_else(|| "revision record is missing manifest".to_string())?; + let key = (manifest.model_id.clone(), manifest.target_version.clone()); + let mut state = self.lock(); + match state.revisions.get(&key) { + Some(existing) if existing.manifest == record.manifest => { + Ok(PublishReadyOutcome::Existing(existing.clone())) + } + Some(_) => Ok(PublishReadyOutcome::Conflict), + None => { + state.revisions.insert(key, record.clone()); + Ok(PublishReadyOutcome::Created(record)) + } + } + } + + async fn get_revision( + &self, + model_id: &str, + target_version: &str, + ) -> CatalogResult> { + Ok(self + .lock() + .revisions + .get(&(model_id.to_string(), target_version.to_string())) + .cloned()) + } + + async fn commit_revision( + &self, + model_id: &str, + target_version: &str, + ) -> CatalogResult { + let key = (model_id.to_string(), target_version.to_string()); + let mut state = self.lock(); + let Some(record) = state.revisions.get_mut(&key) else { + return Ok(CommitOutcome::NotFound); + }; + match RevisionState::try_from(record.state).ok() { + Some(RevisionState::Ready) => { + record.state = RevisionState::Committed as i32; + Ok(CommitOutcome::Committed(record.clone())) + } + Some(RevisionState::Committed) => Ok(CommitOutcome::AlreadyCommitted(record.clone())), + _ => Ok(CommitOutcome::InvalidState(record.clone())), + } + } +} diff --git a/modelexpress_server/src/revision/service.rs b/modelexpress_server/src/revision/service.rs new file mode 100644 index 000000000..fec94e6b2 --- /dev/null +++ b/modelexpress_server/src/revision/service.rs @@ -0,0 +1,206 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +use std::sync::Arc; + +use modelexpress_common::grpc::revision::revision_catalog_service_server::RevisionCatalogService; +use modelexpress_common::grpc::revision::{ + CommitRevisionRequest, GetRevisionRequest, PublishRevisionRequest, RevisionRecord, +}; +use tonic::{Request, Response, Status}; + +use super::state::{CatalogError, RevisionCatalogState}; + +#[derive(Clone)] +pub struct RevisionCatalogServiceImpl { + state: Arc, +} + +impl RevisionCatalogServiceImpl { + #[must_use] + pub fn new(state: Arc) -> Self { + Self { state } + } +} + +#[derive(Debug, Clone, Copy)] +enum InputError { + Missing(&'static str), +} + +impl From for Status { + fn from(error: InputError) -> Self { + match error { + InputError::Missing(field) => Status::invalid_argument(format!("{field} is required")), + } + } +} + +fn require_text(value: &str, field: &'static str) -> Result<(), InputError> { + if value.trim().is_empty() { + Err(InputError::Missing(field)) + } else { + Ok(()) + } +} + +fn map_catalog_error(error: CatalogError) -> Status { + match error { + CatalogError::InvalidManifest(message) => Status::invalid_argument(message), + CatalogError::ManifestConflict { .. } => Status::already_exists(error.to_string()), + CatalogError::RevisionNotFound { .. } => Status::not_found(error.to_string()), + CatalogError::InvalidLifecycle { .. } => Status::failed_precondition(error.to_string()), + CatalogError::Backend(message) => { + tracing::error!(error = %message, "revision catalog backend failure"); + Status::internal("revision catalog backend failure") + } + } +} + +#[tonic::async_trait] +impl RevisionCatalogService for RevisionCatalogServiceImpl { + async fn publish_revision( + &self, + request: Request, + ) -> Result, Status> { + let manifest = request + .into_inner() + .manifest + .ok_or_else(|| Status::invalid_argument("manifest is required"))?; + let published = self + .state + .publish(manifest) + .await + .map_err(map_catalog_error)?; + Ok(Response::new(published.record)) + } + + async fn get_revision( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + require_text(&request.model_id, "model_id")?; + require_text(&request.target_version, "target_version")?; + let record = self + .state + .get(&request.model_id, &request.target_version) + .await + .map_err(map_catalog_error)? + .ok_or_else(|| { + Status::not_found(format!( + "revision '{}/{}' was not found", + request.model_id, request.target_version + )) + })?; + Ok(Response::new(record)) + } + + async fn commit_revision( + &self, + request: Request, + ) -> Result, Status> { + let request = request.into_inner(); + require_text(&request.model_id, "model_id")?; + require_text(&request.target_version, "target_version")?; + let record = self + .state + .commit(&request.model_id, &request.target_version) + .await + .map_err(map_catalog_error)?; + Ok(Response::new(record)) + } +} + +#[cfg(test)] +#[allow(clippy::expect_used)] +mod tests { + use super::*; + use modelexpress_common::grpc::revision::{RevisionManifest, RevisionState}; + + fn launch_manifest() -> RevisionManifest { + RevisionManifest { + model_id: "model".to_string(), + target_version: "0".to_string(), + target_digest: "sha256:target-0".to_string(), + format_digest: "sha256:format".to_string(), + ..Default::default() + } + } + + fn service() -> RevisionCatalogServiceImpl { + RevisionCatalogServiceImpl::new(Arc::new(RevisionCatalogState::for_tests())) + } + + #[tokio::test] + async fn publish_get_and_commit_one_exact_revision() { + let service = service(); + let published = service + .publish_revision(Request::new(PublishRevisionRequest { + manifest: Some(launch_manifest()), + })) + .await + .expect("publish") + .into_inner(); + assert_eq!(published.state, RevisionState::Ready as i32); + + let fetched = service + .get_revision(Request::new(GetRevisionRequest { + model_id: "model".to_string(), + target_version: "0".to_string(), + })) + .await + .expect("get") + .into_inner(); + assert_eq!(fetched, published); + + let committed = service + .commit_revision(Request::new(CommitRevisionRequest { + model_id: "model".to_string(), + target_version: "0".to_string(), + })) + .await + .expect("commit") + .into_inner(); + assert_eq!(committed.state, RevisionState::Committed as i32); + } + + #[test] + fn backend_errors_are_not_exposed_to_clients() { + let status = map_catalog_error(CatalogError::Backend( + "redis://user:secret@internal-host:6379 failed".to_string(), + )); + + assert_eq!(status.code(), tonic::Code::Internal); + assert_eq!(status.message(), "revision catalog backend failure"); + assert!(!status.message().contains("secret")); + } + + #[tokio::test] + async fn malformed_requests_are_rejected_at_the_service_boundary() { + let service = service(); + let missing_manifest = service + .publish_revision(Request::new(PublishRevisionRequest { manifest: None })) + .await + .expect_err("missing manifest"); + assert_eq!(missing_manifest.code(), tonic::Code::InvalidArgument); + + let missing_version = service + .get_revision(Request::new(GetRevisionRequest { + model_id: "model".to_string(), + target_version: String::new(), + })) + .await + .expect_err("missing target version"); + assert_eq!(missing_version.code(), tonic::Code::InvalidArgument); + + let whitespace_model = service + .get_revision(Request::new(GetRevisionRequest { + model_id: " ".to_string(), + target_version: "0".to_string(), + })) + .await + .expect_err("whitespace-only model id"); + assert_eq!(whitespace_model.code(), tonic::Code::InvalidArgument); + } +} diff --git a/modelexpress_server/src/revision/state.rs b/modelexpress_server/src/revision/state.rs new file mode 100644 index 000000000..ea44e1b58 --- /dev/null +++ b/modelexpress_server/src/revision/state.rs @@ -0,0 +1,216 @@ +// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +// SPDX-License-Identifier: Apache-2.0 + +#[cfg(any(test, feature = "integration-tests"))] +use std::sync::Arc; + +use modelexpress_common::grpc::revision::{RevisionManifest, RevisionRecord, RevisionState}; +use modelexpress_common::revision::validate_revision_manifest; +use thiserror::Error; + +use super::backend::{CommitOutcome, DynRevisionCatalogBackend, PublishReadyOutcome}; + +#[derive(Debug, Clone, PartialEq)] +pub struct PublicationResult { + pub record: RevisionRecord, + pub created: bool, +} + +#[derive(Debug, Error)] +pub enum CatalogError { + #[error("invalid revision manifest: {0}")] + InvalidManifest(String), + #[error("revision '{model_id}/{target_version}' already exists with a different manifest")] + ManifestConflict { + model_id: String, + target_version: String, + }, + #[error("revision '{model_id}/{target_version}' was not found")] + RevisionNotFound { + model_id: String, + target_version: String, + }, + #[error("revision '{model_id}/{target_version}' has invalid lifecycle state {state}")] + InvalidLifecycle { + model_id: String, + target_version: String, + state: i32, + }, + #[error("revision catalog backend error: {0}")] + Backend(String), +} + +#[derive(Clone)] +pub struct RevisionCatalogState { + backend: DynRevisionCatalogBackend, +} + +impl RevisionCatalogState { + #[must_use] + pub fn with_backend(backend: DynRevisionCatalogBackend) -> Self { + Self { backend } + } + + #[cfg(any(test, feature = "integration-tests"))] + #[must_use] + pub fn for_tests() -> Self { + Self::with_backend(Arc::new( + super::backend::testing::TestRevisionCatalogBackend::new(), + )) + } + + pub async fn publish( + &self, + manifest: RevisionManifest, + ) -> Result { + validate_revision_manifest(&manifest) + .map_err(|error| CatalogError::InvalidManifest(error.to_string()))?; + let model_id = manifest.model_id.clone(); + let target_version = manifest.target_version.clone(); + let record = RevisionRecord { + manifest: Some(manifest), + state: RevisionState::Ready as i32, + }; + match self + .backend + .publish_ready(record) + .await + .map_err(|error| CatalogError::Backend(error.to_string()))? + { + PublishReadyOutcome::Created(record) => Ok(PublicationResult { + record, + created: true, + }), + PublishReadyOutcome::Existing(record) => Ok(PublicationResult { + record, + created: false, + }), + PublishReadyOutcome::Conflict => Err(CatalogError::ManifestConflict { + model_id, + target_version, + }), + } + } + + pub async fn get( + &self, + model_id: &str, + target_version: &str, + ) -> Result, CatalogError> { + self.backend + .get_revision(model_id, target_version) + .await + .map_err(|error| CatalogError::Backend(error.to_string())) + } + + pub async fn commit( + &self, + model_id: &str, + target_version: &str, + ) -> Result { + match self + .backend + .commit_revision(model_id, target_version) + .await + .map_err(|error| CatalogError::Backend(error.to_string()))? + { + CommitOutcome::Committed(record) | CommitOutcome::AlreadyCommitted(record) => { + Ok(record) + } + CommitOutcome::NotFound => Err(CatalogError::RevisionNotFound { + model_id: model_id.to_string(), + target_version: target_version.to_string(), + }), + CommitOutcome::InvalidState(record) => Err(CatalogError::InvalidLifecycle { + model_id: model_id.to_string(), + target_version: target_version.to_string(), + state: record.state, + }), + } + } +} + +#[cfg(test)] +#[allow(clippy::expect_used)] +mod tests { + use super::*; + use modelexpress_common::grpc::revision::S3Object; + + fn launch_manifest() -> RevisionManifest { + RevisionManifest { + model_id: "model".to_string(), + target_version: "0".to_string(), + target_digest: "sha256:target-0".to_string(), + format_digest: "sha256:format".to_string(), + ..Default::default() + } + } + + fn target_manifest() -> RevisionManifest { + RevisionManifest { + model_id: "model".to_string(), + target_version: "1".to_string(), + base_version: Some("0".to_string()), + base_digest: Some("sha256:target-0".to_string()), + target_digest: "sha256:target-1".to_string(), + format_digest: "sha256:format".to_string(), + payload: Some(S3Object { + bucket: "bucket".to_string(), + key: "model/1/index.json".to_string(), + object_version: None, + checksum: "crc32c:01020304".to_string(), + }), + } + } + + #[tokio::test] + async fn immutable_publication_is_idempotent_and_rejects_conflicts() { + let state = RevisionCatalogState::for_tests(); + let first = state.publish(target_manifest()).await.expect("publish"); + assert!(first.created); + assert_eq!(first.record.state, RevisionState::Ready as i32); + + let retry = state.publish(target_manifest()).await.expect("retry"); + assert!(!retry.created); + assert_eq!(retry.record, first.record); + + let mut conflicting = target_manifest(); + conflicting.target_digest = "sha256:different".to_string(); + assert!(matches!( + state.publish(conflicting).await, + Err(CatalogError::ManifestConflict { .. }) + )); + } + + #[tokio::test] + async fn commit_is_an_idempotent_ready_to_committed_transition() { + let state = RevisionCatalogState::for_tests(); + state.publish(launch_manifest()).await.expect("publish"); + + let committed = state.commit("model", "0").await.expect("commit"); + assert_eq!(committed.state, RevisionState::Committed as i32); + assert_eq!( + state.commit("model", "0").await.expect("idempotent"), + committed + ); + assert!(matches!( + state.commit("model", "missing").await, + Err(CatalogError::RevisionNotFound { .. }) + )); + } + + #[tokio::test] + async fn commit_rejects_an_unrecognized_lifecycle_state() { + let backend = Arc::new(super::super::backend::testing::TestRevisionCatalogBackend::new()); + backend.insert(RevisionRecord { + manifest: Some(launch_manifest()), + state: 99, + }); + let state = RevisionCatalogState::with_backend(backend); + + assert!(matches!( + state.commit("model", "0").await, + Err(CatalogError::InvalidLifecycle { state: 99, .. }) + )); + } +} diff --git a/modelexpress_server/src/server.rs b/modelexpress_server/src/server.rs index adfeb7c35..c3b3c8848 100644 --- a/modelexpress_server/src/server.rs +++ b/modelexpress_server/src/server.rs @@ -2,8 +2,8 @@ // SPDX-License-Identifier: Apache-2.0 //! Reusable server entrypoint. `main` is a thin shell over [`run_server`] so the -//! whole startup path (registry, P2P state, health, reaper, graceful shutdown) can -//! be embedded by a downstream binary that provides its own configuration or services. +//! whole startup path (registry, optional revision catalog, P2P state, health, reaper, graceful +//! shutdown) can be embedded by a downstream binary with its own configuration or services. use std::future::Future; use std::sync::Arc; @@ -11,6 +11,7 @@ use std::sync::Arc; use modelexpress_common::grpc::{ api::api_service_server::ApiServiceServer, health::health_service_server::HealthServiceServer, model::model_service_server::ModelServiceServer, p2p::p2p_service_server::P2pServiceServer, + revision::revision_catalog_service_server::RevisionCatalogServiceServer, weight_sync::weight_sync_service_server::WeightSyncServiceServer, }; use tonic::transport::Server; @@ -23,6 +24,9 @@ use crate::cache::CacheEvictionService; use crate::config::{AuthMode, ServerConfig}; use crate::p2p::{service::P2pServiceImpl, state::P2pStateManager}; use crate::registry::state::RegistryManager; +use crate::revision::backend::create_revision_catalog_backend; +use crate::revision::service::RevisionCatalogServiceImpl; +use crate::revision::state::RevisionCatalogState; use crate::services::{ApiServiceImpl, HealthServiceImpl, ModelDownloadTracker, ModelServiceImpl}; use crate::weight_sync::WeightSyncServiceImpl; @@ -32,8 +36,9 @@ const MAX_MESSAGE_SIZE: usize = 100 * 1024 * 1024; /// Run the ModelExpress gRPC server to completion. /// -/// Connects the registry and P2P metadata backends (failing fast if either is -/// unreachable), starts the cache-eviction and reaper background tasks, serves all +/// Connects the registry and P2P metadata backends, plus the revision catalog when the +/// selected backend supports it, failing fast if any configured backend is unreachable. +/// It starts the cache-eviction and reaper background tasks, serves all /// gRPC services, and tears everything down once `shutdown` resolves. Logging is the /// caller's responsibility: install a subscriber before calling this. /// @@ -115,6 +120,34 @@ pub async fn run_server( .set_serving::>() .await; + let revision_backend = match tokio::time::timeout( + std::time::Duration::from_secs(10), + create_revision_catalog_backend(backend.clone()), + ) + .await + { + Ok(Ok(revision_backend)) => revision_backend, + Ok(Err(error)) => { + error!("Failed to connect to revision catalog backend: {error}"); + return Err(error); + } + Err(_) => { + error!("Timed out connecting to revision catalog backend"); + return Err("revision catalog backend connection timed out".into()); + } + }; + let revision_service = revision_backend.map(|revision_backend| { + let revision_state = Arc::new(RevisionCatalogState::with_backend(revision_backend)); + RevisionCatalogServiceImpl::new(revision_state) + }); + if revision_service.is_some() { + health_reporter + .set_serving::>() + .await; + } else { + info!("Revision catalog service disabled for the selected metadata backend"); + } + // Initialize P2P state manager — fails fast if backend is misconfigured or unreachable let p2p_state = Arc::new(P2pStateManager::with_config(backend)); @@ -188,6 +221,11 @@ pub async fn run_server( let weight_sync = WeightSyncServiceServer::new(weight_sync_service) .max_decoding_message_size(MAX_MESSAGE_SIZE) .max_encoding_message_size(MAX_MESSAGE_SIZE); + let revision = revision_service.map(|revision_service| { + RevisionCatalogServiceServer::new(revision_service) + .max_decoding_message_size(MAX_MESSAGE_SIZE) + .max_encoding_message_size(MAX_MESSAGE_SIZE) + }); info!("Starting gRPC server on: {addr}"); let router = Server::builder() @@ -198,12 +236,14 @@ pub async fn run_server( .add_service(layer.layer(api)) .add_service(layer.layer(model)) .add_service(layer.layer(p2p)) - .add_service(layer.layer(weight_sync)), + .add_service(layer.layer(weight_sync)) + .add_optional_service(revision.map(|revision| layer.layer(revision))), None => router .add_service(api) .add_service(model) .add_service(p2p) - .add_service(weight_sync), + .add_service(weight_sync) + .add_optional_service(revision), }; let server_result = router.serve_with_shutdown(addr, shutdown_signal).await; diff --git a/modelexpress_server/tests/in_process_server.rs b/modelexpress_server/tests/in_process_server.rs index 6d205a084..6d9b1bf7f 100644 --- a/modelexpress_server/tests/in_process_server.rs +++ b/modelexpress_server/tests/in_process_server.rs @@ -16,6 +16,11 @@ use std::time::Duration; use modelexpress_client::Client; use modelexpress_common::client_config::ClientConfig; use modelexpress_common::config::ConnectionConfig; +use modelexpress_common::grpc::revision::revision_catalog_service_client::RevisionCatalogServiceClient; +use modelexpress_common::grpc::revision::{ + CommitRevisionRequest, GetRevisionRequest, PublishRevisionRequest, RevisionManifest, + RevisionState, S3Object, +}; use modelexpress_server::backend_config::BackendConfig; use modelexpress_server::config::ServerConfig; use modelexpress_server::run_server; @@ -90,3 +95,92 @@ async fn server_boots_and_serves_a_client() { async fn another_server_boots_and_serves_a_client() { assert_boots_and_serves().await; } + +#[tokio::test] +async fn revision_catalog_runs_through_the_real_server_router() { + let port = free_port(); + let (shutdown, handle) = start_server(port); + let endpoint = format!("http://127.0.0.1:{port}"); + let mut client = tokio::time::timeout(Duration::from_secs(10), async { + loop { + match RevisionCatalogServiceClient::connect(endpoint.clone()).await { + Ok(client) => break client, + Err(_) => tokio::time::sleep(Duration::from_millis(50)).await, + } + } + }) + .await + .expect("revision catalog service did not accept connections in time"); + let manifest = RevisionManifest { + model_id: "model".to_string(), + target_version: "1".to_string(), + base_version: Some("0".to_string()), + base_digest: Some("sha256:target-0".to_string()), + target_digest: "sha256:target-1".to_string(), + format_digest: "sha256:format".to_string(), + payload: Some(S3Object { + bucket: "bucket".to_string(), + key: "model/1/root.json".to_string(), + object_version: Some("object-v1".to_string()), + checksum: "crc32c:deadbeef".to_string(), + }), + }; + let published = client + .publish_revision(PublishRevisionRequest { + manifest: Some(manifest.clone()), + }) + .await + .expect("publish over network") + .into_inner(); + assert_eq!(published.state, RevisionState::Ready as i32); + + let repeated_publish = client + .publish_revision(PublishRevisionRequest { + manifest: Some(manifest.clone()), + }) + .await + .expect("idempotent publish over network") + .into_inner(); + assert_eq!(repeated_publish, published); + + let mut conflicting_manifest = manifest; + conflicting_manifest.target_digest = "sha256:different-target".to_string(); + let conflict = client + .publish_revision(PublishRevisionRequest { + manifest: Some(conflicting_manifest), + }) + .await + .expect_err("conflicting publish must fail"); + assert_eq!(conflict.code(), tonic::Code::AlreadyExists); + + let committed = client + .commit_revision(CommitRevisionRequest { + model_id: "model".to_string(), + target_version: "1".to_string(), + }) + .await + .expect("commit over network") + .into_inner(); + assert_eq!(committed.state, RevisionState::Committed as i32); + + let repeated_commit = client + .commit_revision(CommitRevisionRequest { + model_id: "model".to_string(), + target_version: "1".to_string(), + }) + .await + .expect("idempotent commit over network") + .into_inner(); + assert_eq!(repeated_commit, committed); + + let fetched = client + .get_revision(GetRevisionRequest { + model_id: "model".to_string(), + target_version: "1".to_string(), + }) + .await + .expect("get over network") + .into_inner(); + assert_eq!(fetched, committed); + stop_and_join(shutdown, handle).await; +}