diff --git a/docs/cli/reference.mdx b/docs/cli/reference.mdx index 634e2b8d53..646365baec 100644 --- a/docs/cli/reference.mdx +++ b/docs/cli/reference.mdx @@ -3471,6 +3471,7 @@ nemo inference virtual-models delete [OPTIONS] NAME **Options:** * `--workspace` +* `--expected-db-version `: Optional database version for optimistic locking. Delete only succeeds if the VirtualModel still has this version. **Help:** diff --git a/openapi/ga/individual/platform.openapi.yaml b/openapi/ga/individual/platform.openapi.yaml index 5c62b70686..f1b5acfbc6 100644 --- a/openapi/ga/individual/platform.openapi.yaml +++ b/openapi/ga/individual/platform.openapi.yaml @@ -903,6 +903,16 @@ paths: title: Parent type: string description: Parent entity ID for nested entities + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. responses: '200': description: Successful Response @@ -3267,11 +3277,23 @@ paths: schema: type: string title: Name + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. responses: '204': description: VirtualModel deleted '404': description: VirtualModel not found. + '409': + description: VirtualModel was modified before it could be deleted. '422': description: Validation Error content: @@ -12412,6 +12434,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -12422,6 +12449,7 @@ components: - updated_by - entity_id - parent + - db_version title: GuardrailConfig description: A guardrail configuration entity. GuardrailConfigFilter: @@ -15981,6 +16009,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -15992,6 +16025,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobStep description: 'A single step within an attempt. @@ -16312,6 +16346,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -16323,6 +16362,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobTask description: 'A task within a step (for parallel execution). @@ -19324,6 +19364,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -19334,6 +19379,7 @@ components: - updated_by - entity_id - parent + - db_version title: VirtualModel description: 'Logical inference route. diff --git a/openapi/ga/openapi.yaml b/openapi/ga/openapi.yaml index 5c62b70686..f1b5acfbc6 100644 --- a/openapi/ga/openapi.yaml +++ b/openapi/ga/openapi.yaml @@ -903,6 +903,16 @@ paths: title: Parent type: string description: Parent entity ID for nested entities + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. responses: '200': description: Successful Response @@ -3267,11 +3277,23 @@ paths: schema: type: string title: Name + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. responses: '204': description: VirtualModel deleted '404': description: VirtualModel not found. + '409': + description: VirtualModel was modified before it could be deleted. '422': description: Validation Error content: @@ -12412,6 +12434,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -12422,6 +12449,7 @@ components: - updated_by - entity_id - parent + - db_version title: GuardrailConfig description: A guardrail configuration entity. GuardrailConfigFilter: @@ -15981,6 +16009,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -15992,6 +16025,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobStep description: 'A single step within an attempt. @@ -16312,6 +16346,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -16323,6 +16362,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobTask description: 'A task within a step (for parallel execution). @@ -19324,6 +19364,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -19334,6 +19379,7 @@ components: - updated_by - entity_id - parent + - db_version title: VirtualModel description: 'Logical inference route. diff --git a/openapi/openapi.yaml b/openapi/openapi.yaml index 5c62b70686..f1b5acfbc6 100644 --- a/openapi/openapi.yaml +++ b/openapi/openapi.yaml @@ -903,6 +903,16 @@ paths: title: Parent type: string description: Parent entity ID for nested entities + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. responses: '200': description: Successful Response @@ -3267,11 +3277,23 @@ paths: schema: type: string title: Name + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. responses: '204': description: VirtualModel deleted '404': description: VirtualModel not found. + '409': + description: VirtualModel was modified before it could be deleted. '422': description: Validation Error content: @@ -12412,6 +12434,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -12422,6 +12449,7 @@ components: - updated_by - entity_id - parent + - db_version title: GuardrailConfig description: A guardrail configuration entity. GuardrailConfigFilter: @@ -15981,6 +16009,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -15992,6 +16025,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobStep description: 'A single step within an attempt. @@ -16312,6 +16346,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -16323,6 +16362,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobTask description: 'A task within a step (for parallel execution). @@ -19324,6 +19364,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -19334,6 +19379,7 @@ components: - updated_by - entity_id - parent + - db_version title: VirtualModel description: 'Logical inference route. diff --git a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/inference/virtual_models.py b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/inference/virtual_models.py index 3eb964d96e..8546c1b29d 100644 --- a/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/inference/virtual_models.py +++ b/packages/nemo_platform_ext/src/nemo_platform_ext/cli/commands/api/inference/virtual_models.py @@ -177,6 +177,13 @@ def delete_virtual_models( ctx: typer.Context, name: Annotated[str, typer.Argument()], workspace: Annotated[str | None, typer.Option("--workspace")] = None, + expected_db_version: Annotated[ + int | None, + typer.Option( + "--expected-db-version", + help="Optional database version for optimistic locking. Delete only succeeds if the VirtualModel still has this version.", + ), + ] = None, ) -> None: """Permanently delete a VirtualModel. @@ -187,6 +194,7 @@ def delete_virtual_models( kwargs = build_kwargs( workspace=workspace, + expected_db_version=expected_db_version, ) client.inference.virtual_models.delete(name, **kwargs) diff --git a/packages/nemo_platform_ext/tests/local/test_services.py b/packages/nemo_platform_ext/tests/local/test_services.py index 1964276ba0..a6716995f0 100644 --- a/packages/nemo_platform_ext/tests/local/test_services.py +++ b/packages/nemo_platform_ext/tests/local/test_services.py @@ -680,20 +680,38 @@ def test_daemonize_services_bounds_probe_and_sleep_by_remaining_deadline( proc = MagicMock() proc.pid = 4242 proc.poll.return_value = None + clock = 0.0 + sleep_calls: list[float] = [] + + def monotonic() -> float: + nonlocal clock + if clock == 0.0: + clock = 4.0 + return 0.0 + return clock + + def probe_status(*_args: object, **_kwargs: object) -> bool: + nonlocal clock + clock = 4.5 + return False + + def sleep(duration: float) -> None: + nonlocal clock + sleep_calls.append(duration) + clock += duration with ( patch("nemo_platform_ext.local.services.require_services_extra"), - patch("nemo_platform_ext.local.services.probe_status", return_value=False) as probe_status, + patch("nemo_platform_ext.local.services.probe_status", side_effect=probe_status) as probe_status_mock, patch("nemo_platform_ext.local.services.subprocess.Popen", return_value=proc), - patch("nemo_platform_ext.local.services.time.monotonic", side_effect=[0.0, 4.0, 4.5, 5.0]), - patch("nemo_platform_ext.local.services.time.sleep") as sleep, + patch("nemo_platform_ext.local.services.time.monotonic", side_effect=monotonic), + patch("nemo_platform_ext.local.services.time.sleep", side_effect=sleep), ): with pytest.raises(services.ServicesStartupTimeoutError): services.daemonize_services(cfg) - assert probe_status.call_args.kwargs["timeout"] == pytest.approx(1.0) - sleep.assert_called_once() - assert sleep.call_args.args[0] == pytest.approx(0.5) + assert probe_status_mock.call_args.kwargs["timeout"] == pytest.approx(1.0) + assert sleep_calls == [pytest.approx(0.5)] proc.terminate.assert_called_once_with() diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/entities.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/entities.py index 0351d6e3e0..700587f8a2 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/entities.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/entities.py @@ -53,7 +53,7 @@ def parse_qualified_name(name: str, default_workspace: str | None = None) -> tup return default_workspace or DEFAULT_WORKSPACE, name -class EntityTypeDefault: +class EntityTypeDefault(str): """Descriptor returning snake_case class name as default for __entity_type__.""" def __get__(self, obj: object | None, objtype: type | None = None) -> str: @@ -105,7 +105,7 @@ class EntityBase(BaseModel): "_db_version", } - __entity_type__: ClassVar[str] = EntityTypeDefault() # type: ignore[assignment] + __entity_type__: ClassVar[str] = EntityTypeDefault() model_config = {"populate_by_name": True} @@ -167,6 +167,7 @@ def parent(self) -> str | None: """Parent entity ID for nested entities.""" return self._parent + @computed_field @property def db_version(self) -> int: """Database version of the entity for optimistic locking.""" @@ -217,42 +218,70 @@ class EntityToken(Protocol): EntityTypeLike = Type[EntityT] | EntityToken -class EntityClientProtocol(Protocol[EntityT]): - """Protocol defining the interface for entity clients.""" +class EntityGetterProtocol(Protocol[EntityT]): + """Protocol for entity clients that can fetch entities by workspace/name.""" + + async def get(self, entity_type: Type[EntityT], *, name: str, workspace: str) -> EntityT: ... + + +class EntityDeleteClientProtocol(EntityGetterProtocol[EntityT], Protocol[EntityT]): + """Protocol for entity clients that can list and delete entities.""" - async def create(self, entity: EntityT) -> EntityT: ... async def list( self, - entity_type: EntityTypeLike, + entity_type: Type[EntityT], *, - workspace: Optional[str] = None, + workspace: str, filter_operation: Optional[FilterOperation] = None, - filter_str: Optional[str] = None, sort: Optional[str] = None, - filter_obj: Optional[Dict[str, Any]] = None, page: int = 1, page_size: int = 100, ) -> ListResponse[EntityT]: ... - async def count_by( + + async def delete( self, - entity_type: EntityTypeLike, - field: str, + entity_type: Type[EntityT], + name: str, *, - workspace: str = DEFAULT_WORKSPACE, - filter_obj: dict[str, Any] | None = None, - ) -> dict[str, int]: ... - async def get(self, entity_type: EntityTypeLike, name: str, *, workspace: Optional[str] = None) -> EntityT: ... - async def get_by_id(self, entity_type: EntityTypeLike, entity_id: str) -> EntityT: ... - async def update(self, entity: EntityT, *, original_name: str | None = None) -> EntityT: ... + workspace: str, + expected_db_version: Optional[int] = None, + ) -> object: ... + + +class EntityClientProtocol(EntityDeleteClientProtocol[EntityT], Protocol[EntityT]): + """Protocol for the common entity CRUD operations used by plugins.""" + + async def create(self, entity: EntityT) -> EntityT: ... + + +class AnyEntityGetterProtocol(Protocol): + """Protocol for clients that can fetch any entity model type.""" + + async def get(self, entity_type: Type[EntityT], *, name: str, workspace: str) -> EntityT: ... + + +class AnyEntityDeleteClientProtocol(AnyEntityGetterProtocol, Protocol): + """Protocol for clients that can list and delete any entity model type.""" + + async def list( + self, + entity_type: Type[EntityT], + *, + workspace: str, + filter_operation: Optional[FilterOperation] = None, + sort: Optional[str] = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[EntityT]: ... + async def delete( - self, entity_type: EntityTypeLike, name: str, *, workspace: Optional[str] = None - ) -> DeleteResponse: ... - async def delete_by_id(self, entity_type: EntityTypeLike, entity_id: str) -> DeleteResponse: ... - async def save(self, entity: EntityT) -> EntityT: ... - async def add(self, entity: EntityT) -> EntityT: ... - async def get_by_field( - self, entity_type: EntityTypeLike, *, workspace: Optional[str] = None, **field_filters: Any - ) -> EntityT: ... + self, + entity_type: Type[EntityT], + name: str, + *, + workspace: str, + expected_db_version: Optional[int] = None, + ) -> object: ... # Error types @@ -415,7 +444,8 @@ def _convert_api_entity_to_model(self, entity: Entity, entity_type: EntityTypeLi # the real caller is in X-NMP-Principal-On-Behalf-Of. if hasattr(result, "_auth_context"): sdk_headers = self.entities_api._client.default_headers - effective = sdk_headers.get("X-NMP-Principal-On-Behalf-Of") or sdk_headers.get("X-NMP-Principal-Id", "") + raw_effective = sdk_headers.get("X-NMP-Principal-On-Behalf-Of") or sdk_headers.get("X-NMP-Principal-Id", "") + effective = raw_effective if isinstance(raw_effective, str) else "" if not effective.startswith("service:"): setattr(result, "_auth_context", None) @@ -590,7 +620,7 @@ async def get( entity_name, workspace=ws, entity_type=_get_entity_type(entity_type), - parent=parent, + parent=parent if parent is not None else omit, ) return self._convert_api_entity_to_model(response, entity_type) except NotFoundError as e: @@ -649,7 +679,7 @@ async def update(self, entity: EntityT, *, original_name: str | None = None) -> entity_type=_get_entity_type(entity_type), data=entity._get_data_fields(), new_name=entity.name if original_name else omit, - parent=entity._parent, + parent=entity._parent if entity._parent is not None else omit, project=entity.project or omit, expected_db_version=entity.db_version, ) @@ -669,6 +699,7 @@ async def delete( *, workspace: Optional[str] = None, parent: Optional[str] = None, + expected_db_version: Optional[int] = None, ) -> DeleteResponse: """Delete an entity by name. @@ -679,12 +710,14 @@ async def delete( name: Entity name (can be workspace-qualified) workspace: Optional workspace override parent: Optional parent entity ID for nested entities + expected_db_version: Optional expected database version for optimistic locking Returns: Deleted entity response Raises: EntityNotFoundError: Entity not found + EntityConflictError: Version mismatch (entity was modified by another request) """ ws, entity_name = parse_qualified_name(name, default_workspace=workspace) try: @@ -692,10 +725,13 @@ async def delete( entity_name, workspace=ws, entity_type=_get_entity_type(entity_type), - parent=parent, + parent=parent if parent is not None else omit, + expected_db_version=expected_db_version if expected_db_version is not None else omit, ) except NotFoundError as e: raise EntityNotFoundError(f"Entity '{entity_name}' not found in workspace '{ws}'") from e + except ConflictError as e: + raise EntityConflictError(str(e)) from e async def delete_by_id( self, @@ -715,6 +751,7 @@ async def delete_by_id( Raises: EntityNotFoundError: Entity not found + EntityConflictError: Version mismatch (entity was modified by another request) """ try: entity = await self.entities_api.get_entity_by_id(entity_id) @@ -722,10 +759,13 @@ async def delete_by_id( entity.name, workspace=entity.workspace, entity_type=entity.entity_type, - parent=entity.parent, + parent=entity.parent if entity.parent is not None else omit, + expected_db_version=entity.db_version, ) except NotFoundError as e: raise EntityNotFoundError(f"Entity with id '{entity_id}' not found") from e + except ConflictError as e: + raise EntityConflictError(str(e)) from e async def save(self, entity: EntityT) -> EntityT: """ diff --git a/packages/nemo_platform_plugin/src/nemo_platform_plugin/entity_client.py b/packages/nemo_platform_plugin/src/nemo_platform_plugin/entity_client.py index 732481a4e9..bf6c6fc681 100644 --- a/packages/nemo_platform_plugin/src/nemo_platform_plugin/entity_client.py +++ b/packages/nemo_platform_plugin/src/nemo_platform_plugin/entity_client.py @@ -29,26 +29,48 @@ injects the real implementation at startup via ``app.dependency_overrides``. """ -from nemo_platform_plugin.dependencies import ( - get_entity_client, +from nemo_platform_plugin.dependencies import get_entity_client +from nemo_platform_plugin.entities import ( + AnyEntityDeleteClientProtocol as NemoAnyEntityDeleteClientProtocol, +) +from nemo_platform_plugin.entities import ( + AnyEntityGetterProtocol as NemoAnyEntityGetterProtocol, ) from nemo_platform_plugin.entities import ( EntityClient as NemoEntitiesClient, ) +from nemo_platform_plugin.entities import ( + EntityClientProtocol as NemoEntitiesClientProtocol, +) from nemo_platform_plugin.entities import ( EntityConflictError as NemoEntityConflictError, ) +from nemo_platform_plugin.entities import ( + EntityDeleteClientProtocol as NemoEntityDeleteClientProtocol, +) +from nemo_platform_plugin.entities import ( + EntityGetterProtocol as NemoEntityGetterProtocol, +) from nemo_platform_plugin.entities import ( EntityNotFoundError as NemoEntityNotFoundError, ) +from nemo_platform_plugin.entities import ( + EntityValidationError as NemoEntityValidationError, +) from nemo_platform_plugin.entities import ( PaginationInfo as NemoPaginationInfo, ) __all__ = [ "NemoEntitiesClient", + "NemoEntitiesClientProtocol", + "NemoAnyEntityDeleteClientProtocol", + "NemoAnyEntityGetterProtocol", + "NemoEntityDeleteClientProtocol", + "NemoEntityGetterProtocol", "NemoEntityConflictError", "NemoEntityNotFoundError", + "NemoEntityValidationError", "NemoPaginationInfo", "get_entity_client", ] diff --git a/packages/nmp_common/tests/entities/test_client.py b/packages/nmp_common/tests/entities/test_client.py index 594a072b93..0083a7dd9b 100644 --- a/packages/nmp_common/tests/entities/test_client.py +++ b/packages/nmp_common/tests/entities/test_client.py @@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, Mock import pytest +from nemo_platform import omit from nemo_platform.types.entities import EntitiesPage, Entity from nemo_platform.types.shared.pagination_data import PaginationData from nemo_platform_plugin.entities import _convert_filter_obj_to_filter_str @@ -962,6 +963,98 @@ class TestEntity(EntityBase): assert entity.db_version == 1 +@pytest.mark.asyncio +async def test_delete_passes_expected_db_version(): + """Test delete includes db_version for optimistic locking when supplied.""" + + class TestEntity(EntityBase): + field_1: str + + mock_api = Mock() + mock_api.delete_entity_by_name = AsyncMock() + client = EntityClient(mock_api) + + await client.delete(TestEntity, "test-entity", workspace="test-workspace", expected_db_version=5) + + mock_api.delete_entity_by_name.assert_awaited_once() + call_kwargs = mock_api.delete_entity_by_name.call_args.kwargs + assert call_kwargs["expected_db_version"] == 5 + + +@pytest.mark.asyncio +async def test_delete_omits_expected_db_version_by_default(): + """Test delete remains unconditional unless a version guard is supplied.""" + + class TestEntity(EntityBase): + field_1: str + + mock_api = Mock() + mock_api.delete_entity_by_name = AsyncMock() + client = EntityClient(mock_api) + + await client.delete(TestEntity, "test-entity", workspace="test-workspace") + + mock_api.delete_entity_by_name.assert_awaited_once() + call_kwargs = mock_api.delete_entity_by_name.call_args.kwargs + assert call_kwargs["expected_db_version"] is omit + + +@pytest.mark.asyncio +async def test_delete_by_id_passes_fetched_db_version(): + """Test delete_by_id deletes the version of the entity it resolved by ID.""" + + class TestEntity(EntityBase): + field_1: str + + now = datetime.now() + mock_api = Mock() + mock_api.get_entity_by_id = AsyncMock( + return_value=Entity( + entity_type="test_entity", + name="test-entity", + workspace="test-workspace", + id="123", + created_at=now, + updated_at=now, + db_version=7, + data={"field_1": "value"}, + ) + ) + mock_api.delete_entity_by_name = AsyncMock() + client = EntityClient(mock_api) + + await client.delete_by_id(TestEntity, "123") + + mock_api.delete_entity_by_name.assert_awaited_once() + call_kwargs = mock_api.delete_entity_by_name.call_args.kwargs + assert call_kwargs["expected_db_version"] == 7 + + +@pytest.mark.asyncio +async def test_delete_version_mismatch_raises_conflict(): + """Test delete maps API conflicts to EntityConflictError.""" + + class TestEntity(EntityBase): + field_1: str + + from nemo_platform import ConflictError + + mock_response = Mock() + mock_response.status_code = 409 + mock_api = Mock() + mock_api.delete_entity_by_name = AsyncMock( + side_effect=ConflictError( + message="Entity was modified by another request.", + response=mock_response, + body=None, + ) + ) + client = EntityClient(mock_api) + + with pytest.raises(EntityConflictError): + await client.delete(TestEntity, "test-entity", workspace="test-workspace", expected_db_version=5) + + @pytest.mark.asyncio async def test_update_with_automatic_version_check(): """Test update automatically includes db_version for optimistic locking when entity was fetched.""" diff --git a/packages/nmp_testing/src/nmp/testing/utils.py b/packages/nmp_testing/src/nmp/testing/utils.py index 0b789a5a4d..064216ba49 100644 --- a/packages/nmp_testing/src/nmp/testing/utils.py +++ b/packages/nmp_testing/src/nmp/testing/utils.py @@ -673,7 +673,7 @@ def add_mock_provider( from nemo_platform.types.inference.virtual_model import VirtualModel as _SDKVirtualModel virtual_model_cache = global_virtual_model_cache() - now_iso = _datetime.now().isoformat() + now = _datetime.now() for entity_name in served_models: key = (workspace, entity_name) if key in virtual_model_cache.virtual_model_map: @@ -684,10 +684,11 @@ def add_mock_provider( workspace=workspace, name=entity_name, parent=workspace, + db_version=1, default_model_entity=f"{workspace}/{entity_name}", autoprovisioned=should_autoprovision_virtual_model, - created_at=now_iso, - updated_at=now_iso, + created_at=now, + updated_at=now, ) except RuntimeError: # From E2E tests, the local cache is not available (app runs in a separate process). diff --git a/plugins/example-plugin/src/nemo_example_plugin/middleware_service.py b/plugins/example-plugin/src/nemo_example_plugin/middleware_service.py index 061098bd79..358f41b686 100644 --- a/plugins/example-plugin/src/nemo_example_plugin/middleware_service.py +++ b/plugins/example-plugin/src/nemo_example_plugin/middleware_service.py @@ -256,6 +256,14 @@ async def delete_config( status_code=status.HTTP_404_NOT_FOUND, detail=f"Middleware config '{name}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=( + f"Middleware config '{name}' was modified by another request in workspace '{workspace}'. " + "Refresh the config and try again." + ), + ) from exc except Exception: logger.exception("Failed to delete middleware config '%s'", name) raise HTTPException(status_code=500, detail="Failed to delete middleware config.") diff --git a/plugins/example-plugin/src/nemo_example_plugin/service.py b/plugins/example-plugin/src/nemo_example_plugin/service.py index 1bb493a726..d1c2608a64 100644 --- a/plugins/example-plugin/src/nemo_example_plugin/service.py +++ b/plugins/example-plugin/src/nemo_example_plugin/service.py @@ -485,6 +485,14 @@ async def delete_item( status_code=404, detail=f"Item '{name}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=409, + detail=( + f"Item '{name}' was modified by another request in workspace '{workspace}'. " + "Refresh the item and try again." + ), + ) from exc except Exception as exc: logger.exception("Failed to delete item '%s'", name) raise HTTPException(status_code=500, detail="Failed to delete item.") from exc diff --git a/plugins/example-plugin/tests/test_service.py b/plugins/example-plugin/tests/test_service.py index 07e1d42667..5acbb90373 100644 --- a/plugins/example-plugin/tests/test_service.py +++ b/plugins/example-plugin/tests/test_service.py @@ -16,6 +16,8 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from nemo_example_plugin.entities import ExampleItem +from nemo_example_plugin.middleware_config import ExampleMiddlewareConfig +from nemo_example_plugin.middleware_service import _get_entity_client as _get_middleware_entity_client from nemo_example_plugin.service import ExampleService, _get_entity_client from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError, NemoPaginationInfo @@ -35,6 +37,14 @@ def _make_item(name: str = "item1", workspace: str = "default") -> ExampleItem: return item +def _make_middleware_config(name: str = "cfg-1", workspace: str = "default") -> ExampleMiddlewareConfig: + """Build a fake persisted ExampleMiddlewareConfig with store-populated fields.""" + config = ExampleMiddlewareConfig(name=name, workspace=workspace) + config._id = f"id-{name}" + config._created_at = NOW + return config + + def _make_list_response(items: list[ExampleItem]) -> MagicMock: """Build a fake ListResponse as returned by entity_client.list().""" resp = MagicMock() @@ -56,6 +66,7 @@ def _make_app(mock_client: AsyncMock) -> FastAPI: for spec in service.get_routers(): app.include_router(spec.router, prefix=spec.prefix) app.dependency_overrides[_get_entity_client] = lambda: mock_client + app.dependency_overrides[_get_middleware_entity_client] = lambda: mock_client return app @@ -229,12 +240,18 @@ def test_update_item_404() -> None: def test_delete_item_204() -> None: mock = AsyncMock() + mock.get.return_value = _make_item("widget") mock.delete.return_value = None client = TestClient(_make_app(mock)) resp = client.delete("/v2/workspaces/default/items/widget") assert resp.status_code == 204 + mock.delete.assert_awaited_once_with( + ExampleItem, + name="widget", + workspace="default", + ) def test_delete_item_404() -> None: @@ -247,6 +264,41 @@ def test_delete_item_404() -> None: assert resp.status_code == 404 +def test_delete_item_409_when_changed() -> None: + mock = AsyncMock() + mock.delete.side_effect = NemoEntityConflictError("changed") + + client = TestClient(_make_app(mock)) + resp = client.delete("/v2/workspaces/default/items/widget") + + assert resp.status_code == 409 + + +def test_delete_middleware_config_204() -> None: + mock = AsyncMock() + mock.delete.return_value = None + + client = TestClient(_make_app(mock)) + resp = client.delete("/v2/workspaces/default/middleware-configs/cfg-1") + + assert resp.status_code == 204 + mock.delete.assert_awaited_once_with( + ExampleMiddlewareConfig, + name="cfg-1", + workspace="default", + ) + + +def test_delete_middleware_config_409_when_changed() -> None: + mock = AsyncMock() + mock.delete.side_effect = NemoEntityConflictError("changed") + + client = TestClient(_make_app(mock)) + resp = client.delete("/v2/workspaces/default/middleware-configs/cfg-1") + + assert resp.status_code == 409 + + # --------------------------------------------------------------------------- # Entity computed fields — id and created_at are present # --------------------------------------------------------------------------- diff --git a/plugins/nemo-agents/openapi/openapi.yaml b/plugins/nemo-agents/openapi/openapi.yaml index 78f7368988..94866c05ae 100644 --- a/plugins/nemo-agents/openapi/openapi.yaml +++ b/plugins/nemo-agents/openapi/openapi.yaml @@ -2322,6 +2322,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -2332,6 +2337,7 @@ components: - updated_by - entity_id - parent + - db_version title: Agent description: "An agent definition \u2014 stores agent config and metadata.\n\ \nEntity type: ``agent``\nPrimary lookup: by ``name`` within a ``workspace``.\n\ @@ -2460,6 +2466,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -2470,6 +2481,7 @@ components: - updated_by - entity_id - parent + - db_version title: AgentDeployment description: "A running (or pending) deployment of an Agent.\n\nEntity type:\ \ ``agent_deployment``\nLifecycle: pending \u2192 starting \u2192 running\ diff --git a/plugins/nemo-agents/src/nemo_agents_plugin/api/v2/agents.py b/plugins/nemo-agents/src/nemo_agents_plugin/api/v2/agents.py index e294576939..50c8eb1dc8 100644 --- a/plugins/nemo-agents/src/nemo_agents_plugin/api/v2/agents.py +++ b/plugins/nemo-agents/src/nemo_agents_plugin/api/v2/agents.py @@ -180,6 +180,14 @@ async def delete_agent( status_code=404, detail=f"Agent '{name}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=409, + detail=( + f"Agent '{name}' was modified by another request in workspace '{workspace}'. " + "Refresh the agent and try again." + ), + ) from exc except Exception as exc: logger.exception("Failed to delete agent '%s'", name) raise HTTPException(status_code=500, detail="Failed to delete agent.") from exc diff --git a/plugins/nemo-agents/src/nemo_agents_plugin/runner/controller.py b/plugins/nemo-agents/src/nemo_agents_plugin/runner/controller.py index 5c94a70938..3f6f257a07 100644 --- a/plugins/nemo-agents/src/nemo_agents_plugin/runner/controller.py +++ b/plugins/nemo-agents/src/nemo_agents_plugin/runner/controller.py @@ -364,7 +364,12 @@ async def _delete_deployment(self, dep: AgentDeployment) -> None: self._starting_since.pop((dep.workspace, dep.name), None) try: - await self.entities.delete(AgentDeployment, name=dep.name, workspace=dep.workspace) + await self.entities.delete( + AgentDeployment, + name=dep.name, + workspace=dep.workspace, + expected_db_version=dep.db_version, + ) except Exception: logger.exception("Failed to delete deployment entity '%s'", dep.name) else: diff --git a/plugins/nemo-agents/src/nemo_agents_plugin/runner/deployments_backend.py b/plugins/nemo-agents/src/nemo_agents_plugin/runner/deployments_backend.py index bc4aac021c..0fd1918880 100644 --- a/plugins/nemo-agents/src/nemo_agents_plugin/runner/deployments_backend.py +++ b/plugins/nemo-agents/src/nemo_agents_plugin/runner/deployments_backend.py @@ -417,7 +417,12 @@ async def create_deployment( except Exception: # Avoid orphaning the config if Deployment create fails. try: - await entities.delete(DeploymentConfig, name=name, workspace=workspace) + await entities.delete( + DeploymentConfig, + name=name, + workspace=workspace, + expected_db_version=deployment_config.db_version, + ) except Exception: logger.exception( "Failed to clean up DeploymentConfig '%s/%s' after Deployment create failure", @@ -476,7 +481,13 @@ async def delete_deployment(self, workspace: str, name: str) -> bool: return False try: - await entities.delete(DeploymentConfig, name=name, workspace=workspace) + deployment_config = await entities.get(DeploymentConfig, name=name, workspace=workspace) + await entities.delete( + DeploymentConfig, + name=name, + workspace=workspace, + expected_db_version=deployment_config.db_version, + ) except NemoEntityNotFoundError: pass return True diff --git a/plugins/nemo-agents/tests/unit/test_agents_api.py b/plugins/nemo-agents/tests/unit/test_agents_api.py index fb2bc9846b..5da43f89a7 100644 --- a/plugins/nemo-agents/tests/unit/test_agents_api.py +++ b/plugins/nemo-agents/tests/unit/test_agents_api.py @@ -391,6 +391,7 @@ def _make_deployment( class TestDeleteAgent: def test_delete_existing_returns_204(self, client: TestClient, mock_entity_client: AsyncMock) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) + mock_entity_client.get = AsyncMock(return_value=_make_agent("calc")) mock_entity_client.delete = AsyncMock(return_value=None) resp = client.delete("/apis/agents/v2/workspaces/default/agents/calc") @@ -399,11 +400,16 @@ def test_delete_existing_returns_204(self, client: TestClient, mock_entity_clien def test_delete_calls_entity_client_delete(self, client: TestClient, mock_entity_client: AsyncMock) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) + mock_entity_client.get = AsyncMock(return_value=_make_agent("calc")) mock_entity_client.delete = AsyncMock(return_value=None) client.delete("/apis/agents/v2/workspaces/default/agents/calc") - mock_entity_client.delete.assert_called_once_with(Agent, name="calc", workspace="default") + mock_entity_client.delete.assert_called_once_with( + Agent, + name="calc", + workspace="default", + ) def test_delete_not_found_returns_404(self, client: TestClient, mock_entity_client: AsyncMock) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) @@ -413,6 +419,14 @@ def test_delete_not_found_returns_404(self, client: TestClient, mock_entity_clie assert resp.status_code == 404 + def test_delete_conflict_returns_409(self, client: TestClient, mock_entity_client: AsyncMock) -> None: + mock_entity_client.list = AsyncMock(return_value=_list_response([])) + mock_entity_client.delete = AsyncMock(side_effect=NemoEntityConflictError("changed")) + + resp = client.delete("/apis/agents/v2/workspaces/default/agents/calc") + + assert resp.status_code == 409 + def test_delete_blocks_on_running_deployment(self, client: TestClient, mock_entity_client: AsyncMock) -> None: """DELETE /agents/{name} returns 409 when a running deployment references the agent.""" dep = _make_deployment(agent="calc", status="running") @@ -445,6 +459,7 @@ def test_delete_allowed_when_only_failed_deployments( """failed deployments do not block deletion — they are terminal.""" dep = _make_deployment(agent="calc", status="failed") mock_entity_client.list = AsyncMock(return_value=_list_response([dep])) + mock_entity_client.get = AsyncMock(return_value=_make_agent("calc")) mock_entity_client.delete = AsyncMock(return_value=None) resp = client.delete("/apis/agents/v2/workspaces/default/agents/calc") @@ -457,6 +472,7 @@ def test_delete_allowed_when_only_deleting_deployments( """deleting deployments do not block — they are already being cleaned up.""" dep = _make_deployment(agent="calc", status="deleting") mock_entity_client.list = AsyncMock(return_value=_list_response([dep])) + mock_entity_client.get = AsyncMock(return_value=_make_agent("calc")) mock_entity_client.delete = AsyncMock(return_value=None) resp = client.delete("/apis/agents/v2/workspaces/default/agents/calc") @@ -469,6 +485,7 @@ def test_delete_only_checks_deployments_for_this_agent( """A running deployment for a *different* agent must not block deletion.""" other_dep = _make_deployment(name="other-dep", agent="other-agent", status="running") mock_entity_client.list = AsyncMock(return_value=_list_response([other_dep])) + mock_entity_client.get = AsyncMock(return_value=_make_agent("calc")) mock_entity_client.delete = AsyncMock(return_value=None) resp = client.delete("/apis/agents/v2/workspaces/default/agents/calc") @@ -477,6 +494,7 @@ def test_delete_only_checks_deployments_for_this_agent( def test_delete_server_error_on_delete_returns_500(self, client: TestClient, mock_entity_client: AsyncMock) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) + mock_entity_client.get = AsyncMock(return_value=_make_agent("calc")) mock_entity_client.delete = AsyncMock(side_effect=RuntimeError("db error")) resp = client.delete("/apis/agents/v2/workspaces/default/agents/calc") diff --git a/plugins/nemo-agents/tests/unit/test_runner_deployments.py b/plugins/nemo-agents/tests/unit/test_runner_deployments.py index 21e3e15d19..0c6d039017 100644 --- a/plugins/nemo-agents/tests/unit/test_runner_deployments.py +++ b/plugins/nemo-agents/tests/unit/test_runner_deployments.py @@ -419,6 +419,7 @@ async def test_create_deployment_cleans_config_on_deployment_failure() -> None: delete_call = entities.delete.await_args assert delete_call is not None assert delete_call.args[0] is DeploymentConfig + assert delete_call.kwargs["expected_db_version"] == 1 @pytest.mark.asyncio @@ -453,8 +454,10 @@ async def test_delete_waits_for_deployment_gone_before_config_delete() -> None: deployment_config="hello-dep", status="READY", ) - # First get returns the deployment; subsequent gets in the wait loop raise NotFound. - entities.get = AsyncMock(side_effect=[deployment, NemoEntityNotFoundError("gone")]) + deployment_config = DeploymentConfig(name="hello-dep", workspace="default") + # First get returns the deployment; the wait-loop get raises NotFound; final get + # fetches the config version used for the conditional delete. + entities.get = AsyncMock(side_effect=[deployment, NemoEntityNotFoundError("gone"), deployment_config]) entities.update = AsyncMock() entities.delete = AsyncMock() backend._entities = entities @@ -469,6 +472,7 @@ async def test_delete_waits_for_deployment_gone_before_config_delete() -> None: delete_call = entities.delete.await_args assert delete_call is not None assert delete_call.args[0] is DeploymentConfig + assert delete_call.kwargs["expected_db_version"] == 1 @pytest.mark.asyncio diff --git a/plugins/nemo-auditor/openapi/openapi.yaml b/plugins/nemo-auditor/openapi/openapi.yaml index 8977bea627..daefbaa577 100644 --- a/plugins/nemo-auditor/openapi/openapi.yaml +++ b/plugins/nemo-auditor/openapi/openapi.yaml @@ -890,6 +890,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -900,6 +905,7 @@ components: - updated_by - entity_id - parent + - db_version title: AuditConfigOutput description: Audit configuration stored in the entity store. AuditInputSpec: @@ -1427,6 +1433,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -1439,6 +1450,7 @@ components: - updated_by - entity_id - parent + - db_version title: AuditTargetOutput description: Audit target (model under test) stored in the entity store. ConfigFilter: diff --git a/plugins/nemo-auditor/src/nemo_auditor/api/v2/configs.py b/plugins/nemo-auditor/src/nemo_auditor/api/v2/configs.py index 6a7641cbc8..acc67003cd 100644 --- a/plugins/nemo-auditor/src/nemo_auditor/api/v2/configs.py +++ b/plugins/nemo-auditor/src/nemo_auditor/api/v2/configs.py @@ -203,6 +203,14 @@ async def delete_config( status_code=404, detail=f"AuditConfig '{name}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=409, + detail=( + f"AuditConfig '{name}' was modified by another request in workspace '{workspace}'. " + "Refresh the config and try again." + ), + ) from exc except Exception as exc: logger.exception("Failed to delete audit config '%s'", name) raise HTTPException(status_code=500, detail="Failed to delete audit config.") from exc diff --git a/plugins/nemo-auditor/src/nemo_auditor/api/v2/targets.py b/plugins/nemo-auditor/src/nemo_auditor/api/v2/targets.py index e42429498e..91e0fa743e 100644 --- a/plugins/nemo-auditor/src/nemo_auditor/api/v2/targets.py +++ b/plugins/nemo-auditor/src/nemo_auditor/api/v2/targets.py @@ -200,6 +200,14 @@ async def delete_target( status_code=404, detail=f"AuditTarget '{name}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=409, + detail=( + f"AuditTarget '{name}' was modified by another request in workspace '{workspace}'. " + "Refresh the target and try again." + ), + ) from exc except Exception as exc: logger.exception("Failed to delete audit target '%s'", name) raise HTTPException(status_code=500, detail="Failed to delete audit target.") from exc diff --git a/plugins/nemo-auditor/tests/test_api_configs.py b/plugins/nemo-auditor/tests/test_api_configs.py index 25e6124144..9c43e4b4da 100644 --- a/plugins/nemo-auditor/tests/test_api_configs.py +++ b/plugins/nemo-auditor/tests/test_api_configs.py @@ -13,6 +13,7 @@ import logging from datetime import datetime, timezone +from typing import Any from unittest.mock import AsyncMock, MagicMock import pytest @@ -52,6 +53,12 @@ def _list_response(items): return resp +def _await_args(mock: AsyncMock) -> Any: + args = mock.await_args + assert args is not None + return args + + @pytest.fixture def mock_entity_client() -> AsyncMock: return AsyncMock() @@ -97,7 +104,7 @@ def test_constructs_audit_config_with_path_workspace(self, client, mock_entity_c json={"name": "cfg-1"}, ) assert resp.status_code == 201, resp.text - sent = mock_entity_client.create.await_args.args[0] + sent = _await_args(mock_entity_client.create).args[0] assert isinstance(sent, AuditConfig) assert sent.name == "cfg-1" assert sent.workspace == "prod" @@ -162,7 +169,7 @@ def test_forwards_pagination_params(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) resp = client.get("/apis/auditor/v2/workspaces/default/configs?page=3&page_size=5&sort=name") assert resp.status_code == 200 - kwargs = mock_entity_client.list.await_args.kwargs + kwargs = _await_args(mock_entity_client.list).kwargs assert kwargs["page"] == 3 assert kwargs["page_size"] == 5 assert kwargs["sort"] == "name" @@ -180,7 +187,7 @@ def test_replaces_fields_and_returns_200(self, client, mock_entity_client) -> No ) assert resp.status_code == 200, resp.text assert resp.json()["description"] == "new" - sent = mock_entity_client.update.await_args.args[0] + sent = _await_args(mock_entity_client.update).args[0] assert sent.description == "new" assert sent.name == "cfg-1" @@ -238,23 +245,34 @@ def test_conflict_sanitizes_log_fields(self, client, mock_entity_client, caplog) class TestDeleteConfig: def test_returns_204(self, client, mock_entity_client) -> None: + mock_entity_client.get = AsyncMock(return_value=_make_config("cfg-1")) mock_entity_client.delete = AsyncMock(return_value=None) resp = client.delete("/apis/auditor/v2/workspaces/default/configs/cfg-1") assert resp.status_code == 204 assert resp.content == b"" + mock_entity_client.delete.assert_awaited_once_with( + AuditConfig, + name="cfg-1", + workspace="default", + ) def test_404_when_missing(self, client, mock_entity_client) -> None: mock_entity_client.delete = AsyncMock(side_effect=NemoEntityNotFoundError("nope")) resp = client.delete("/apis/auditor/v2/workspaces/default/configs/missing") assert resp.status_code == 404 + def test_409_when_changed(self, client, mock_entity_client) -> None: + mock_entity_client.delete = AsyncMock(side_effect=NemoEntityConflictError("changed")) + resp = client.delete("/apis/auditor/v2/workspaces/default/configs/cfg-1") + assert resp.status_code == 409 + class TestListConfigsFiltering: def test_filter_by_description_forwards_filter_obj(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([_make_config("a", description="prod")])) resp = client.get("/apis/auditor/v2/workspaces/default/configs?filter[description]=prod") assert resp.status_code == 200, resp.text - kwargs = mock_entity_client.list.await_args.kwargs + kwargs = _await_args(mock_entity_client.list).kwargs assert kwargs["filter_obj"] == {"description": "prod"} body = resp.json() assert body["filter"] == {"description": "prod"} @@ -264,7 +282,7 @@ def test_filter_by_project_narrows(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) resp = client.get("/apis/auditor/v2/workspaces/default/configs?filter[project]=team-a") assert resp.status_code == 200, resp.text - assert mock_entity_client.list.await_args.kwargs["filter_obj"] == {"project": "team-a"} + assert _await_args(mock_entity_client.list).kwargs["filter_obj"] == {"project": "team-a"} def test_filter_created_at_range_parses_gte_lte(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) @@ -274,7 +292,7 @@ def test_filter_created_at_range_parses_gte_lte(self, client, mock_entity_client "&filter[created_at][$lte]=2024-12-31T00:00:00Z" ) assert resp.status_code == 200, resp.text - filter_obj = mock_entity_client.list.await_args.kwargs["filter_obj"] + filter_obj = _await_args(mock_entity_client.list).kwargs["filter_obj"] assert set(filter_obj["created_at"].keys()) == {"$gte", "$lte"} assert filter_obj["created_at"]["$gte"].startswith("2024-01-01") assert filter_obj["created_at"]["$lte"].startswith("2024-12-31") @@ -288,7 +306,7 @@ def test_empty_filter_forwards_none(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([_make_config("a"), _make_config("b")])) resp = client.get("/apis/auditor/v2/workspaces/default/configs") assert resp.status_code == 200, resp.text - kwargs = mock_entity_client.list.await_args.kwargs + kwargs = _await_args(mock_entity_client.list).kwargs assert kwargs["filter_obj"] is None body = resp.json() assert [c["name"] for c in body["data"]] == ["a", "b"] diff --git a/plugins/nemo-auditor/tests/test_api_targets.py b/plugins/nemo-auditor/tests/test_api_targets.py index 7fd227ffcf..12b055d72a 100644 --- a/plugins/nemo-auditor/tests/test_api_targets.py +++ b/plugins/nemo-auditor/tests/test_api_targets.py @@ -7,6 +7,7 @@ import logging from datetime import datetime, timezone +from typing import Any from unittest.mock import AsyncMock, MagicMock import pytest @@ -43,6 +44,12 @@ def _list_response(items): return resp +def _await_args(mock: AsyncMock) -> Any: + args = mock.await_args + assert args is not None + return args + + @pytest.fixture def mock_entity_client() -> AsyncMock: return AsyncMock() @@ -183,22 +190,33 @@ def test_conflict_sanitizes_log_fields(self, client, mock_entity_client, caplog) class TestDeleteTarget: def test_returns_204(self, client, mock_entity_client) -> None: + mock_entity_client.get = AsyncMock(return_value=_make_target("tgt-1")) mock_entity_client.delete = AsyncMock(return_value=None) resp = client.delete("/apis/auditor/v2/workspaces/default/targets/tgt-1") assert resp.status_code == 204 + mock_entity_client.delete.assert_awaited_once_with( + AuditTarget, + name="tgt-1", + workspace="default", + ) def test_404_when_missing(self, client, mock_entity_client) -> None: mock_entity_client.delete = AsyncMock(side_effect=NemoEntityNotFoundError("nope")) resp = client.delete("/apis/auditor/v2/workspaces/default/targets/missing") assert resp.status_code == 404 + def test_409_when_changed(self, client, mock_entity_client) -> None: + mock_entity_client.delete = AsyncMock(side_effect=NemoEntityConflictError("changed")) + resp = client.delete("/apis/auditor/v2/workspaces/default/targets/tgt-1") + assert resp.status_code == 409 + class TestListTargetsFiltering: def test_filter_by_type_forwards_filter_obj(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([_make_target("a", type="nim")])) resp = client.get("/apis/auditor/v2/workspaces/default/targets?filter[type]=nim") assert resp.status_code == 200, resp.text - kwargs = mock_entity_client.list.await_args.kwargs + kwargs = _await_args(mock_entity_client.list).kwargs assert kwargs["filter_obj"] == {"type": "nim"} body = resp.json() assert body["filter"] == {"type": "nim"} @@ -208,20 +226,20 @@ def test_filter_by_model_forwards_filter_obj(self, client, mock_entity_client) - mock_entity_client.list = AsyncMock(return_value=_list_response([_make_target("a")])) resp = client.get("/apis/auditor/v2/workspaces/default/targets?filter[model]=meta/llama-3.1-8b-instruct") assert resp.status_code == 200, resp.text - kwargs = mock_entity_client.list.await_args.kwargs + kwargs = _await_args(mock_entity_client.list).kwargs assert kwargs["filter_obj"] == {"model": "meta/llama-3.1-8b-instruct"} def test_filter_by_description_forwards_filter_obj(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) resp = client.get("/apis/auditor/v2/workspaces/default/targets?filter[description]=prod") assert resp.status_code == 200, resp.text - assert mock_entity_client.list.await_args.kwargs["filter_obj"] == {"description": "prod"} + assert _await_args(mock_entity_client.list).kwargs["filter_obj"] == {"description": "prod"} def test_filter_by_project_narrows(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) resp = client.get("/apis/auditor/v2/workspaces/default/targets?filter[project]=team-a") assert resp.status_code == 200, resp.text - assert mock_entity_client.list.await_args.kwargs["filter_obj"] == {"project": "team-a"} + assert _await_args(mock_entity_client.list).kwargs["filter_obj"] == {"project": "team-a"} def test_filter_created_at_range_parses_gte_lte(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([])) @@ -231,7 +249,7 @@ def test_filter_created_at_range_parses_gte_lte(self, client, mock_entity_client "&filter[created_at][$lte]=2024-12-31T00:00:00Z" ) assert resp.status_code == 200, resp.text - filter_obj = mock_entity_client.list.await_args.kwargs["filter_obj"] + filter_obj = _await_args(mock_entity_client.list).kwargs["filter_obj"] assert set(filter_obj["created_at"].keys()) == {"$gte", "$lte"} assert filter_obj["created_at"]["$gte"].startswith("2024-01-01") assert filter_obj["created_at"]["$lte"].startswith("2024-12-31") @@ -245,7 +263,7 @@ def test_empty_filter_forwards_none(self, client, mock_entity_client) -> None: mock_entity_client.list = AsyncMock(return_value=_list_response([_make_target("a"), _make_target("b")])) resp = client.get("/apis/auditor/v2/workspaces/default/targets") assert resp.status_code == 200, resp.text - kwargs = mock_entity_client.list.await_args.kwargs + kwargs = _await_args(mock_entity_client.list).kwargs assert kwargs["filter_obj"] is None body = resp.json() assert [t["name"] for t in body["data"]] == ["a", "b"] diff --git a/plugins/nemo-deployments/openapi/openapi.yaml b/plugins/nemo-deployments/openapi/openapi.yaml index ffeacffba6..7947d04cec 100644 --- a/plugins/nemo-deployments/openapi/openapi.yaml +++ b/plugins/nemo-deployments/openapi/openapi.yaml @@ -822,6 +822,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -833,6 +838,7 @@ components: - updated_by - entity_id - parent + - db_version title: Deployment description: Desired and observed deployment state. DeploymentBackendConfigInput: @@ -956,6 +962,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -966,6 +977,7 @@ components: - updated_by - entity_id - parent + - db_version title: DeploymentConfig description: Immutable PodSpec-shaped deployment template. DockerDeploymentConfig: @@ -1619,6 +1631,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -1629,6 +1646,7 @@ components: - updated_by - entity_id - parent + - db_version title: Volume description: Persistent volume request and observed state. VolumeBackendConfig: diff --git a/plugins/nemo-deployments/src/nemo_deployments_plugin/api/v2/deployment_configs.py b/plugins/nemo-deployments/src/nemo_deployments_plugin/api/v2/deployment_configs.py index 02d9e9da5b..960e6f08a9 100644 --- a/plugins/nemo-deployments/src/nemo_deployments_plugin/api/v2/deployment_configs.py +++ b/plugins/nemo-deployments/src/nemo_deployments_plugin/api/v2/deployment_configs.py @@ -111,8 +111,9 @@ async def delete_deployment_config( name: str, entity_client: NemoEntitiesClient = Depends(get_entity_client), ) -> None: - # Best-effort referential check before delete. Entity store has no conditional - # delete API today, so this cannot be fully atomic against concurrent creates. + # Best-effort referential check before delete. The version check below protects the + # config row itself, but cannot make this relationship check atomic against a + # concurrent deployment create. referencing = await deployment_names_using_config( entity_client, workspace=workspace, @@ -134,3 +135,11 @@ async def delete_deployment_config( status_code=404, detail=f"DeploymentConfig '{name}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=409, + detail=( + f"DeploymentConfig '{name}' was modified by another request in workspace '{workspace}'. " + "Refresh the config and try again." + ), + ) from exc diff --git a/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/deployment_reconciler.py b/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/deployment_reconciler.py index cb59eced05..f0caf2fb2c 100644 --- a/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/deployment_reconciler.py +++ b/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/deployment_reconciler.py @@ -94,6 +94,22 @@ async def reconcile_one( *, deployments_by_name: dict[tuple[str, str], Deployment], volumes_by_name: dict[tuple[str, str], Volume], + ) -> None: + try: + await self._reconcile_one( + deployment, + deployments_by_name=deployments_by_name, + volumes_by_name=volumes_by_name, + ) + except NemoEntityConflictError: + logger.debug("Optimistic lock conflict on deployment %s - retry next cycle.", deployment_id(deployment)) + + async def _reconcile_one( + self, + deployment: Deployment, + *, + deployments_by_name: dict[tuple[str, str], Deployment], + volumes_by_name: dict[tuple[str, str], Volume], ) -> None: if deployment.desired_state == "STOPPED" or deployment.status == "DELETING": await self._reconcile_delete(deployment) @@ -179,13 +195,12 @@ async def _reconcile_create( labels=labels, backend_config=config.backend_config.model_dump(by_alias=True, exclude_none=True), ) - logger.info("Created deployment %s: %s", dep_id, status_update.status) - await self._update_deployment_status(deployment, status_update) - except NemoEntityConflictError: - raise except Exception as exc: logger.exception("Failed to create deployment %s", dep_id) await self._update_deployment_status_failure(deployment, f"Failed to create deployment: {exc}") + return + logger.info("Created deployment %s: %s", dep_id, status_update.status) + await self._update_deployment_status(deployment, status_update) async def _reconcile_delete(self, deployment: Deployment) -> None: dep_id = deployment_id(deployment) @@ -222,12 +237,17 @@ async def _reconcile_delete(self, deployment: Deployment) -> None: logger.warning("No executor for delete of %s — removing entity only", dep_id) try: - await self._entities.delete(Deployment, name=deployment.name, workspace=deployment.workspace) + await self._entities.delete( + Deployment, + name=deployment.name, + workspace=deployment.workspace, + expected_db_version=deployment.db_version, + ) logger.info("Deleted deployment entity %s", dep_id) except NemoEntityNotFoundError: logger.debug("Deployment entity %s already deleted", dep_id) except NemoEntityConflictError: - raise + logger.debug("Optimistic lock conflict deleting deployment %s - retry next cycle.", dep_id) except Exception: logger.exception("Failed to delete deployment entity %s", dep_id) @@ -282,14 +302,6 @@ async def _handle_drift( labels=labels, backend_config=config.backend_config.model_dump(by_alias=True, exclude_none=True), ) - message = ( - f"Recovering deployment — backend resources recreated " - f"(attempt {attempt}/{limits.max_attempts}). {status_update.status_message}" - ) - status_update = status_update.model_copy(update={"status_message": message}) - await self._update_deployment_status(deployment, status_update) - except NemoEntityConflictError: - raise except Exception as exc: logger.exception("Drift recovery failed for %s", dep_id) await self._update_deployment_status( @@ -299,6 +311,13 @@ async def _handle_drift( status_message=(f"Recovery attempt {attempt}/{limits.max_attempts} failed: {exc}. Will retry."), ), ) + return + message = ( + f"Recovering deployment — backend resources recreated " + f"(attempt {attempt}/{limits.max_attempts}). {status_update.status_message}" + ) + status_update = status_update.model_copy(update={"status_message": message}) + await self._update_deployment_status(deployment, status_update) def _controller_recovery_limits(self) -> DriftRecoveryLimits: ctrl = self._controller_config @@ -439,10 +458,7 @@ async def _update_deployment_status(self, deployment: Deployment, update: Backen await self._save(deployment) async def _save(self, deployment: Deployment) -> None: - try: - await self._entities.update(deployment) - except NemoEntityConflictError: - raise + await self._entities.update(deployment) def _starting_timestamp(deployment: Deployment) -> datetime | None: diff --git a/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/volume_reconciler.py b/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/volume_reconciler.py index 9e64df6885..2c68d0f835 100644 --- a/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/volume_reconciler.py +++ b/plugins/nemo-deployments/src/nemo_deployments_plugin/reconciler/volume_reconciler.py @@ -23,6 +23,12 @@ def __init__(self, entities: NemoEntitiesClient, registry: ExecutorRegistry) -> self._registry = registry async def reconcile_one(self, volume: Volume) -> None: + try: + await self._reconcile_one(volume) + except NemoEntityConflictError: + logger.debug("Optimistic lock conflict on volume %s/%s - retry next cycle.", volume.workspace, volume.name) + + async def _reconcile_one(self, volume: Volume) -> None: if volume.status == "DELETING": await self._reconcile_delete(volume) return @@ -57,12 +63,17 @@ async def _reconcile_delete(self, volume: Volume) -> None: return try: - await self._entities.delete(Volume, name=volume.name, workspace=volume.workspace) + await self._entities.delete( + Volume, + name=volume.name, + workspace=volume.workspace, + expected_db_version=volume.db_version, + ) logger.info("Deleted volume entity %s", volume_id) except NemoEntityNotFoundError: logger.debug("Volume entity %s already deleted", volume_id) except NemoEntityConflictError: - raise + logger.debug("Optimistic lock conflict deleting volume %s - retry next cycle.", volume_id) except Exception: logger.exception("Failed to delete volume entity %s", volume_id) @@ -76,16 +87,15 @@ async def _reconcile_create(self, volume: Volume, backend: DeploymentBackend) -> access_modes=list(volume.access_modes), backend_config=backend_config, ) - await self._update_volume_status(volume, update) - logger.info("Volume %s/%s created: %s", volume.workspace, volume.name, update.status) - except NemoEntityConflictError: - raise except Exception as exc: logger.exception("Failed to create volume %s/%s", volume.workspace, volume.name) await self._update_volume_status( volume, VolumeStatusUpdate(status="FAILED", status_message=f"Failed to create volume: {exc}"), ) + return + await self._update_volume_status(volume, update) + logger.info("Volume %s/%s created: %s", volume.workspace, volume.name, update.status) async def _reconcile_read(self, volume: Volume, backend: DeploymentBackend) -> None: backend_config = volume.backend_config.model_dump(by_alias=True, exclude_none=True) @@ -95,15 +105,14 @@ async def _reconcile_read(self, volume: Volume, backend: DeploymentBackend) -> N name=volume.name, backend_config=backend_config, ) - await self._update_volume_status(volume, update) - except NemoEntityConflictError: - raise except Exception as exc: logger.exception("Failed to read volume status %s/%s", volume.workspace, volume.name) await self._update_volume_status( volume, VolumeStatusUpdate(status="FAILED", status_message=f"Failed to read volume status: {exc}"), ) + return + await self._update_volume_status(volume, update) async def _update_volume_status(self, volume: Volume, update: VolumeStatusUpdate) -> None: if ( @@ -118,7 +127,4 @@ async def _update_volume_status(self, volume: Volume, update: VolumeStatusUpdate await self._save(volume) async def _save(self, volume: Volume) -> None: - try: - await self._entities.update(volume) - except NemoEntityConflictError: - raise + await self._entities.update(volume) diff --git a/plugins/nemo-deployments/tests/unit/reconciler/test_deployment_reconciler.py b/plugins/nemo-deployments/tests/unit/reconciler/test_deployment_reconciler.py index 4288e611d6..35a55b16a9 100644 --- a/plugins/nemo-deployments/tests/unit/reconciler/test_deployment_reconciler.py +++ b/plugins/nemo-deployments/tests/unit/reconciler/test_deployment_reconciler.py @@ -333,6 +333,29 @@ async def test_desired_stopped_deletes( mock_entities.delete.assert_awaited_once() +@pytest.mark.asyncio +async def test_desired_stopped_delete_conflict_is_handled_for_retry( + deployment_reconciler: DeploymentReconciler, + mock_backend: MockDeploymentBackend, + mock_entities: AsyncMock, +) -> None: + dep = make_deployment() + dep.desired_state = "STOPPED" + cfg = make_deployment_config() + deployment_reconciler.set_config_cache({("default", "cfg1"): cfg}) + mock_entities.delete.side_effect = NemoEntityConflictError("conflict") + + await deployment_reconciler.reconcile_one(dep, deployments_by_name={}, volumes_by_name=NO_VOLUMES) + + assert mock_backend.deployment_delete_calls == [("default", "dep1")] + mock_entities.delete.assert_awaited_once_with( + Deployment, + name="dep1", + workspace="default", + expected_db_version=dep.db_version, + ) + + @pytest.mark.asyncio async def test_delete_proceeds_when_config_missing( deployment_reconciler: DeploymentReconciler, @@ -535,7 +558,7 @@ async def test_drift_recovery_exhausted( @pytest.mark.asyncio -async def test_conflict_propagates_from_save( +async def test_conflict_from_save_is_handled_for_retry( deployment_reconciler: DeploymentReconciler, mock_entities: AsyncMock, ) -> None: @@ -544,12 +567,13 @@ async def test_conflict_propagates_from_save( deployment_reconciler.set_config_cache({("default", "cfg1"): cfg}) mock_entities.update.side_effect = NemoEntityConflictError("conflict") - with pytest.raises(NemoEntityConflictError): - await deployment_reconciler.reconcile_one( - dep, - deployments_by_name={}, - volumes_by_name=NO_VOLUMES, - ) + await deployment_reconciler.reconcile_one( + dep, + deployments_by_name={}, + volumes_by_name=NO_VOLUMES, + ) + + mock_entities.update.assert_awaited() @pytest.mark.asyncio diff --git a/plugins/nemo-deployments/tests/unit/reconciler/test_volume_reconciler.py b/plugins/nemo-deployments/tests/unit/reconciler/test_volume_reconciler.py index cc02353269..af35c8d4b6 100644 --- a/plugins/nemo-deployments/tests/unit/reconciler/test_volume_reconciler.py +++ b/plugins/nemo-deployments/tests/unit/reconciler/test_volume_reconciler.py @@ -7,6 +7,7 @@ from helpers import make_volume from nemo_deployments_plugin.backends.base import VolumeStatusUpdate from nemo_deployments_plugin.reconciler.volume_reconciler import VolumeReconciler +from nemo_platform_plugin.entity_client import NemoEntityConflictError from reconciler.conftest import MockDeploymentBackend @@ -58,6 +59,27 @@ async def test_deleting_volume_removes_backend_then_entity( mock_entities.delete.assert_awaited_once() +@pytest.mark.asyncio +async def test_deleting_volume_delete_conflict_is_handled_for_retry( + volume_reconciler: VolumeReconciler, + mock_backend: MockDeploymentBackend, + mock_entities: AsyncMock, +) -> None: + vol = make_volume() + vol.status = "DELETING" + mock_entities.delete.side_effect = NemoEntityConflictError("conflict") + + await volume_reconciler.reconcile_one(vol) + + assert mock_backend.volume_delete_calls == [("default", "vol1")] + mock_entities.delete.assert_awaited_once_with( + type(vol), + name="vol1", + workspace="default", + expected_db_version=vol.db_version, + ) + + @pytest.mark.asyncio async def test_deleting_volume_waits_for_executor( volume_reconciler: VolumeReconciler, @@ -73,3 +95,16 @@ async def test_deleting_volume_waits_for_executor( await reconciler.reconcile_one(vol) mock_entities.delete.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_volume_update_conflict_is_handled_for_retry( + volume_reconciler: VolumeReconciler, + mock_entities: AsyncMock, +) -> None: + vol = make_volume() + mock_entities.update.side_effect = NemoEntityConflictError("conflict") + + await volume_reconciler.reconcile_one(vol) + + mock_entities.update.assert_awaited() diff --git a/plugins/nemo-deployments/tests/unit/test_api_deployment_configs.py b/plugins/nemo-deployments/tests/unit/test_api_deployment_configs.py index 3333bb8df6..9c0742d849 100644 --- a/plugins/nemo-deployments/tests/unit/test_api_deployment_configs.py +++ b/plugins/nemo-deployments/tests/unit/test_api_deployment_configs.py @@ -50,6 +50,11 @@ def test_delete_deployment_config_204(client: TestClient, mock_entity_client: As mock_entity_client.list.return_value = list_response([]) resp = client.delete("/apis/deployments/v2/workspaces/default/deployment-configs/cfg1") assert resp.status_code == 204 + mock_entity_client.delete.assert_awaited_once_with( + configs_module.DeploymentConfig, + name="cfg1", + workspace="default", + ) def test_delete_deployment_config_409_when_referenced(client: TestClient, mock_entity_client: AsyncMock) -> None: @@ -60,6 +65,13 @@ def test_delete_deployment_config_409_when_referenced(client: TestClient, mock_e mock_entity_client.delete.assert_not_awaited() +def test_delete_deployment_config_409_when_changed(client: TestClient, mock_entity_client: AsyncMock) -> None: + mock_entity_client.list.return_value = list_response([]) + mock_entity_client.delete.side_effect = NemoEntityConflictError("changed") + resp = client.delete("/apis/deployments/v2/workspaces/default/deployment-configs/cfg1") + assert resp.status_code == 409 + + def test_create_deployment_config_409(client: TestClient, mock_entity_client: AsyncMock) -> None: mock_entity_client.create.side_effect = NemoEntityConflictError("exists") resp = client.post( diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/dependencies.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/dependencies.py index c039b88c47..05f39a8c46 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/dependencies.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/dependencies.py @@ -11,12 +11,12 @@ from nemo_evaluator.api.service.task_service import TaskService from nemo_evaluator.api.service.taskset_service import TasksetService from nemo_platform import AsyncNeMoPlatform -from nemo_platform_plugin.dependencies import get_entity_client, get_sdk_client -from nemo_platform_plugin.entities import EntityClient +from nemo_platform_plugin.dependencies import get_sdk_client +from nemo_platform_plugin.entity_client import NemoEntitiesClient, get_entity_client def get_metric_service( - entity_client: EntityClient = Depends(get_entity_client), + entity_client: NemoEntitiesClient = Depends(get_entity_client), sdk: AsyncNeMoPlatform = Depends(get_sdk_client), ) -> MetricService: """Provide a MetricService wired to the Entity Store and Files service.""" @@ -24,14 +24,14 @@ def get_metric_service( def get_result_service( - entity_client: EntityClient = Depends(get_entity_client), + entity_client: NemoEntitiesClient = Depends(get_entity_client), ) -> ResultService: """Provide a ResultService wired to the Entity Store (read-only over result entities).""" return ResultService(entity_client) def get_task_service( - entity_client: EntityClient = Depends(get_entity_client), + entity_client: NemoEntitiesClient = Depends(get_entity_client), metric_service: MetricService = Depends(get_metric_service), ) -> TaskService: """Provide a TaskService. It uses the MetricService to normalize inline task metrics into @@ -40,7 +40,7 @@ def get_task_service( def get_taskset_service( - entity_client: EntityClient = Depends(get_entity_client), + entity_client: NemoEntitiesClient = Depends(get_entity_client), task_service: TaskService = Depends(get_task_service), ) -> TasksetService: """Provide a TasksetService. It uses the TaskService to validate that each referenced task diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/metric_service.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/metric_service.py index 1fb0741651..dc70e08316 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/metric_service.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/metric_service.py @@ -34,10 +34,10 @@ from nemo_evaluator.shared.metric_bundles.bundles import MetricBundle as RuntimeMetricBundle from nemo_platform import AsyncNeMoPlatform from nemo_platform_plugin.api.filter import ComparisonOperation, FilterOperator, LogicalOperation -from nemo_platform_plugin.entities import ( - EntityClient, - EntityConflictError, - EntityNotFoundError, +from nemo_platform_plugin.entity_client import ( + NemoEntitiesClientProtocol, + NemoEntityConflictError, + NemoEntityNotFoundError, ) from nemo_platform_plugin.filter_ops import FilterOperation from nemo_platform_plugin.log_utils import sanitize_for_log @@ -124,7 +124,7 @@ def _entity_from_bundle( class MetricService: """Service layer for stored metric CRUD.""" - def __init__(self, entity_client: EntityClient, sdk: AsyncNeMoPlatform): + def __init__(self, entity_client: NemoEntitiesClientProtocol[MetricBundleEntity], sdk: AsyncNeMoPlatform): self.entity_client = entity_client self.sdk = sdk @@ -147,7 +147,7 @@ async def create_metric( try: await self.entity_client.get(MetricBundleEntity, name=name, workspace=workspace) raise ValueError(f"Metric with name '{name}' already exists in workspace '{workspace}'") - except EntityNotFoundError: + except NemoEntityNotFoundError: pass # Convert the wire DTO into the runtime bundle (JSON round-trip keeps the @@ -164,7 +164,7 @@ async def create_metric( try: created = await self.entity_client.create(entity) - except EntityConflictError as e: + except NemoEntityConflictError as e: # We own this uniquely-named fileset and the entity was not created; # deleting it cannot affect the existing metric's data. await self._discard_bundle(bundle_ref) @@ -201,7 +201,7 @@ async def store_derived_metric(self, metric: MetricInline, *, workspace: str) -> try: await self.entity_client.get(MetricBundleEntity, name=name, workspace=workspace) return ref - except EntityNotFoundError: + except NemoEntityNotFoundError: pass bundle_ref = await store_bundle(self.sdk, workspace, name, runtime_bundle) @@ -210,7 +210,7 @@ async def store_derived_metric(self, metric: MetricInline, *, workspace: str) -> ) try: await self.entity_client.create(entity) - except EntityConflictError: + except NemoEntityConflictError: # Raced another writer to the same content-addressed name; theirs is byte-identical, so # drop the fileset we just uploaded and reuse the existing entry. await self._discard_bundle(bundle_ref) @@ -224,7 +224,7 @@ async def get_metric(self, workspace: str, name: str) -> Metric | None: try: entity = await self.entity_client.get(MetricBundleEntity, workspace=workspace, name=name) return _entity_to_schema(entity) - except EntityNotFoundError: + except NemoEntityNotFoundError: return None async def list_metrics( @@ -269,12 +269,17 @@ async def delete_metric(self, workspace: str, name: str) -> bool: """Delete a stored metric and its backing bundle. Returns False if not found.""" try: entity = await self.entity_client.get(MetricBundleEntity, workspace=workspace, name=name) - except EntityNotFoundError: + except NemoEntityNotFoundError: return False try: - await self.entity_client.delete(MetricBundleEntity, name, workspace=workspace) - except EntityNotFoundError: + await self.entity_client.delete( + MetricBundleEntity, + entity.name, + workspace=workspace, + expected_db_version=entity.db_version, + ) + except NemoEntityNotFoundError: # Lost a delete race: another request already removed it. return False await self._discard_bundle(entity.bundle_ref) diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/result_service.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/result_service.py index aa559fdea5..a412867d27 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/result_service.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/result_service.py @@ -14,10 +14,12 @@ from __future__ import annotations from datetime import datetime +from typing import TypeVar from nemo_evaluator.api.schemas import AgentEvalResult, EvaluateResult from nemo_evaluator.entities import AgentEvalResultEntity, EvaluateResultEntity -from nemo_platform_plugin.entities import EntityBase, EntityClient, EntityNotFoundError, PaginationInfo +from nemo_platform_plugin.entities import PaginationInfo +from nemo_platform_plugin.entity_client import NemoAnyEntityDeleteClientProtocol, NemoEntityNotFoundError from nemo_platform_plugin.filter_ops import FilterOperation from nemo_platform_plugin.schema import Page, PaginationData @@ -80,10 +82,13 @@ def _pagination(src: PaginationInfo, current_page_size: int) -> PaginationData: ) +_ResultEntityT = TypeVar("_ResultEntityT", AgentEvalResultEntity, EvaluateResultEntity) + + class ResultService: """List/get/delete for persisted eval-result entities, exposed as API DTOs.""" - def __init__(self, entity_client: EntityClient): + def __init__(self, entity_client: NemoAnyEntityDeleteClientProtocol): self.entity_client = entity_client # --- agent-eval results -------------------------------------------------- @@ -111,7 +116,7 @@ async def list_agent_eval_results( async def get_agent_eval_result(self, workspace: str, name: str) -> AgentEvalResult | None: try: entity = await self.entity_client.get(AgentEvalResultEntity, workspace=workspace, name=name) - except EntityNotFoundError: + except NemoEntityNotFoundError: return None return _to_agent_eval(entity) @@ -143,17 +148,17 @@ async def list_eval_results( async def get_eval_result(self, workspace: str, name: str) -> EvaluateResult | None: try: entity = await self.entity_client.get(EvaluateResultEntity, workspace=workspace, name=name) - except EntityNotFoundError: + except NemoEntityNotFoundError: return None return _to_evaluate(entity) async def delete_eval_result(self, workspace: str, name: str) -> bool: return await self._delete(EvaluateResultEntity, workspace, name) - async def _delete(self, entity_cls: type[EntityBase], workspace: str, name: str) -> bool: + async def _delete(self, entity_cls: type[_ResultEntityT], workspace: str, name: str) -> bool: """Delete by workspace/name; ``False`` if absent. Type-agnostic (delete takes no body).""" try: await self.entity_client.delete(entity_cls, name, workspace=workspace) - except EntityNotFoundError: + except NemoEntityNotFoundError: return False return True diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/task_service.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/task_service.py index f76dc2dc92..a64e4ceb58 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/task_service.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/task_service.py @@ -13,11 +13,16 @@ from __future__ import annotations import logging +from typing import Protocol from nemo_evaluator.api.schemas import MetricInline, MetricRef, Task, TaskInput, parse_entity_ref -from nemo_evaluator.api.service.metric_service import MetricService from nemo_evaluator.entities import TaskEntity -from nemo_platform_plugin.entities import EntityClient, EntityConflictError, EntityNotFoundError, PaginationInfo +from nemo_platform_plugin.entities import PaginationInfo +from nemo_platform_plugin.entity_client import ( + NemoEntitiesClientProtocol, + NemoEntityConflictError, + NemoEntityNotFoundError, +) from nemo_platform_plugin.filter_ops import FilterOperation from nemo_platform_plugin.log_utils import sanitize_for_log from nemo_platform_plugin.schema import Page, PaginationData @@ -25,6 +30,12 @@ logger = logging.getLogger(__name__) +class _MetricService(Protocol): + async def store_derived_metric(self, metric: MetricInline, *, workspace: str) -> MetricRef: ... + + async def get_metric(self, workspace: str, name: str) -> object | None: ... + + class MetricRefNotFoundError(ValueError): """A task references a stored metric that does not exist.""" @@ -64,7 +75,7 @@ def _pagination(src: PaginationInfo, current_page_size: int) -> PaginationData: class TaskService: """Create/get/list/delete for persisted agent-eval task entities, exposed as the ``Task`` DTO.""" - def __init__(self, entity_client: EntityClient, metric_service: MetricService): + def __init__(self, entity_client: NemoEntitiesClientProtocol[TaskEntity], metric_service: _MetricService): self.entity_client = entity_client self.metric_service = metric_service @@ -102,7 +113,7 @@ async def create_task( ) try: created = await self.entity_client.create(entity) - except EntityConflictError as exc: + except NemoEntityConflictError as exc: raise ValueError(f"Task '{workspace}/{name}' already exists") from exc logger.info( "Task created", extra={"workspace": sanitize_for_log(workspace), "task_name": sanitize_for_log(name)} @@ -112,7 +123,7 @@ async def create_task( async def get_task(self, workspace: str, name: str) -> Task | None: try: entity = await self.entity_client.get(TaskEntity, workspace=workspace, name=name) - except EntityNotFoundError: + except NemoEntityNotFoundError: return None return _entity_to_task(entity) @@ -140,7 +151,7 @@ async def delete_task(self, workspace: str, name: str) -> bool: """Delete a stored task; ``False`` if absent.""" try: await self.entity_client.delete(TaskEntity, name, workspace=workspace) - except EntityNotFoundError: + except NemoEntityNotFoundError: return False logger.info( "Task deleted", extra={"workspace": sanitize_for_log(workspace), "task_name": sanitize_for_log(name)} diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/taskset_service.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/taskset_service.py index d1ecfcf84a..adb404c9ea 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/service/taskset_service.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/service/taskset_service.py @@ -16,11 +16,16 @@ from __future__ import annotations import logging +from typing import Protocol from nemo_evaluator.api.schemas import TaskRef, Taskset, TasksetInput, parse_entity_ref -from nemo_evaluator.api.service.task_service import TaskService from nemo_evaluator.entities import TasksetEntity -from nemo_platform_plugin.entities import EntityClient, EntityConflictError, EntityNotFoundError, PaginationInfo +from nemo_platform_plugin.entities import PaginationInfo +from nemo_platform_plugin.entity_client import ( + NemoEntitiesClientProtocol, + NemoEntityConflictError, + NemoEntityNotFoundError, +) from nemo_platform_plugin.filter_ops import FilterOperation from nemo_platform_plugin.log_utils import sanitize_for_log from nemo_platform_plugin.schema import Page, PaginationData @@ -28,6 +33,10 @@ logger = logging.getLogger(__name__) +class _TaskService(Protocol): + async def get_task(self, workspace: str, name: str) -> object | None: ... + + class TaskRefNotFoundError(ValueError): """A taskset references a task that does not exist. @@ -86,7 +95,7 @@ def _pagination(src: PaginationInfo, current_page_size: int) -> PaginationData: class TasksetService: """Create/get/list/delete for persisted taskset entities, exposed as the ``Taskset`` DTO.""" - def __init__(self, entity_client: EntityClient, task_service: TaskService): + def __init__(self, entity_client: NemoEntitiesClientProtocol[TasksetEntity], task_service: _TaskService): self.entity_client = entity_client self.task_service = task_service @@ -128,7 +137,7 @@ async def create_taskset( ) try: created = await self.entity_client.create(entity) - except EntityConflictError as exc: + except NemoEntityConflictError as exc: raise TasksetExistsError(f"Taskset '{workspace}/{name}' already exists") from exc logger.info( "Taskset created", @@ -139,7 +148,7 @@ async def create_taskset( async def get_taskset(self, workspace: str, name: str) -> Taskset | None: try: entity = await self.entity_client.get(TasksetEntity, workspace=workspace, name=name) - except EntityNotFoundError: + except NemoEntityNotFoundError: return None return _entity_to_taskset(entity) @@ -167,7 +176,7 @@ async def delete_taskset(self, workspace: str, name: str) -> bool: """Delete a stored taskset; ``False`` if absent.""" try: await self.entity_client.delete(TasksetEntity, name, workspace=workspace) - except EntityNotFoundError: + except NemoEntityNotFoundError: return False logger.info( "Taskset deleted", diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/metrics.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/metrics.py index 135c706de3..6b12570ad1 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/metrics.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/metrics.py @@ -22,6 +22,7 @@ from nemo_platform_plugin.api.parsed_filter import ParsedFilter, make_filter_dep from nemo_platform_plugin.authz import CallerKind, PermissionSet, path_rule, perm from nemo_platform_plugin.entities import EntityValidationError +from nemo_platform_plugin.entity_client import NemoEntityConflictError from nemo_platform_plugin.jobs.openapi_utils import generate_openapi_extra_params from nemo_platform_plugin.log_utils import sanitize_for_log from nemo_platform_plugin.schema import Page @@ -185,6 +186,10 @@ async def delete_metric( return None except HTTPException: raise + except NemoEntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception(f"Failed to delete metric {sanitize_for_log(workspace)}/{sanitize_for_log(name)}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Internal server error") diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/results.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/results.py index 7bb5810a5d..57075b142c 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/results.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/results.py @@ -25,6 +25,7 @@ from nemo_evaluator.authz import scope from nemo_platform_plugin.api.parsed_filter import ParsedFilter, make_filter_dep from nemo_platform_plugin.authz import CallerKind, PermissionSet, path_rule, perm +from nemo_platform_plugin.entity_client import NemoEntityConflictError from nemo_platform_plugin.jobs.openapi_utils import generate_openapi_extra_params from nemo_platform_plugin.log_utils import sanitize_for_log from nemo_platform_plugin.schema import DatetimeFilter, Page @@ -164,6 +165,11 @@ async def delete_agent_eval_result( return None except HTTPException: raise + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Result was modified by another request: {workspace}/{name}. Refresh and try again.", + ) from exc except Exception: logger.exception(f"Failed to delete agent-eval result {sanitize_for_log(workspace)}/{sanitize_for_log(name)}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Internal server error") @@ -252,6 +258,11 @@ async def delete_eval_result( return None except HTTPException: raise + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Result was modified by another request: {workspace}/{name}. Refresh and try again.", + ) from exc except Exception: logger.exception(f"Failed to delete eval result {sanitize_for_log(workspace)}/{sanitize_for_log(name)}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Internal server error") diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasks.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasks.py index 344190073a..a07bd3c4c3 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasks.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasks.py @@ -17,6 +17,7 @@ from nemo_platform_plugin.api.parsed_filter import ParsedFilter, make_filter_dep from nemo_platform_plugin.authz import CallerKind, PermissionSet, path_rule, perm from nemo_platform_plugin.entities import EntityValidationError +from nemo_platform_plugin.entity_client import NemoEntityConflictError from nemo_platform_plugin.jobs.openapi_utils import generate_openapi_extra_params from nemo_platform_plugin.log_utils import sanitize_for_log from nemo_platform_plugin.schema import Page @@ -178,6 +179,11 @@ async def delete_task( return None except HTTPException: raise + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Task was modified by another request: {workspace}/{name}. Refresh and try again.", + ) from exc except Exception: logger.exception(f"Failed to delete task {sanitize_for_log(workspace)}/{sanitize_for_log(name)}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Internal server error") diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasksets.py b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasksets.py index f1cf350527..697e3dbc43 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasksets.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/api/v2/tasksets.py @@ -22,6 +22,7 @@ from nemo_platform_plugin.api.parsed_filter import ParsedFilter, make_filter_dep from nemo_platform_plugin.authz import CallerKind, PermissionSet, path_rule, perm from nemo_platform_plugin.entities import EntityValidationError +from nemo_platform_plugin.entity_client import NemoEntityConflictError from nemo_platform_plugin.jobs.openapi_utils import generate_openapi_extra_params from nemo_platform_plugin.log_utils import sanitize_for_log from nemo_platform_plugin.schema import Page @@ -174,6 +175,11 @@ async def delete_taskset( # Wrap only the service call so the 404 below is raised outside the try (no catch-and-re-raise). try: deleted = await service.delete_taskset(workspace, name) + except NemoEntityConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=f"Taskset was modified by another request: {workspace}/{name}. Refresh and try again.", + ) from exc except Exception: logger.exception(f"Failed to delete taskset {sanitize_for_log(workspace)}/{sanitize_for_log(name)}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Internal server error") diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/metric_refs.py b/plugins/nemo-evaluator/src/nemo_evaluator/metric_refs.py index 5c6daeb63a..e46f6901f7 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/metric_refs.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/metric_refs.py @@ -12,8 +12,6 @@ from __future__ import annotations -from typing import Any - # ``MetricRef`` / ``MetricRefOrInline`` are defined in ``api.schemas`` (next to ``MetricInline``) so # entity/DTO modules can reference them without importing this module's entities-dependent resolution # logic (which would create an import cycle). Imported here for use below and re-exported for the @@ -23,7 +21,7 @@ from nemo_evaluator.metric_storage import load_bundle from nemo_evaluator.shared.metric_bundles.bundles import MetricBundle from nemo_platform import AsyncNeMoPlatform -from nemo_platform_plugin.entities import EntityNotFoundError +from nemo_platform_plugin.entity_client import NemoEntityGetterProtocol, NemoEntityNotFoundError def parse_metric_ref(root: str, default_workspace: str) -> tuple[str, str]: @@ -39,7 +37,7 @@ async def resolve_metric_ref( ref: MetricRef, *, workspace: str, - entity_client: Any, + entity_client: NemoEntityGetterProtocol[MetricBundleEntity] | None, async_sdk: AsyncNeMoPlatform | None, ) -> MetricBundle: """Load and reconstruct the stored metric a reference points at.""" @@ -51,7 +49,7 @@ async def resolve_metric_ref( ref_workspace, name = parse_metric_ref(ref.root, workspace) try: entity = await entity_client.get(MetricBundleEntity, name=name, workspace=ref_workspace) - except EntityNotFoundError as exc: + except NemoEntityNotFoundError as exc: raise ValueError( f"Metric reference '{ref.root}' not found. " f"Ensure a stored metric named '{name}' exists in workspace '{ref_workspace}', " @@ -64,7 +62,7 @@ async def resolve_metric_specs( metrics: list[MetricRefOrInline], *, workspace: str, - entity_client: Any, + entity_client: NemoEntityGetterProtocol[MetricBundleEntity] | None, async_sdk: AsyncNeMoPlatform | None, ) -> list[MetricBundle]: """Resolve a wire metric list into runtime bundles. diff --git a/plugins/nemo-evaluator/src/nemo_evaluator/task_refs.py b/plugins/nemo-evaluator/src/nemo_evaluator/task_refs.py index 7447d1462e..0ce535fdc5 100644 --- a/plugins/nemo-evaluator/src/nemo_evaluator/task_refs.py +++ b/plugins/nemo-evaluator/src/nemo_evaluator/task_refs.py @@ -19,7 +19,7 @@ from nemo_evaluator.api.schemas import TasksetRef, parse_entity_ref from nemo_evaluator.entities import TaskEntity, TasksetEntity from nemo_evaluator.jobs.agent_spec import AgentEvalTaskInput -from nemo_platform_plugin.entities import EntityClient, EntityNotFoundError +from nemo_platform_plugin.entity_client import NemoAnyEntityGetterProtocol, NemoEntityNotFoundError def _entity_to_task_input(entity: TaskEntity) -> AgentEvalTaskInput: @@ -44,7 +44,7 @@ async def resolve_taskset_ref( ref: TasksetRef, *, workspace: str, - entity_client: EntityClient | None, + entity_client: NemoAnyEntityGetterProtocol | None, ) -> list[AgentEvalTaskInput]: """Load a stored taskset and expand its members into inline task DTOs. @@ -59,7 +59,7 @@ async def resolve_taskset_ref( ref_workspace, name = parse_entity_ref(ref.root, workspace) try: taskset = await entity_client.get(TasksetEntity, name=name, workspace=ref_workspace) - except EntityNotFoundError as exc: + except NemoEntityNotFoundError as exc: raise ValueError( f"Taskset reference '{ref.root}' not found. " f"Ensure a stored taskset named '{name}' exists in workspace '{ref_workspace}', " @@ -75,7 +75,7 @@ async def resolve_taskset_ref( task_workspace, task_name = parse_entity_ref(task_ref.root, ref_workspace) try: entity = await entity_client.get(TaskEntity, name=task_name, workspace=task_workspace) - except EntityNotFoundError as exc: + except NemoEntityNotFoundError as exc: raise ValueError( f"Task '{task_ref.root}' referenced by taskset '{ref.root}' was not found; " "the stored task may have been deleted after the taskset was created." @@ -97,7 +97,7 @@ async def resolve_agent_eval_tasks( tasks: TasksetRef | list[AgentEvalTaskInput], *, workspace: str, - entity_client: EntityClient | None, + entity_client: NemoAnyEntityGetterProtocol | None, ) -> list[AgentEvalTaskInput]: """Normalize an agent-eval ``tasks`` field to an inline task list. diff --git a/plugins/nemo-evaluator/tests/api/service/test_metric_service.py b/plugins/nemo-evaluator/tests/api/service/test_metric_service.py index edf617f6f7..7e38f597f2 100644 --- a/plugins/nemo-evaluator/tests/api/service/test_metric_service.py +++ b/plugins/nemo-evaluator/tests/api/service/test_metric_service.py @@ -3,8 +3,9 @@ from __future__ import annotations +from collections.abc import Iterator from datetime import datetime, timezone -from unittest.mock import AsyncMock, patch +from unittest.mock import patch import pytest from nemo_evaluator.api.schemas import MetricInline @@ -14,12 +15,11 @@ from nemo_evaluator.shared.metric_bundles.bundles import bundle_metric from nemo_evaluator.shared.metric_bundles.cloudpickle import CloudpickleMetricBundlePackager from nemo_evaluator_sdk.metrics.exact_match import ExactMatchMetric -from nemo_platform_plugin.entities import ( - EntityConflictError, - EntityNotFoundError, - ListResponse, - PaginationInfo, -) +from nemo_platform import AsyncNeMoPlatform +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError +from nemo_platform_plugin.files.types import CreateFilesetRequest +from nemo_platform_plugin.filter_ops import FilterOperation # ---- in-memory fakes ------------------------------------------------------- @@ -32,40 +32,49 @@ async def read(self) -> bytes: return self._data +class _FakeOperationResponse: + def data(self) -> object: + return object() + + class _FakeAsyncFilesClient: def __init__(self) -> None: self._store: dict[tuple[str, str], dict[str, bytes]] = {} - async def create_fileset(self, *, body, workspace=None, exist_ok=False): - self._store.setdefault((workspace, body.name), {}) - return AsyncMock(data=lambda: object()) + async def create_fileset( + self, *, body: CreateFilesetRequest, workspace: str | None = None, exist_ok: bool = False + ) -> _FakeOperationResponse: + self._store.setdefault((workspace or "default", body.name), {}) + return _FakeOperationResponse() - async def delete_fileset(self, *, name, workspace=None): - self._store.pop((workspace, name), None) - return AsyncMock(data=lambda: object()) + async def delete_fileset(self, *, name: str, workspace: str | None = None) -> _FakeOperationResponse: + self._store.pop((workspace or "default", name), None) + return _FakeOperationResponse() - async def upload_file(self, *, path, content, workspace, name): + async def upload_file(self, *, path: str, content: bytes, workspace: str, name: str) -> _FakeOperationResponse: self._store.setdefault((workspace, name), {})[path] = bytes(content) - return AsyncMock(data=lambda: object()) + return _FakeOperationResponse() - async def download_file(self, *, path, workspace, name): + async def download_file(self, *, path: str, workspace: str, name: str) -> _FakeResponse: return _FakeResponse(self._store[(workspace, name)][path]) class _FakeEntityClient: def __init__(self) -> None: self.entities: dict[tuple[str, str], MetricBundleEntity] = {} + self.delete_error: Exception | None = None + self.list_filter_operations: list[FilterOperation | None] = [] - async def get(self, entity_cls, *, workspace, name): + async def get(self, entity_type: type[MetricBundleEntity], *, workspace: str, name: str) -> MetricBundleEntity: key = (workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") return self.entities[key] - async def create(self, entity): + async def create(self, entity: MetricBundleEntity) -> MetricBundleEntity: key = (entity.workspace, entity.name) if key in self.entities: - raise EntityConflictError(f"{key} exists") + raise NemoEntityConflictError(f"{key} exists") now = datetime.now(timezone.utc) entity._id = f"metric_bundle-{entity.name}" entity._created_at = now @@ -73,10 +82,29 @@ async def create(self, entity): self.entities[key] = entity return entity - async def delete(self, entity_cls, name, *, workspace): + async def delete( + self, + entity_type: type[MetricBundleEntity], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + if self.delete_error is not None: + raise self.delete_error self.entities.pop((workspace, name), None) - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): + async def list( + self, + entity_type: type[MetricBundleEntity], + *, + workspace: str, + filter_operation: FilterOperation | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[MetricBundleEntity]: + self.list_filter_operations.append(filter_operation) items = [e for (ws, _), e in self.entities.items() if ws == workspace] return ListResponse( data=items, @@ -90,15 +118,27 @@ async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, ) +class _FakePlatform(AsyncNeMoPlatform): + pass + + +def _fake_platform() -> _FakePlatform: + return _FakePlatform.__new__(_FakePlatform) + + @pytest.fixture -def fake_files(): +def fake_files() -> _FakeAsyncFilesClient: return _FakeAsyncFilesClient() @pytest.fixture -def service(fake_files): - svc = MetricService(_FakeEntityClient(), object()) - svc._fake_files = fake_files +def fake_entity_client() -> _FakeEntityClient: + return _FakeEntityClient() + + +@pytest.fixture +def service(fake_files: _FakeAsyncFilesClient, fake_entity_client: _FakeEntityClient) -> Iterator[MetricService]: + svc = MetricService(fake_entity_client, _fake_platform()) with patch("nemo_evaluator.metric_storage.client_from_platform", return_value=fake_files): yield svc @@ -157,13 +197,12 @@ async def test_delete_returns_false_when_missing(service: MetricService) -> None assert await service.delete_metric("default", "nope") is False -async def test_delete_handles_concurrent_delete_race(service: MetricService) -> None: +async def test_delete_handles_concurrent_delete_race( + service: MetricService, fake_entity_client: _FakeEntityClient +) -> None: await service.create_metric("m", _bundle(), workspace="default") - async def _already_deleted(*_args, **_kwargs): - raise EntityNotFoundError("deleted concurrently") - - service.entity_client.delete = _already_deleted + fake_entity_client.delete_error = NemoEntityNotFoundError("deleted concurrently") assert await service.delete_metric("default", "m") is False @@ -181,7 +220,9 @@ async def test_list_returns_workspace_metrics(service: MetricService) -> None: # ---- derived metrics ------------------------------------------------------- -async def test_store_derived_metric_names_by_digest_and_marks_derived(service: MetricService, fake_files) -> None: +async def test_store_derived_metric_names_by_digest_and_marks_derived( + service: MetricService, fake_files: _FakeAsyncFilesClient, fake_entity_client: _FakeEntityClient +) -> None: from nemo_evaluator.api.service.metric_service import _MAX_ENTITY_NAME_LENGTH ref = await service.store_derived_metric(_bundle(), workspace="default") @@ -190,12 +231,14 @@ async def test_store_derived_metric_names_by_digest_and_marks_derived(service: M assert workspace == "default" assert name.startswith("derived.") assert len(name) <= _MAX_ENTITY_NAME_LENGTH - entity = service.entity_client.entities[("default", name)] + entity = fake_entity_client.entities[("default", name)] assert entity.derived is True assert _fileset_of(service, entity.bundle_ref) in fake_files._store -async def test_store_derived_metric_distinguishes_full_contract(service: MetricService) -> None: +async def test_store_derived_metric_distinguishes_full_contract( + service: MetricService, fake_entity_client: _FakeEntityClient +) -> None: bundle = _bundle() variant = bundle.model_copy(update={"metadata": bundle.metadata.model_copy(update={"description": "different"})}) assert bundle.payload.digest == variant.payload.digest @@ -204,32 +247,25 @@ async def test_store_derived_metric_distinguishes_full_contract(service: MetricS second = await service.store_derived_metric(variant, workspace="default") assert first.root != second.root - assert len(service.entity_client.entities) == 2 + assert len(fake_entity_client.entities) == 2 -async def test_store_derived_metric_is_content_addressed_dedup(service: MetricService, fake_files) -> None: +async def test_store_derived_metric_is_content_addressed_dedup( + service: MetricService, fake_files: _FakeAsyncFilesClient, fake_entity_client: _FakeEntityClient +) -> None: bundle = _bundle() first = await service.store_derived_metric(bundle, workspace="default") second = await service.store_derived_metric(bundle, workspace="default") assert first.root == second.root - assert len(service.entity_client.entities) == 1 + assert len(fake_entity_client.entities) == 1 assert len(fake_files._store) == 1 -async def test_list_excludes_derived_by_default(service: MetricService) -> None: - captured: list[object] = [] - original_list = service.entity_client.list - - async def _spy(entity_cls, *, filter_operation=None, **kwargs): - captured.append(filter_operation) - return await original_list(entity_cls, filter_operation=filter_operation, **kwargs) - - service.entity_client.list = _spy - +async def test_list_excludes_derived_by_default(service: MetricService, fake_entity_client: _FakeEntityClient) -> None: await service.list_metrics("default") await service.list_metrics("default", include_derived=True) - assert captured[0] is not None - assert captured[1] is None + assert fake_entity_client.list_filter_operations[0] is not None + assert fake_entity_client.list_filter_operations[1] is None diff --git a/plugins/nemo-evaluator/tests/api/service/test_result_service.py b/plugins/nemo-evaluator/tests/api/service/test_result_service.py index b0b7f37874..ec1cecc661 100644 --- a/plugins/nemo-evaluator/tests/api/service/test_result_service.py +++ b/plugins/nemo-evaluator/tests/api/service/test_result_service.py @@ -10,22 +10,26 @@ from __future__ import annotations from datetime import datetime, timezone +from typing import TypeVar import pytest from nemo_evaluator.api.schemas import AgentEvalResult, EvaluateResult from nemo_evaluator.api.service.result_service import ResultService from nemo_evaluator.entities import AgentEvalResultEntity, EvaluateResultEntity from nemo_evaluator_sdk.values.results import AggregatedMetricResult -from nemo_platform_plugin.entities import EntityBase, EntityNotFoundError, ListResponse, PaginationInfo +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityNotFoundError + +_ResultEntityT = TypeVar("_ResultEntityT", AgentEvalResultEntity, EvaluateResultEntity) class _FakeEntityClient: """In-memory store keyed by (entity_type, workspace, name).""" def __init__(self) -> None: - self.entities: dict[tuple[str, str, str], EntityBase] = {} + self.entities: dict[tuple[str, str, str], AgentEvalResultEntity | EvaluateResultEntity] = {} - def seed(self, entity: EntityBase) -> EntityBase: + def seed(self, entity: _ResultEntityT) -> _ResultEntityT: now = datetime.now(timezone.utc) entity._id = f"{entity.__entity_type__}-{entity.name}" entity._created_at = now @@ -33,23 +37,43 @@ def seed(self, entity: EntityBase) -> EntityBase: self.entities[(entity.__entity_type__, entity.workspace, entity.name)] = entity return entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[_ResultEntityT], *, workspace: str, name: str) -> _ResultEntityT: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") - return self.entities[key] + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") + entity = self.entities[key] + assert isinstance(entity, entity_type) + return entity - async def delete(self, entity_cls, name, *, workspace): - # Mirror the real EntityClient: raise EntityNotFoundError when absent. - key = (entity_cls.__entity_type__, workspace, name) + async def delete( + self, + entity_type: type[_ResultEntityT], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + # Mirror the real EntityClient: raise NemoEntityNotFoundError when absent. + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") del self.entities[key] - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): - items = [ - e for (etype, ws, _), e in self.entities.items() if etype == entity_cls.__entity_type__ and ws == workspace - ] + async def list( + self, + entity_type: type[_ResultEntityT], + *, + workspace: str, + filter_operation: object | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[_ResultEntityT]: + items: list[_ResultEntityT] = [] + for (etype, ws, _), entity in self.entities.items(): + if etype == entity_type.__entity_type__ and ws == workspace: + assert isinstance(entity, entity_type) + items.append(entity) return ListResponse( data=items, pagination=PaginationInfo( diff --git a/plugins/nemo-evaluator/tests/api/service/test_task_service.py b/plugins/nemo-evaluator/tests/api/service/test_task_service.py index e484377f30..6b5b333f92 100644 --- a/plugins/nemo-evaluator/tests/api/service/test_task_service.py +++ b/plugins/nemo-evaluator/tests/api/service/test_task_service.py @@ -6,18 +6,14 @@ from datetime import datetime, timezone import pytest -from nemo_evaluator.api.schemas import MetricInline, MetricRef, Task, TaskInput +from nemo_evaluator.api.schemas import MetadataItem, MetricInline, MetricRef, Task, TaskInput, TaskInputs from nemo_evaluator.api.service.task_service import MetricRefNotFoundError, TaskService +from nemo_evaluator.entities import TaskEntity from nemo_evaluator.shared.metric_bundles.bundles import bundle_metric from nemo_evaluator.shared.metric_bundles.cloudpickle import CloudpickleMetricBundlePackager from nemo_evaluator_sdk.metrics.exact_match import ExactMatchMetric -from nemo_platform_plugin.entities import ( - EntityBase, - EntityConflictError, - EntityNotFoundError, - ListResponse, - PaginationInfo, -) +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError class _FakeMetricService: @@ -45,12 +41,12 @@ def _inline_metric() -> MetricInline: class _FakeEntityClient: def __init__(self) -> None: - self.entities: dict[tuple[str, str, str], EntityBase] = {} + self.entities: dict[tuple[str, str, str], TaskEntity] = {} - async def create(self, entity): + async def create(self, entity: TaskEntity) -> TaskEntity: key = (entity.__entity_type__, entity.workspace, entity.name) if key in self.entities: - raise EntityConflictError(f"{key} exists") + raise NemoEntityConflictError(f"{key} exists") now = datetime.now(timezone.utc) entity._id = f"{entity.__entity_type__}-{entity.name}" entity._created_at = now @@ -58,21 +54,37 @@ async def create(self, entity): self.entities[key] = entity return entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[TaskEntity], *, workspace: str, name: str) -> TaskEntity: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") return self.entities[key] - async def delete(self, entity_cls, name, *, workspace): - key = (entity_cls.__entity_type__, workspace, name) + async def delete( + self, + entity_type: type[TaskEntity], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") del self.entities[key] - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): + async def list( + self, + entity_type: type[TaskEntity], + *, + workspace: str, + filter_operation: object | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[TaskEntity]: items = [ - e for (etype, ws, _), e in self.entities.items() if etype == entity_cls.__entity_type__ and ws == workspace + e for (etype, ws, _), e in self.entities.items() if etype == entity_type.__entity_type__ and ws == workspace ] return ListResponse( data=items, @@ -89,9 +101,9 @@ async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, def _task_input() -> TaskInput: return TaskInput( intent="Answer the question.", - inputs={"instruction": "What is 2+2?"}, + inputs=TaskInputs(instruction="What is 2+2?"), metrics=[MetricRef("default/stored-metric")], - metadata=[{"key": "suite", "value": "smoke"}], + metadata=[MetadataItem(key="suite", value="smoke")], ) @@ -125,7 +137,7 @@ async def test_create_normalizes_inline_metrics_to_refs( inline = _inline_metric() task_input = TaskInput( intent="Answer the question.", - inputs={"instruction": "What is 2+2?"}, + inputs=TaskInputs(instruction="What is 2+2?"), metrics=[MetricRef("default/stored-metric"), inline], ) @@ -140,14 +152,14 @@ async def test_create_normalizes_inline_metrics_to_refs( async def test_create_rejects_missing_metric_ref(service: TaskService) -> None: - task_input = TaskInput(intent="x", inputs={"instruction": "?"}, metrics=[MetricRef("default/nope")]) + task_input = TaskInput(intent="x", inputs=TaskInputs(instruction="?"), metrics=[MetricRef("default/nope")]) with pytest.raises(MetricRefNotFoundError, match="not found"): await service.create_task("task-1", task_input, workspace="default") async def test_create_canonicalizes_bare_metric_ref(service: TaskService) -> None: # A bare "stored-metric" ref resolves against the task workspace and is persisted as "default/stored-metric". - task_input = TaskInput(intent="x", inputs={"instruction": "?"}, metrics=[MetricRef("stored-metric")]) + task_input = TaskInput(intent="x", inputs=TaskInputs(instruction="?"), metrics=[MetricRef("stored-metric")]) created = await service.create_task("task-1", task_input, workspace="default") assert created.metrics[0].root == "default/stored-metric" diff --git a/plugins/nemo-evaluator/tests/api/service/test_taskset_service.py b/plugins/nemo-evaluator/tests/api/service/test_taskset_service.py index 8ce7d511f3..d1bec5757b 100644 --- a/plugins/nemo-evaluator/tests/api/service/test_taskset_service.py +++ b/plugins/nemo-evaluator/tests/api/service/test_taskset_service.py @@ -6,20 +6,16 @@ from datetime import datetime, timezone import pytest -from nemo_evaluator.api.schemas import TaskRef, Taskset, TasksetInput +from nemo_evaluator.api.schemas import MetadataItem, TaskRef, Taskset, TasksetInput from nemo_evaluator.api.service.taskset_service import ( DuplicateTaskRefError, TaskRefNotFoundError, TasksetExistsError, TasksetService, ) -from nemo_platform_plugin.entities import ( - EntityBase, - EntityConflictError, - EntityNotFoundError, - ListResponse, - PaginationInfo, -) +from nemo_evaluator.entities import TasksetEntity +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError class _FakeTaskService: @@ -34,12 +30,12 @@ async def get_task(self, workspace: str, name: str) -> object | None: class _FakeEntityClient: def __init__(self) -> None: - self.entities: dict[tuple[str, str, str], EntityBase] = {} + self.entities: dict[tuple[str, str, str], TasksetEntity] = {} - async def create(self, entity): + async def create(self, entity: TasksetEntity) -> TasksetEntity: key = (entity.__entity_type__, entity.workspace, entity.name) if key in self.entities: - raise EntityConflictError(f"{key} exists") + raise NemoEntityConflictError(f"{key} exists") now = datetime.now(timezone.utc) entity._id = f"{entity.__entity_type__}-{entity.name}" entity._created_at = now @@ -47,21 +43,37 @@ async def create(self, entity): self.entities[key] = entity return entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[TasksetEntity], *, workspace: str, name: str) -> TasksetEntity: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") return self.entities[key] - async def delete(self, entity_cls, name, *, workspace): - key = (entity_cls.__entity_type__, workspace, name) + async def delete( + self, + entity_type: type[TasksetEntity], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") del self.entities[key] - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): + async def list( + self, + entity_type: type[TasksetEntity], + *, + workspace: str, + filter_operation: object | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[TasksetEntity]: items = [ - e for (etype, ws, _), e in self.entities.items() if etype == entity_cls.__entity_type__ and ws == workspace + e for (etype, ws, _), e in self.entities.items() if etype == entity_type.__entity_type__ and ws == workspace ] return ListResponse( data=items, @@ -79,7 +91,7 @@ def _taskset_input() -> TasksetInput: return TasksetInput( description="A smoke-test grouping.", tasks=[TaskRef("task-a"), TaskRef("default/task-b")], - metadata=[{"key": "suite", "value": "smoke"}], + metadata=[MetadataItem(key="suite", value="smoke")], ) diff --git a/plugins/nemo-evaluator/tests/api/v2/test_metrics_routes.py b/plugins/nemo-evaluator/tests/api/v2/test_metrics_routes.py index 0c4fe7e320..e9ff7b1591 100644 --- a/plugins/nemo-evaluator/tests/api/v2/test_metrics_routes.py +++ b/plugins/nemo-evaluator/tests/api/v2/test_metrics_routes.py @@ -10,8 +10,10 @@ from __future__ import annotations +from collections.abc import Iterator +from dataclasses import dataclass from datetime import datetime, timezone -from unittest.mock import AsyncMock, patch +from unittest.mock import patch import pytest from fastapi import FastAPI @@ -25,12 +27,11 @@ from nemo_evaluator.shared.metric_bundles.cloudpickle import CloudpickleMetricBundlePackager from nemo_evaluator.shared.metric_bundles.inline import InlineMetricBundlePackager from nemo_evaluator_sdk.metrics.exact_match import ExactMatchMetric -from nemo_platform_plugin.entities import ( - EntityConflictError, - EntityNotFoundError, - ListResponse, - PaginationInfo, -) +from nemo_platform import AsyncNeMoPlatform +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError +from nemo_platform_plugin.files.types import CreateFilesetRequest +from nemo_platform_plugin.filter_ops import FilterOperation # ---- in-memory fakes ------------------------------------------------------- @@ -39,53 +40,94 @@ class _FakeAsyncFilesClient: def __init__(self) -> None: self._store: dict[tuple[str, str], dict[str, bytes]] = {} - async def create_fileset(self, *, body, workspace=None, exist_ok=False): - self._store.setdefault((workspace, body.name), {}) - return AsyncMock(data=lambda: object()) + async def create_fileset( + self, *, body: CreateFilesetRequest, workspace: str | None = None, exist_ok: bool = False + ) -> _FakeOperationResponse: + self._store.setdefault((workspace or "default", body.name), {}) + return _FakeOperationResponse() - async def delete_fileset(self, *, name, workspace=None): - self._store.pop((workspace, name), None) - return AsyncMock(data=lambda: object()) + async def delete_fileset(self, *, name: str, workspace: str | None = None) -> _FakeOperationResponse: + self._store.pop((workspace or "default", name), None) + return _FakeOperationResponse() - async def upload_file(self, *, path, content, workspace, name): + async def upload_file(self, *, path: str, content: bytes, workspace: str, name: str) -> _FakeOperationResponse: self._store.setdefault((workspace, name), {})[path] = bytes(content) - return AsyncMock(data=lambda: object()) + return _FakeOperationResponse() - async def download_file(self, *, path, workspace, name): - class _Resp: - async def read(self): - return self._data + async def download_file(self, *, path: str, workspace: str, name: str) -> _FakeResponse: + return _FakeResponse(self._store[(workspace, name)][path]) - resp = _Resp() - resp._data = self._store[(workspace, name)][path] - return resp + +class _FakeResponse: + def __init__(self, data: bytes) -> None: + self._data = data + + async def read(self) -> bytes: + return self._data + + +class _FakeOperationResponse: + def data(self) -> object: + return object() class _FakeEntityClient: def __init__(self) -> None: self.entities: dict[tuple[str, str], MetricBundleEntity] = {} + self.entity_versions: dict[tuple[str, str], int] = {} + self.bump_version_on_next_delete = False + self.delete_expected_db_versions: list[int | None] = [] - async def get(self, entity_cls, *, workspace, name): + async def get(self, entity_type: type[MetricBundleEntity], *, workspace: str, name: str) -> MetricBundleEntity: key = (workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") return self.entities[key] - async def create(self, entity): + async def create(self, entity: MetricBundleEntity) -> MetricBundleEntity: key = (entity.workspace, entity.name) if key in self.entities: - raise EntityConflictError(f"{key} exists") + raise NemoEntityConflictError(f"{key} exists") now = datetime.now(timezone.utc) entity._id = f"metric_bundle-{entity.name}" entity._created_at = now entity._updated_at = now + entity._db_version = 1 self.entities[key] = entity + self.entity_versions[key] = entity.db_version return entity - async def delete(self, entity_cls, name, *, workspace): - self.entities.pop((workspace, name), None) - - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): + async def delete( + self, + entity_type: type[MetricBundleEntity], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + key = (workspace, name) + if key not in self.entities: + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") + if self.bump_version_on_next_delete: + self.entity_versions[key] += 1 + self.entities[key]._db_version = self.entity_versions[key] + self.bump_version_on_next_delete = False + self.delete_expected_db_versions.append(expected_db_version) + if expected_db_version is not None and self.entity_versions[key] != expected_db_version: + raise NemoEntityConflictError(f"{workspace}/{name} changed") + del self.entities[key] + del self.entity_versions[key] + + async def list( + self, + entity_type: type[MetricBundleEntity], + *, + workspace: str, + filter_operation: FilterOperation | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[MetricBundleEntity]: items = [e for (ws, _), e in self.entities.items() if ws == workspace] return ListResponse( data=items, @@ -99,15 +141,35 @@ async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, ) +class _FakePlatform(AsyncNeMoPlatform): + pass + + +def _fake_platform() -> _FakePlatform: + return _FakePlatform.__new__(_FakePlatform) + + +@dataclass(frozen=True) +class _MetricsRouteHarness: + client: TestClient + entity_client: _FakeEntityClient + + @pytest.fixture -def client() -> TestClient: +def metrics_route_harness() -> Iterator[_MetricsRouteHarness]: app = FastAPI() app.include_router(metrics_routes.router, prefix="/v2/workspaces/{workspace}") fake_files = _FakeAsyncFilesClient() - service = MetricService(_FakeEntityClient(), object()) + entity_client = _FakeEntityClient() + service = MetricService(entity_client, _fake_platform()) app.dependency_overrides[get_metric_service] = lambda: service with patch("nemo_evaluator.metric_storage.client_from_platform", return_value=fake_files): - yield TestClient(app) + yield _MetricsRouteHarness(TestClient(app), entity_client) + + +@pytest.fixture +def client(metrics_route_harness: _MetricsRouteHarness) -> TestClient: + return metrics_route_harness.client def _create_body() -> dict: @@ -168,6 +230,20 @@ def test_create_then_delete(client: TestClient) -> None: assert client.get(f"{_BASE}/exact").status_code == 404 +def test_delete_stale_version_returns_409_and_keeps_metric(metrics_route_harness: _MetricsRouteHarness) -> None: + client = metrics_route_harness.client + entity_client = metrics_route_harness.entity_client + assert client.post(f"{_BASE}/exact", json=_create_body()).status_code == 201 + + entity_client.bump_version_on_next_delete = True + deleted = client.delete(f"{_BASE}/exact") + + assert deleted.status_code == 409 + assert entity_client.delete_expected_db_versions == [1] + assert ("default", "exact") in entity_client.entities + assert client.get(f"{_BASE}/exact").status_code == 200 + + def test_delete_missing_returns_404(client: TestClient) -> None: assert client.delete(f"{_BASE}/nope").status_code == 404 diff --git a/plugins/nemo-evaluator/tests/api/v2/test_results_routes.py b/plugins/nemo-evaluator/tests/api/v2/test_results_routes.py index 02da07a415..24cc48f1e8 100644 --- a/plugins/nemo-evaluator/tests/api/v2/test_results_routes.py +++ b/plugins/nemo-evaluator/tests/api/v2/test_results_routes.py @@ -11,6 +11,7 @@ from __future__ import annotations from datetime import datetime, timezone +from typing import TypeVar import pytest from fastapi import FastAPI @@ -20,14 +21,17 @@ from nemo_evaluator.api.v2 import results as results_routes from nemo_evaluator.entities import AgentEvalResultEntity, EvaluateResultEntity from nemo_evaluator_sdk.values.results import AggregatedMetricResult -from nemo_platform_plugin.entities import EntityBase, EntityNotFoundError, ListResponse, PaginationInfo +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError + +_ResultEntityT = TypeVar("_ResultEntityT", AgentEvalResultEntity, EvaluateResultEntity) class _FakeEntityClient: def __init__(self) -> None: - self.entities: dict[tuple[str, str, str], EntityBase] = {} + self.entities: dict[tuple[str, str, str], AgentEvalResultEntity | EvaluateResultEntity] = {} - def seed(self, entity: EntityBase) -> EntityBase: + def seed(self, entity: _ResultEntityT) -> _ResultEntityT: now = datetime.now(timezone.utc) entity._id = f"{entity.__entity_type__}-{entity.name}" entity._created_at = now @@ -35,23 +39,43 @@ def seed(self, entity: EntityBase) -> EntityBase: self.entities[(entity.__entity_type__, entity.workspace, entity.name)] = entity return entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[_ResultEntityT], *, workspace: str, name: str) -> _ResultEntityT: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") - return self.entities[key] + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") + entity = self.entities[key] + assert isinstance(entity, entity_type) + return entity - async def delete(self, entity_cls, name, *, workspace): - # Mirror the real EntityClient: raise EntityNotFoundError when absent. - key = (entity_cls.__entity_type__, workspace, name) + async def delete( + self, + entity_type: type[_ResultEntityT], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + # Mirror the real EntityClient: raise NemoEntityNotFoundError when absent. + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") del self.entities[key] - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): - items = [ - e for (etype, ws, _), e in self.entities.items() if etype == entity_cls.__entity_type__ and ws == workspace - ] + async def list( + self, + entity_type: type[_ResultEntityT], + *, + workspace: str, + filter_operation: object | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[_ResultEntityT]: + items: list[_ResultEntityT] = [] + for (etype, ws, _), entity in self.entities.items(): + if etype == entity_type.__entity_type__ and ws == workspace: + assert isinstance(entity, entity_type) + items.append(entity) return ListResponse( data=items, pagination=PaginationInfo( @@ -150,6 +174,32 @@ def test_delete_missing_returns_404(client: TestClient) -> None: assert client.delete(f"{_EVAL}/nope").status_code == 404 +def test_delete_agent_eval_result_conflict_returns_409() -> None: + class _Service: + async def delete_agent_eval_result(self, workspace: str, name: str) -> bool: + raise NemoEntityConflictError("changed") + + app = FastAPI() + app.include_router(results_routes.agent_eval_results_router, prefix="/v2/workspaces/{workspace}") + app.dependency_overrides[get_result_service] = lambda: _Service() + client = TestClient(app) + + assert client.delete(f"{_AGENT}/job-1").status_code == 409 + + +def test_delete_eval_result_conflict_returns_409() -> None: + class _Service: + async def delete_eval_result(self, workspace: str, name: str) -> bool: + raise NemoEntityConflictError("changed") + + app = FastAPI() + app.include_router(results_routes.evaluate_results_router, prefix="/v2/workspaces/{workspace}") + app.dependency_overrides[get_result_service] = lambda: _Service() + client = TestClient(app) + + assert client.delete(f"{_EVAL}/job-1").status_code == 409 + + def test_filter_translates_custom_fields_to_data_namespace() -> None: # Custom (non-base) trait fields must be rewritten to data.* for the entity store; base columns # (workspace, created_at) pass through. The plain Filter does no translation and the store 500s. diff --git a/plugins/nemo-evaluator/tests/api/v2/test_tasks_routes.py b/plugins/nemo-evaluator/tests/api/v2/test_tasks_routes.py index 4d77b283f6..e7b09ada9c 100644 --- a/plugins/nemo-evaluator/tests/api/v2/test_tasks_routes.py +++ b/plugins/nemo-evaluator/tests/api/v2/test_tasks_routes.py @@ -15,26 +15,23 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from nemo_evaluator.api.dependencies import get_task_service -from nemo_evaluator.api.schemas import MetricRef, TaskInput +from nemo_evaluator.api.schemas import MetricInline, MetricRef, TaskInput, TaskInputs from nemo_evaluator.api.service.task_service import TaskService from nemo_evaluator.api.v2 import tasks as tasks_routes -from nemo_platform_plugin.entities import ( - EntityBase, - EntityConflictError, - EntityNotFoundError, - ListResponse, - PaginationInfo, -) +from nemo_evaluator.entities import TaskEntity +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError +from nemo_platform_plugin.filter_ops import FilterOperation class _FakeEntityClient: def __init__(self) -> None: - self.entities: dict[tuple[str, str, str], EntityBase] = {} + self.entities: dict[tuple[str, str, str], TaskEntity] = {} - async def create(self, entity): + async def create(self, entity: TaskEntity) -> TaskEntity: key = (entity.__entity_type__, entity.workspace, entity.name) if key in self.entities: - raise EntityConflictError(f"{key} exists") + raise NemoEntityConflictError(f"{key} exists") now = datetime.now(timezone.utc) entity._id = f"{entity.__entity_type__}-{entity.name}" entity._created_at = now @@ -42,21 +39,37 @@ async def create(self, entity): self.entities[key] = entity return entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[TaskEntity], *, workspace: str, name: str) -> TaskEntity: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") return self.entities[key] - async def delete(self, entity_cls, name, *, workspace): - key = (entity_cls.__entity_type__, workspace, name) + async def delete( + self, + entity_type: type[TaskEntity], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") del self.entities[key] - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): + async def list( + self, + entity_type: type[TaskEntity], + *, + workspace: str, + filter_operation: FilterOperation | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[TaskEntity]: items = [ - e for (etype, ws, _), e in self.entities.items() if etype == entity_cls.__entity_type__ and ws == workspace + e for (etype, ws, _), e in self.entities.items() if etype == entity_type.__entity_type__ and ws == workspace ] return ListResponse( data=items, @@ -70,7 +83,7 @@ class _FakeMetricService: """Normalizes inline metrics to derived refs and resolves the ``default/stored-metric`` ref the route bodies submit, so task-create metric-ref validation passes for the happy-path tests.""" - async def store_derived_metric(self, metric, *, workspace: str) -> MetricRef: + async def store_derived_metric(self, metric: MetricInline, *, workspace: str) -> MetricRef: return MetricRef(f"{workspace}/derived.{metric.payload.digest}") async def get_metric(self, workspace: str, name: str) -> object | None: @@ -89,7 +102,7 @@ def client() -> TestClient: def _body() -> dict: return TaskInput( intent="Answer the question.", - inputs={"instruction": "What is 2+2?"}, + inputs=TaskInputs(instruction="What is 2+2?"), metrics=[MetricRef("default/stored-metric")], ).model_dump(mode="json") @@ -162,3 +175,16 @@ def test_delete_then_get_404(client: TestClient) -> None: def test_delete_missing_returns_404(client: TestClient) -> None: assert client.delete(f"{_BASE}/nope").status_code == 404 + + +def test_delete_conflict_returns_409() -> None: + class _Service: + async def delete_task(self, workspace: str, name: str) -> bool: + raise NemoEntityConflictError("changed") + + app = FastAPI() + app.include_router(tasks_routes.router, prefix="/v2/workspaces/{workspace}") + app.dependency_overrides[get_task_service] = lambda: _Service() + client = TestClient(app) + + assert client.delete(f"{_BASE}/task-1").status_code == 409 diff --git a/plugins/nemo-evaluator/tests/api/v2/test_tasksets_routes.py b/plugins/nemo-evaluator/tests/api/v2/test_tasksets_routes.py index c7e94b2460..277276c686 100644 --- a/plugins/nemo-evaluator/tests/api/v2/test_tasksets_routes.py +++ b/plugins/nemo-evaluator/tests/api/v2/test_tasksets_routes.py @@ -18,23 +18,20 @@ from nemo_evaluator.api.schemas import TaskRef, TasksetInput from nemo_evaluator.api.service.taskset_service import TasksetService from nemo_evaluator.api.v2 import tasksets as tasksets_routes -from nemo_platform_plugin.entities import ( - EntityBase, - EntityConflictError, - EntityNotFoundError, - ListResponse, - PaginationInfo, -) +from nemo_evaluator.entities import TasksetEntity +from nemo_platform_plugin.entities import ListResponse, PaginationInfo +from nemo_platform_plugin.entity_client import NemoEntityConflictError, NemoEntityNotFoundError +from nemo_platform_plugin.filter_ops import FilterOperation class _FakeEntityClient: def __init__(self) -> None: - self.entities: dict[tuple[str, str, str], EntityBase] = {} + self.entities: dict[tuple[str, str, str], TasksetEntity] = {} - async def create(self, entity): + async def create(self, entity: TasksetEntity) -> TasksetEntity: key = (entity.__entity_type__, entity.workspace, entity.name) if key in self.entities: - raise EntityConflictError(f"{key} exists") + raise NemoEntityConflictError(f"{key} exists") now = datetime.now(timezone.utc) entity._id = f"{entity.__entity_type__}-{entity.name}" entity._created_at = now @@ -42,21 +39,37 @@ async def create(self, entity): self.entities[key] = entity return entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[TasksetEntity], *, workspace: str, name: str) -> TasksetEntity: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") return self.entities[key] - async def delete(self, entity_cls, name, *, workspace): - key = (entity_cls.__entity_type__, workspace, name) + async def delete( + self, + entity_type: type[TasksetEntity], + name: str, + *, + workspace: str, + expected_db_version: int | None = None, + ) -> None: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") del self.entities[key] - async def list(self, entity_cls, *, workspace, filter_operation=None, sort=None, page=1, page_size=100): + async def list( + self, + entity_type: type[TasksetEntity], + *, + workspace: str, + filter_operation: FilterOperation | None = None, + sort: str | None = None, + page: int = 1, + page_size: int = 100, + ) -> ListResponse[TasksetEntity]: items = [ - e for (etype, ws, _), e in self.entities.items() if etype == entity_cls.__entity_type__ and ws == workspace + e for (etype, ws, _), e in self.entities.items() if etype == entity_type.__entity_type__ and ws == workspace ] return ListResponse( data=items, @@ -170,3 +183,16 @@ def test_delete_then_get_404(client: TestClient) -> None: def test_delete_missing_returns_404(client: TestClient) -> None: assert client.delete(f"{_BASE}/nope").status_code == 404 + + +def test_delete_conflict_returns_409() -> None: + class _Service: + async def delete_taskset(self, workspace: str, name: str) -> bool: + raise NemoEntityConflictError("changed") + + app = FastAPI() + app.include_router(tasksets_routes.router, prefix="/v2/workspaces/{workspace}") + app.dependency_overrides[get_taskset_service] = lambda: _Service() + client = TestClient(app) + + assert client.delete(f"{_BASE}/ts-1").status_code == 409 diff --git a/plugins/nemo-evaluator/tests/test_metric_refs.py b/plugins/nemo-evaluator/tests/test_metric_refs.py index d5e3616d76..2d41e553c8 100644 --- a/plugins/nemo-evaluator/tests/test_metric_refs.py +++ b/plugins/nemo-evaluator/tests/test_metric_refs.py @@ -15,10 +15,12 @@ resolve_metric_specs, ) from nemo_evaluator.metric_storage import store_bundle -from nemo_evaluator.shared.metric_bundles.bundles import bundle_metric +from nemo_evaluator.shared.metric_bundles.bundles import MetricBundle, bundle_metric from nemo_evaluator.shared.metric_bundles.cloudpickle import CloudpickleMetricBundlePackager from nemo_evaluator_sdk.metrics.exact_match import ExactMatchMetric -from nemo_platform_plugin.entities import EntityNotFoundError +from nemo_platform import AsyncNeMoPlatform +from nemo_platform_plugin.entity_client import NemoEntityNotFoundError +from nemo_platform_plugin.files.types import CreateFilesetRequest from pydantic import ValidationError # ---- in-memory fakes (mirror the storage round-trip) ----------------------- @@ -36,19 +38,21 @@ class _FakeAsyncFilesClient: def __init__(self) -> None: self._store: dict[tuple[str, str], dict[str, bytes]] = {} - async def create_fileset(self, *, body, workspace=None, exist_ok=False): - self._store.setdefault((workspace, body.name), {}) + async def create_fileset( + self, *, body: CreateFilesetRequest, workspace: str | None = None, exist_ok: bool = False + ) -> AsyncMock: + self._store.setdefault((workspace or "default", body.name), {}) return AsyncMock(data=lambda: object()) - async def delete_fileset(self, *, name, workspace=None): - self._store.pop((workspace, name), None) + async def delete_fileset(self, *, name: str, workspace: str | None = None) -> AsyncMock: + self._store.pop((workspace or "default", name), None) return AsyncMock(data=lambda: object()) - async def upload_file(self, *, path, content, workspace, name): + async def upload_file(self, *, path: str, content: bytes, workspace: str, name: str) -> AsyncMock: self._store.setdefault((workspace, name), {})[path] = bytes(content) return AsyncMock(data=lambda: object()) - async def download_file(self, *, path, workspace, name): + async def download_file(self, *, path: str, workspace: str, name: str) -> _FakeResponse: return _FakeResponse(self._store[(workspace, name)][path]) @@ -56,11 +60,11 @@ class _FakeEntityClient: def __init__(self) -> None: self.entities: dict[tuple[str, str], MetricBundleEntity] = {} - async def get(self, entity_cls, *, workspace, name): + async def get(self, entity_type: type[MetricBundleEntity], *, workspace: str, name: str) -> MetricBundleEntity: try: return self.entities[(workspace, name)] except KeyError: - raise EntityNotFoundError(f"{workspace}/{name} not found") + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") def _bundle(): @@ -73,10 +77,16 @@ def _metric_inline() -> MetricInline: return MetricInline.model_validate_json(_bundle().model_dump_json()) -async def _stored(fake_client: _FakeAsyncFilesClient, entity_client: _FakeEntityClient, workspace: str, name: str): +def _fake_platform() -> AsyncNeMoPlatform: + return AsyncMock(spec=AsyncNeMoPlatform) + + +async def _stored( + fake_client: _FakeAsyncFilesClient, entity_client: _FakeEntityClient, workspace: str, name: str +) -> MetricBundle: bundle = _bundle() with patch("nemo_evaluator.metric_storage.client_from_platform", return_value=fake_client): - ref = await store_bundle(object(), workspace, name, bundle) + ref = await store_bundle(_fake_platform(), workspace, name, bundle) entity_client.entities[(workspace, name)] = MetricBundleEntity( name=name, workspace=workspace, @@ -127,7 +137,7 @@ async def test_resolve_loads_referenced_bundle() -> None: [MetricRef(root="default/exact")], workspace="default", entity_client=entity_client, - async_sdk=object(), + async_sdk=_fake_platform(), ) assert len(result) == 1 @@ -145,7 +155,7 @@ async def test_resolve_mixes_refs_and_inline_preserving_order() -> None: [MetricRef(root="exact"), inline], workspace="default", entity_client=entity_client, - async_sdk=object(), + async_sdk=_fake_platform(), ) assert len(result) == 2 @@ -168,7 +178,7 @@ async def test_resolve_missing_metric_raises_clear_error() -> None: [MetricRef(root="default/no-such-metric")], workspace="default", entity_client=_FakeEntityClient(), - async_sdk=object(), + async_sdk=_fake_platform(), ) @@ -178,7 +188,7 @@ async def test_resolve_ref_without_entity_client_raises() -> None: [MetricRef(root="default/exact")], workspace="default", entity_client=None, - async_sdk=object(), + async_sdk=_fake_platform(), ) @@ -186,9 +196,11 @@ async def test_resolve_ref_without_entity_client_raises() -> None: def test_input_spec_accepts_ref_and_inline() -> None: - spec = EvaluateInputSpec( - metrics=["default/stored-metric", _metric_inline()], - dataset=[{"expected": "a", "output": "a"}], + spec = EvaluateInputSpec.model_validate( + { + "metrics": ["default/stored-metric", _metric_inline()], + "dataset": [{"expected": "a", "output": "a"}], + } ) assert isinstance(spec.metrics[0], MetricRef) assert spec.metrics[0].root == "default/stored-metric" @@ -196,7 +208,9 @@ def test_input_spec_accepts_ref_and_inline() -> None: def test_canonical_spec_rejects_unresolved_ref() -> None: with pytest.raises(ValidationError): - EvaluateSpec( - metrics=["default/stored-metric"], - dataset=[{"expected": "a", "output": "a"}], + EvaluateSpec.model_validate( + { + "metrics": ["default/stored-metric"], + "dataset": [{"expected": "a", "output": "a"}], + } ) diff --git a/plugins/nemo-evaluator/tests/test_task_refs.py b/plugins/nemo-evaluator/tests/test_task_refs.py index a0b635c3b5..34e8353e10 100644 --- a/plugins/nemo-evaluator/tests/test_task_refs.py +++ b/plugins/nemo-evaluator/tests/test_task_refs.py @@ -5,12 +5,17 @@ from __future__ import annotations +from typing import TypeVar + import pytest from nemo_evaluator.api.schemas import MetadataItem, MetricRef, TaskInputs, TaskRef, TasksetRef from nemo_evaluator.entities import TaskEntity, TasksetEntity from nemo_evaluator.jobs.agent_spec import AgentEvalTaskInput from nemo_evaluator.task_refs import resolve_agent_eval_tasks, resolve_taskset_ref -from nemo_platform_plugin.entities import EntityBase, EntityNotFoundError +from nemo_platform_plugin.entities import EntityBase +from nemo_platform_plugin.entity_client import NemoEntityNotFoundError + +_EntityT = TypeVar("_EntityT", bound=EntityBase) class _FakeEntityClient: @@ -22,11 +27,13 @@ def __init__(self) -> None: def add(self, entity: EntityBase) -> None: self.entities[(entity.__entity_type__, entity.workspace, entity.name)] = entity - async def get(self, entity_cls, *, workspace, name): - key = (entity_cls.__entity_type__, workspace, name) + async def get(self, entity_type: type[_EntityT], *, workspace: str, name: str) -> _EntityT: + key = (entity_type.__entity_type__, workspace, name) if key not in self.entities: - raise EntityNotFoundError(f"{workspace}/{name} not found") - return self.entities[key] + raise NemoEntityNotFoundError(f"{workspace}/{name} not found") + entity = self.entities[key] + assert isinstance(entity, entity_type) + return entity def _task(name: str, *, workspace: str = "default", metric: str = "default/m") -> TaskEntity: diff --git a/plugins/nemo-guardrails/tests/integration/utils.py b/plugins/nemo-guardrails/tests/integration/utils.py index 44bf501949..d7f8be80e0 100644 --- a/plugins/nemo-guardrails/tests/integration/utils.py +++ b/plugins/nemo-guardrails/tests/integration/utils.py @@ -102,6 +102,7 @@ def make_guardrail_config( "parent": workspace, "workspace": workspace, "name": name, + "db_version": 1, "created_at": "2026-01-01T00:00:00Z", "updated_at": "2026-01-01T00:00:00Z", "data": data, diff --git a/plugins/nemo-guardrails/tests/unit/test_middleware.py b/plugins/nemo-guardrails/tests/unit/test_middleware.py index b486f87b39..2f0d5182e6 100644 --- a/plugins/nemo-guardrails/tests/unit/test_middleware.py +++ b/plugins/nemo-guardrails/tests/unit/test_middleware.py @@ -156,6 +156,7 @@ def _make_entity( "parent": "", "workspace": workspace, "name": name, + "db_version": 1, "created_at": "2026-01-01T00:00:00Z", "updated_at": updated_at, "data": rails.model_dump(exclude_none=True), @@ -653,7 +654,7 @@ async def test_input_masking_writes_back_last_user_message(self, middleware: Gua with patch.object(middleware, "_run_rails", new=AsyncMock(return_value=generation_response)): result = await _process_request(middleware, request_body, {}, _entity_source()) - assert isinstance(result, dict) + assert not isinstance(result, ImmediateResponse) assert result["messages"] == [ {"role": "user", "content": "earlier turn"}, {"role": "assistant", "content": "ok"}, @@ -681,7 +682,7 @@ async def test_input_masking_skips_non_string_user_content(self, middleware: Gua with patch.object(middleware, "_run_rails", new=AsyncMock(return_value=generation_response)): result = await _process_request(middleware, request_body, {}, _entity_source()) - assert isinstance(result, dict) + assert not isinstance(result, ImmediateResponse) assert result["messages"][0]["content"] == multimodal_content async def test_user_log_options_forwarded_to_run_rails(self, middleware: GuardrailsMiddleware) -> None: diff --git a/plugins/nemo-insights/src/nemo_insights_plugin/service.py b/plugins/nemo-insights/src/nemo_insights_plugin/service.py index ae70b231f9..6c70c77791 100644 --- a/plugins/nemo-insights/src/nemo_insights_plugin/service.py +++ b/plugins/nemo-insights/src/nemo_insights_plugin/service.py @@ -303,6 +303,8 @@ async def delete_insight( status_code=404, detail=f"Insight '{insight_id}' not found in workspace '{workspace}'.", ) from exc + except NemoEntityConflictError as exc: + raise HTTPException(status_code=409, detail="Concurrent modification - please retry.") from exc except Exception as exc: logger.exception("Failed to delete insight") raise HTTPException(status_code=500, detail="Failed to delete insight.") from exc diff --git a/plugins/nemo-insights/tests/test_cli_profile.py b/plugins/nemo-insights/tests/test_cli_profile.py index fea8fe460e..418fcece24 100644 --- a/plugins/nemo-insights/tests/test_cli_profile.py +++ b/plugins/nemo-insights/tests/test_cli_profile.py @@ -2,6 +2,7 @@ # SPDX-License-Identifier: Apache-2.0 import os +from collections.abc import Iterator from pathlib import Path import httpx @@ -15,6 +16,7 @@ from typer.testing import CliRunner runner = CliRunner() +_PROFILE_ENV_KEYS = ("NMP_BASE_URL", "INFERENCE_API_KEY") class AnalystRecorder: @@ -32,7 +34,21 @@ def app() -> typer.Typer: @pytest.fixture(autouse=True) -def quiet_preflight(monkeypatch: pytest.MonkeyPatch) -> None: +def restore_profile_env() -> Iterator[None]: + original = {key: os.environ.get(key) for key in _PROFILE_ENV_KEYS} + missing = {key for key in _PROFILE_ENV_KEYS if key not in os.environ} + yield + for key in _PROFILE_ENV_KEYS: + if key in missing: + os.environ.pop(key, None) + else: + value = original[key] + if value is not None: + os.environ[key] = value + + +@pytest.fixture(autouse=True) +def quiet_preflight(monkeypatch: pytest.MonkeyPatch, restore_profile_env: None) -> None: async def queryable(base_url: str, workspace: str, agent: str) -> bool: return True diff --git a/sdk/python/nemo-platform/.nmpcontext/openapi.yaml b/sdk/python/nemo-platform/.nmpcontext/openapi.yaml index 5c62b70686..f1b5acfbc6 100644 --- a/sdk/python/nemo-platform/.nmpcontext/openapi.yaml +++ b/sdk/python/nemo-platform/.nmpcontext/openapi.yaml @@ -903,6 +903,16 @@ paths: title: Parent type: string description: Parent entity ID for nested entities + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the entity still has this version. responses: '200': description: Successful Response @@ -3267,11 +3277,23 @@ paths: schema: type: string title: Name + - name: expected_db_version + in: query + required: false + schema: + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. + title: Expected Db Version + type: integer + description: Optional database version for optimistic locking. Delete only + succeeds if the VirtualModel still has this version. responses: '204': description: VirtualModel deleted '404': description: VirtualModel not found. + '409': + description: VirtualModel was modified before it could be deleted. '422': description: Validation Error content: @@ -12412,6 +12434,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -12422,6 +12449,7 @@ components: - updated_by - entity_id - parent + - db_version title: GuardrailConfig description: A guardrail configuration entity. GuardrailConfigFilter: @@ -15981,6 +16009,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -15992,6 +16025,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobStep description: 'A single step within an attempt. @@ -16312,6 +16346,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -16323,6 +16362,7 @@ components: - updated_by - entity_id - parent + - db_version title: PlatformJobTask description: 'A task within a step (for parallel execution). @@ -19324,6 +19364,11 @@ components: description: Parent entity ID for nested entities. readOnly: true type: string + db_version: + type: integer + title: Db Version + description: Database version of the entity for optimistic locking. + readOnly: true type: object required: - workspace @@ -19334,6 +19379,7 @@ components: - updated_by - entity_id - parent + - db_version title: VirtualModel description: 'Logical inference route. diff --git a/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/inference/virtual_models.py b/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/inference/virtual_models.py index ea1b6a2336..7536b725fb 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/inference/virtual_models.py +++ b/sdk/python/nemo-platform/src/nemo_platform/cli/commands/api/inference/virtual_models.py @@ -177,6 +177,13 @@ def delete_virtual_models( ctx: typer.Context, name: Annotated[str, typer.Argument()], workspace: Annotated[str | None, typer.Option("--workspace")] = None, + expected_db_version: Annotated[ + int | None, + typer.Option( + "--expected-db-version", + help="Optional database version for optimistic locking. Delete only succeeds if the VirtualModel still has this version.", + ), + ] = None, ) -> None: """Permanently delete a VirtualModel. @@ -187,6 +194,7 @@ def delete_virtual_models( kwargs = build_kwargs( workspace=workspace, + expected_db_version=expected_db_version, ) client.inference.virtual_models.delete(name, **kwargs) diff --git a/sdk/python/nemo-platform/src/nemo_platform/resources/entities/entities.py b/sdk/python/nemo-platform/src/nemo_platform/resources/entities/entities.py index 881cc756b8..469cab7213 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/resources/entities/entities.py +++ b/sdk/python/nemo-platform/src/nemo_platform/resources/entities/entities.py @@ -242,6 +242,7 @@ def delete_entity_by_name( *, workspace: str | None = None, entity_type: str, + expected_db_version: int | Omit = omit, parent: str | Omit = omit, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. @@ -260,6 +261,9 @@ def delete_entity_by_name( ``` Args: + expected_db_version: Optional database version for optimistic locking. Delete only succeeds if the + entity still has this version. + parent: Parent entity ID for nested entities extra_headers: Send extra headers @@ -291,7 +295,11 @@ def delete_entity_by_name( extra_body=extra_body, timeout=timeout, query=maybe_transform( - {"parent": parent}, entity_delete_entity_by_name_params.EntityDeleteEntityByNameParams + { + "expected_db_version": expected_db_version, + "parent": parent, + }, + entity_delete_entity_by_name_params.EntityDeleteEntityByNameParams, ), ), cast_to=DeleteResponse, @@ -681,6 +689,7 @@ async def delete_entity_by_name( *, workspace: str | None = None, entity_type: str, + expected_db_version: int | Omit = omit, parent: str | Omit = omit, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. @@ -699,6 +708,9 @@ async def delete_entity_by_name( ``` Args: + expected_db_version: Optional database version for optimistic locking. Delete only succeeds if the + entity still has this version. + parent: Parent entity ID for nested entities extra_headers: Send extra headers @@ -730,7 +742,11 @@ async def delete_entity_by_name( extra_body=extra_body, timeout=timeout, query=await async_maybe_transform( - {"parent": parent}, entity_delete_entity_by_name_params.EntityDeleteEntityByNameParams + { + "expected_db_version": expected_db_version, + "parent": parent, + }, + entity_delete_entity_by_name_params.EntityDeleteEntityByNameParams, ), ), cast_to=DeleteResponse, diff --git a/sdk/python/nemo-platform/src/nemo_platform/resources/inference/api.md b/sdk/python/nemo-platform/src/nemo_platform/resources/inference/api.md index 394f32b395..e50ab3c913 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/resources/inference/api.md +++ b/sdk/python/nemo-platform/src/nemo_platform/resources/inference/api.md @@ -21,7 +21,7 @@ Methods: - client.inference.virtual_models.create(\*, workspace, \*\*params) -> VirtualModel - client.inference.virtual_models.retrieve(name, \*, workspace) -> VirtualModel - client.inference.virtual_models.list(\*, workspace, \*\*params) -> SyncDefaultPagination[VirtualModel] -- client.inference.virtual_models.delete(name, \*, workspace) -> None +- client.inference.virtual_models.delete(name, \*, workspace, \*\*params) -> None - client.inference.virtual_models.patch(name, \*, workspace, \*\*params) -> VirtualModel ## Models diff --git a/sdk/python/nemo-platform/src/nemo_platform/resources/inference/virtual_models.py b/sdk/python/nemo-platform/src/nemo_platform/resources/inference/virtual_models.py index 24494a9f3a..a1956e42b6 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/resources/inference/virtual_models.py +++ b/sdk/python/nemo-platform/src/nemo_platform/resources/inference/virtual_models.py @@ -32,17 +32,18 @@ async_to_streamed_response_wrapper, ) from ...pagination import SyncDefaultPagination, AsyncDefaultPagination +from ..._exceptions import ConflictError from ..._base_client import AsyncPaginator, make_request_options from ...types.inference import ( virtual_model_list_params, virtual_model_patch_params, virtual_model_create_params, + virtual_model_delete_params, ) from ...types.inference.virtual_model import VirtualModel from ...types.inference.middleware_call_param import MiddlewareCallParam from ...types.inference.virtual_model_filter_param import VirtualModelFilterParam from ...types.inference.virtual_model_inference_config_param import VirtualModelInferenceConfigParam -from ..._exceptions import ConflictError __all__ = ["VirtualModelsResource", "AsyncVirtualModelsResource"] @@ -165,7 +166,7 @@ def create( except ConflictError: if not exist_ok: raise - return self.retrieve(name = name, workspace = workspace) + return self.retrieve(name=name, workspace=workspace) def retrieve( self, @@ -282,6 +283,7 @@ def delete( name: str, *, workspace: str | None = None, + expected_db_version: int | Omit = omit, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. extra_headers: Headers | None = None, @@ -296,6 +298,9 @@ def delete( VirtualModel. IGW's model cache is refreshed on its next polling cycle. Args: + expected_db_version: Optional database version for optimistic locking. Delete only succeeds if the + VirtualModel still has this version. + extra_headers: Send extra headers extra_query: Add additional query parameters to the request @@ -318,7 +323,13 @@ def delete( name=name, ), options=make_request_options( - extra_headers=extra_headers, extra_query=extra_query, extra_body=extra_body, timeout=timeout + extra_headers=extra_headers, + extra_query=extra_query, + extra_body=extra_body, + timeout=timeout, + query=maybe_transform( + {"expected_db_version": expected_db_version}, virtual_model_delete_params.VirtualModelDeleteParams + ), ), cast_to=NoneType, ) @@ -535,7 +546,7 @@ async def create( except ConflictError: if not exist_ok: raise - return await self.retrieve(name = name, workspace = workspace) + return await self.retrieve(name=name, workspace=workspace) async def retrieve( self, @@ -652,6 +663,7 @@ async def delete( name: str, *, workspace: str | None = None, + expected_db_version: int | Omit = omit, # Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs. # The extra values given here take precedence over values defined on the client or passed to this method. extra_headers: Headers | None = None, @@ -666,6 +678,9 @@ async def delete( VirtualModel. IGW's model cache is refreshed on its next polling cycle. Args: + expected_db_version: Optional database version for optimistic locking. Delete only succeeds if the + VirtualModel still has this version. + extra_headers: Send extra headers extra_query: Add additional query parameters to the request @@ -688,7 +703,13 @@ async def delete( name=name, ), options=make_request_options( - extra_headers=extra_headers, extra_query=extra_query, extra_body=extra_body, timeout=timeout + extra_headers=extra_headers, + extra_query=extra_query, + extra_body=extra_body, + timeout=timeout, + query=await async_maybe_transform( + {"expected_db_version": expected_db_version}, virtual_model_delete_params.VirtualModelDeleteParams + ), ), cast_to=NoneType, ) diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/entities/entity_delete_entity_by_name_params.py b/sdk/python/nemo-platform/src/nemo_platform/types/entities/entity_delete_entity_by_name_params.py index 2df520f36f..23a1e395c6 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/entities/entity_delete_entity_by_name_params.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/entities/entity_delete_entity_by_name_params.py @@ -27,5 +27,11 @@ class EntityDeleteEntityByNameParams(TypedDict, total=False): entity_type: Required[str] + expected_db_version: int + """Optional database version for optimistic locking. + + Delete only succeeds if the entity still has this version. + """ + parent: str """Parent entity ID for nested entities""" diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/guardrail/guardrail_config.py b/sdk/python/nemo-platform/src/nemo_platform/types/guardrail/guardrail_config.py index 605e53cf32..f9ac51543b 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/guardrail/guardrail_config.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/guardrail/guardrail_config.py @@ -33,6 +33,9 @@ class GuardrailConfig(BaseModel): created_by: Optional[str] = None + db_version: int + """Database version of the entity for optimistic locking.""" + entity_id: str """Alias for id for backwards compatibility.""" diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/inference/__init__.py b/sdk/python/nemo-platform/src/nemo_platform/types/inference/__init__.py index 2e2b34faaa..1d2d80d999 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/inference/__init__.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/inference/__init__.py @@ -60,6 +60,7 @@ from .virtual_model_patch_params import VirtualModelPatchParams as VirtualModelPatchParams from .model_provider_filter_param import ModelProviderFilterParam as ModelProviderFilterParam from .virtual_model_create_params import VirtualModelCreateParams as VirtualModelCreateParams +from .virtual_model_delete_params import VirtualModelDeleteParams as VirtualModelDeleteParams from .deployment_config_list_params import DeploymentConfigListParams as DeploymentConfigListParams from .k8s_nim_operator_config_param import K8sNIMOperatorConfigParam as K8sNIMOperatorConfigParam from .model_deployment_configs_page import ModelDeploymentConfigsPage as ModelDeploymentConfigsPage diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model.py b/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model.py index ddb72178f4..5350be63fc 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model.py @@ -52,6 +52,9 @@ class VirtualModel(BaseModel): created_by: Optional[str] = None + db_version: int + """Database version of the entity for optimistic locking.""" + entity_id: str """Alias for id for backwards compatibility.""" diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model_delete_params.py b/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model_delete_params.py new file mode 100644 index 0000000000..67c8ca6aa8 --- /dev/null +++ b/sdk/python/nemo-platform/src/nemo_platform/types/inference/virtual_model_delete_params.py @@ -0,0 +1,32 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +# File generated from our OpenAPI spec by Stainless. See CONTRIBUTING.md for details. + +from __future__ import annotations + +from typing_extensions import TypedDict + +__all__ = ["VirtualModelDeleteParams"] + + +class VirtualModelDeleteParams(TypedDict, total=False): + workspace: str + + expected_db_version: int + """Optional database version for optimistic locking. + + Delete only succeeds if the VirtualModel still has this version. + """ diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_step.py b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_step.py index 19a01b2529..632d94bf78 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_step.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_step.py @@ -39,6 +39,9 @@ class PlatformJobStep(BaseModel): created_by: Optional[str] = None + db_version: int + """Database version of the entity for optimistic locking.""" + entity_id: str """Alias for id for backwards compatibility.""" diff --git a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_task.py b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_task.py index f76f9566c7..5134adb72c 100644 --- a/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_task.py +++ b/sdk/python/nemo-platform/src/nemo_platform/types/jobs/platform_job_task.py @@ -36,6 +36,9 @@ class PlatformJobTask(BaseModel): created_by: Optional[str] = None + db_version: int + """Database version of the entity for optimistic locking.""" + entity_id: str """Alias for id for backwards compatibility.""" diff --git a/sdk/python/nemo-platform/tests/api_resources/inference/test_virtual_models.py b/sdk/python/nemo-platform/tests/api_resources/inference/test_virtual_models.py index f2ce1cf4e6..c18a4eef51 100644 --- a/sdk/python/nemo-platform/tests/api_resources/inference/test_virtual_models.py +++ b/sdk/python/nemo-platform/tests/api_resources/inference/test_virtual_models.py @@ -263,6 +263,16 @@ def test_method_delete(self, client: NeMoPlatform) -> None: ) assert virtual_model is None + @pytest.mark.skip(reason="Mock server tests are disabled") + @parametrize + def test_method_delete_with_all_params(self, client: NeMoPlatform) -> None: + virtual_model = client.inference.virtual_models.delete( + name="name", + workspace="workspace", + expected_db_version=0, + ) + assert virtual_model is None + @pytest.mark.skip(reason="Mock server tests are disabled") @parametrize def test_raw_response_delete(self, client: NeMoPlatform) -> None: @@ -633,6 +643,16 @@ async def test_method_delete(self, async_client: AsyncNeMoPlatform) -> None: ) assert virtual_model is None + @pytest.mark.skip(reason="Mock server tests are disabled") + @parametrize + async def test_method_delete_with_all_params(self, async_client: AsyncNeMoPlatform) -> None: + virtual_model = await async_client.inference.virtual_models.delete( + name="name", + workspace="workspace", + expected_db_version=0, + ) + assert virtual_model is None + @pytest.mark.skip(reason="Mock server tests are disabled") @parametrize async def test_raw_response_delete(self, async_client: AsyncNeMoPlatform) -> None: diff --git a/sdk/python/nemo-platform/tests/api_resources/test_entities.py b/sdk/python/nemo-platform/tests/api_resources/test_entities.py index 57781ebf66..393f01db93 100644 --- a/sdk/python/nemo-platform/tests/api_resources/test_entities.py +++ b/sdk/python/nemo-platform/tests/api_resources/test_entities.py @@ -189,6 +189,7 @@ def test_method_delete_entity_by_name_with_all_params(self, client: NeMoPlatform name="name", workspace="workspace", entity_type="entity_type", + expected_db_version=0, parent="parent", ) assert_matches_type(DeleteResponse, entity, path=["response"]) @@ -608,6 +609,7 @@ async def test_method_delete_entity_by_name_with_all_params(self, async_client: name="name", workspace="workspace", entity_type="entity_type", + expected_db_version=0, parent="parent", ) assert_matches_type(DeleteResponse, entity, path=["response"]) diff --git a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services.py b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services.py index e2b2bca695..a17a2c2bdb 100644 --- a/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services.py +++ b/sdk/python/nemo-platform/tests/vendored/nemo_platform_ext/local/test_services.py @@ -680,20 +680,38 @@ def test_daemonize_services_bounds_probe_and_sleep_by_remaining_deadline( proc = MagicMock() proc.pid = 4242 proc.poll.return_value = None + clock = 0.0 + sleep_calls: list[float] = [] + + def monotonic() -> float: + nonlocal clock + if clock == 0.0: + clock = 4.0 + return 0.0 + return clock + + def probe_status(*_args: object, **_kwargs: object) -> bool: + nonlocal clock + clock = 4.5 + return False + + def sleep(duration: float) -> None: + nonlocal clock + sleep_calls.append(duration) + clock += duration with ( patch("nemo_platform.local.services.require_services_extra"), - patch("nemo_platform.local.services.probe_status", return_value=False) as probe_status, + patch("nemo_platform.local.services.probe_status", side_effect=probe_status) as probe_status_mock, patch("nemo_platform.local.services.subprocess.Popen", return_value=proc), - patch("nemo_platform.local.services.time.monotonic", side_effect=[0.0, 4.0, 4.5, 5.0]), - patch("nemo_platform.local.services.time.sleep") as sleep, + patch("nemo_platform.local.services.time.monotonic", side_effect=monotonic), + patch("nemo_platform.local.services.time.sleep", side_effect=sleep), ): with pytest.raises(services.ServicesStartupTimeoutError): services.daemonize_services(cfg) - assert probe_status.call_args.kwargs["timeout"] == pytest.approx(1.0) - sleep.assert_called_once() - assert sleep.call_args.args[0] == pytest.approx(0.5) + assert probe_status_mock.call_args.kwargs["timeout"] == pytest.approx(1.0) + assert sleep_calls == [pytest.approx(0.5)] proc.terminate.assert_called_once_with() diff --git a/services/core/entities/src/nmp/core/entities/api/v2/entities/endpoints.py b/services/core/entities/src/nmp/core/entities/api/v2/entities/endpoints.py index d0fe28e5c3..78a8a1b001 100644 --- a/services/core/entities/src/nmp/core/entities/api/v2/entities/endpoints.py +++ b/services/core/entities/src/nmp/core/entities/api/v2/entities/endpoints.py @@ -591,6 +591,10 @@ async def delete_entity_by_name( workspace_repository: WorkspaceRepository, auth_client: AuthClientDep, parent: str | None = Query(default=None, description="Parent entity ID for nested entities"), + expected_db_version: int | None = Query( + default=None, + description="Optional database version for optimistic locking. Delete only succeeds if the entity still has this version.", + ), ) -> DeleteResponse: """Delete entity by name.""" # Check if workspace is being deleted (404 for user requests) @@ -603,12 +607,19 @@ async def delete_entity_by_name( await _invalidate_role_binding_cache_if_present(repository, workspace, entity_type, name, parent) - deleted_count = await repository.delete_entity_by_name( - workspace=workspace, - entity_type=entity_type, - name=name, - parent=parent, - ) + try: + deleted_count = await repository.delete_entity_by_name( + workspace=workspace, + entity_type=entity_type, + name=name, + parent=parent, + expected_db_version=expected_db_version, + ) + except EntityVersionConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail=str(e), + ) from e if deleted_count == 0: raise HTTPException(status_code=404, detail="Entity not found") return DeleteResponse(id=f"{workspace}/{entity_type}/{name}") diff --git a/services/core/entities/src/nmp/core/entities/app/repository/entity.py b/services/core/entities/src/nmp/core/entities/app/repository/entity.py index ef8e168a2d..16e5f365f2 100644 --- a/services/core/entities/src/nmp/core/entities/app/repository/entity.py +++ b/services/core/entities/src/nmp/core/entities/app/repository/entity.py @@ -194,6 +194,7 @@ async def delete_entity_by_name( entity_type: str, name: str, parent: Optional[str] = None, + expected_db_version: Optional[int] = None, session: AsyncSession | None = None, ) -> int: """Delete an entity by name. @@ -203,6 +204,8 @@ async def delete_entity_by_name( entity_type: Entity type name: Entity name parent: Optional parent entity ID (None for root entities) + expected_db_version: Optional expected database version for optimistic locking. If provided, + delete only succeeds when the stored entity still has this version. Returns: Number of deleted entities (0 or 1) diff --git a/services/core/entities/src/nmp/core/entities/app/repository/sqlalchemy/entity.py b/services/core/entities/src/nmp/core/entities/app/repository/sqlalchemy/entity.py index 202b29612f..24ce3582a4 100644 --- a/services/core/entities/src/nmp/core/entities/app/repository/sqlalchemy/entity.py +++ b/services/core/entities/src/nmp/core/entities/app/repository/sqlalchemy/entity.py @@ -361,6 +361,7 @@ async def delete_entity_by_name( entity_type: str, name: str, parent: Optional[str] = None, + expected_db_version: int | None = None, session: AsyncSession | None = None, ) -> int: """Delete an entity by name.""" @@ -380,6 +381,18 @@ async def delete_entity_by_name( if existing_entity is None: return 0 - await sess.delete(existing_entity) - await sess.commit() + if expected_db_version is not None and existing_entity.db_version != expected_db_version: + raise EntityVersionConflictError( + f"Entity '{name}' of type '{entity_type}' in workspace '{workspace}' was modified by another request. " + f"Expected version {expected_db_version}, but current version is {existing_entity.db_version}. Please refetch and retry." + ) + + try: + await sess.delete(existing_entity) + await sess.commit() + except StaleDataError as err: + await sess.rollback() + raise EntityVersionConflictError( + f"Entity '{name}' of type '{entity_type}' in workspace '{workspace}' was modified by another request. Please refetch and retry." + ) from err return 1 diff --git a/services/core/entities/tests/repository/test_entity_delete_versioning.py b/services/core/entities/tests/repository/test_entity_delete_versioning.py new file mode 100644 index 0000000000..a03d97077b --- /dev/null +++ b/services/core/entities/tests/repository/test_entity_delete_versioning.py @@ -0,0 +1,118 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Tests for optimistic locking on entity deletes.""" + +import pytest +from nmp.core.entities.app.repository import SQLAlchemyEntityRepository +from nmp.core.entities.app.repository.exceptions import EntityVersionConflictError +from nmp.core.entities.app.repository.sqlalchemy.models import DBEntity +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker + + +@pytest.mark.asyncio +class TestEntityDeleteVersioning: + async def test_delete_entity_by_name_with_matching_expected_db_version( + self, entity_repo: SQLAlchemyEntityRepository, setup_workspaces + ): + entity = await entity_repo.create_entity( + workspace="workspace-1", + entity_type="config", + name="delete-with-version", + data={"value": 1}, + ) + + deleted_count = await entity_repo.delete_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-with-version", + expected_db_version=entity.db_version, + ) + + assert deleted_count == 1 + assert ( + await entity_repo.get_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-with-version", + ) + is None + ) + + async def test_delete_entity_by_name_rejects_stale_expected_db_version( + self, entity_repo: SQLAlchemyEntityRepository, setup_workspaces + ): + entity = await entity_repo.create_entity( + workspace="workspace-1", + entity_type="config", + name="delete-stale-version", + data={"value": 1}, + ) + updated = await entity_repo.update_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-stale-version", + data={"value": 2}, + ) + + with pytest.raises(EntityVersionConflictError): + await entity_repo.delete_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-stale-version", + expected_db_version=entity.db_version, + ) + + assert updated.db_version != entity.db_version + assert ( + await entity_repo.get_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-stale-version", + ) + is not None + ) + + async def test_delete_entity_by_name_rolls_back_stale_commit( + self, + entity_repo: SQLAlchemyEntityRepository, + session_maker: async_sessionmaker[AsyncSession], + setup_workspaces, + ): + entity = await entity_repo.create_entity( + workspace="workspace-1", + entity_type="config", + name="delete-stale-commit", + data={"value": 1}, + ) + + async with session_maker() as shared_session: + result = await shared_session.execute(select(DBEntity).where(DBEntity.id == entity.id)) + stale_entity = result.scalar_one() + await shared_session.commit() + + await entity_repo.update_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-stale-commit", + data={"value": 2}, + ) + + assert stale_entity.db_version == entity.db_version + with pytest.raises(EntityVersionConflictError): + await entity_repo.delete_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-stale-commit", + session=shared_session, + ) + + remaining = await entity_repo.get_entity_by_name( + workspace="workspace-1", + entity_type="config", + name="delete-stale-commit", + session=shared_session, + ) + assert remaining is not None + assert remaining.data == {"value": 2} diff --git a/services/core/files/src/nmp/core/files/app/file_lock.py b/services/core/files/src/nmp/core/files/app/file_lock.py index 914b747b01..9f5fb3be33 100644 --- a/services/core/files/src/nmp/core/files/app/file_lock.py +++ b/services/core/files/src/nmp/core/files/app/file_lock.py @@ -64,8 +64,8 @@ async def acquire(self, path: str) -> AsyncIterator[bool]: True if lock acquired and caller should proceed with write False if lock not acquired and caller should skip write """ - acquired = await self._try_acquire(path) - if not acquired: + lock = await self._try_acquire(path) + if lock is None: logger.debug("Lock not acquired, another request is writing") yield False return @@ -74,16 +74,16 @@ async def acquire(self, path: str) -> AsyncIterator[bool]: try: yield True finally: - await self._release(path) + await self._release(lock) - async def _try_acquire(self, path: str, max_attempts: int = 3) -> bool: + async def _try_acquire(self, path: str, max_attempts: int = 3) -> FileLock | None: """Attempt to acquire the lock, handling conflicts and expiry. Args: path: The file path to lock max_attempts: Maximum number of attempts before giving up - Returns True if lock acquired, False if should skip. + Returns the acquired lock entity if lock acquired, None if should skip. """ lock_name = _path_to_lock_name(path) @@ -96,8 +96,8 @@ async def _try_acquire(self, path: str, max_attempts: int = 3) -> bool: acquired_at=datetime.now(UTC), ) try: - await self.entity_client.create(lock) - return True + created = await self.entity_client.create(lock) + return created or lock except EntityConflictError: pass # Lock exists, check if we should retry @@ -112,7 +112,7 @@ async def _try_acquire(self, path: str, max_attempts: int = 3) -> bool: if expires_at >= datetime.now(UTC): # Lock is fresh - another request is actively writing logger.debug("Fresh lock held by another request") - return False + return None # Lock is stale - atomically take it over by updating with version check. # Only succeeds if db_version hasn't changed (no one else took it). @@ -120,10 +120,10 @@ async def _try_acquire(self, path: str, max_attempts: int = 3) -> bool: try: # Update the entity - db_version is automatically included for optimistic locking existing.acquired_at = datetime.now(UTC) - await self.entity_client.update(existing) + updated = await self.entity_client.update(existing) # Successfully took over the stale lock logger.debug("Successfully took over stale lock") - return True + return updated or existing except EntityConflictError: # Version changed - someone else took the lock, retry on next iteration logger.debug("Update failed, lock was modified by another request") @@ -131,17 +131,24 @@ async def _try_acquire(self, path: str, max_attempts: int = 3) -> bool: pass # Someone else deleted it, that's fine logger.debug("Lock acquisition attempts exhausted") - return False + return None - async def _release(self, path: str) -> None: + async def _release(self, lock: FileLock) -> None: """Release the lock.""" - lock_name = _path_to_lock_name(path) try: - await self.entity_client.delete(FileLock, lock_name, workspace=self.workspace) + await self.entity_client.delete( + FileLock, + lock.name, + workspace=self.workspace, + expected_db_version=lock.db_version, + ) logger.debug("Lock released") except EntityNotFoundError: # Lock already gone (expired and cleaned up by another request) logger.debug("Lock already released (expired)") + except EntityConflictError: + # Another request reacquired the lock after this holder lost ownership. + logger.debug("Lock version changed before release; leaving current lock in place") async def get_active_locks(self, paths: list[str]) -> set[str]: """Get paths that have active (non-expired) locks. diff --git a/services/core/files/tests/test_file_lock.py b/services/core/files/tests/test_file_lock.py index 0b6cac875e..12a49550b5 100644 --- a/services/core/files/tests/test_file_lock.py +++ b/services/core/files/tests/test_file_lock.py @@ -46,7 +46,12 @@ async def test_acquire_succeeds_when_no_existing_lock(lock_manager, mock_entity_ assert created_lock.workspace == "test-workspace" # Should have released the lock - mock_entity_client.delete.assert_called_once_with(FileLock, expected_lock_name, workspace="test-workspace") + mock_entity_client.delete.assert_called_once_with( + FileLock, + expected_lock_name, + workspace="test-workspace", + expected_db_version=1, + ) async def test_acquire_fails_when_fresh_lock_exists(lock_manager, mock_entity_client): @@ -118,7 +123,12 @@ async def test_acquire_succeeds_after_cleaning_stale_lock(lock_manager, mock_ent assert updated_entity.acquired_at > stale_time assert updated_entity._db_version == 2 # db_version from stale lock is preserved # Lock should be released when context manager exits (delete is called in _release) - mock_entity_client.delete.assert_called_once_with(FileLock, lock_name, workspace="test-workspace") + mock_entity_client.delete.assert_called_once_with( + FileLock, + lock_name, + workspace="test-workspace", + expected_db_version=3, + ) async def test_acquire_handles_update_version_conflict(lock_manager, mock_entity_client): @@ -262,7 +272,12 @@ async def test_lock_released_on_success(lock_manager, mock_entity_client): await asyncio.sleep(0.01) # Lock should be released - mock_entity_client.delete.assert_called_with(FileLock, expected_lock_name, workspace="test-workspace") + mock_entity_client.delete.assert_called_with( + FileLock, + expected_lock_name, + workspace="test-workspace", + expected_db_version=1, + ) async def test_lock_released_on_exception(lock_manager, mock_entity_client): @@ -277,7 +292,12 @@ async def test_lock_released_on_exception(lock_manager, mock_entity_client): raise ValueError("test error") # Lock should still be released - mock_entity_client.delete.assert_called_with(FileLock, expected_lock_name, workspace="test-workspace") + mock_entity_client.delete.assert_called_with( + FileLock, + expected_lock_name, + workspace="test-workspace", + expected_db_version=1, + ) async def test_lock_not_released_if_already_gone(lock_manager, mock_entity_client): @@ -293,6 +313,17 @@ async def test_lock_not_released_if_already_gone(lock_manager, mock_entity_clien mock_entity_client.delete.assert_called_once() +async def test_lock_not_released_if_version_changed(lock_manager, mock_entity_client): + """Test that a lock reacquired by another request is not deleted by stale release.""" + mock_entity_client.create.return_value = None + mock_entity_client.delete.side_effect = EntityConflictError("changed") + + async with lock_manager.acquire("test/file.bin") as acquired: + assert acquired is True + + mock_entity_client.delete.assert_called_once() + + async def test_conflict_then_lock_deleted_before_get(lock_manager, mock_entity_client): """Test handling when lock is deleted between create and get.""" # Track create calls - first fails, second (after retry) succeeds diff --git a/services/core/inference-gateway/src/nmp/core/inference_gateway/api/v2/virtual_models.py b/services/core/inference-gateway/src/nmp/core/inference_gateway/api/v2/virtual_models.py index 2869b4333d..cbb8ffb06e 100644 --- a/services/core/inference-gateway/src/nmp/core/inference_gateway/api/v2/virtual_models.py +++ b/services/core/inference-gateway/src/nmp/core/inference_gateway/api/v2/virtual_models.py @@ -511,12 +511,17 @@ async def update_virtual_model( status_code=status.HTTP_204_NO_CONTENT, responses={ 404: {"description": "VirtualModel not found."}, + 409: {"description": "VirtualModel was modified before it could be deleted."}, }, ) async def delete_virtual_model( workspace: str, name: str, entity_client: EntityClientDep, + expected_db_version: int | None = Query( + default=None, + description="Optional database version for optimistic locking. Delete only succeeds if the VirtualModel still has this version.", + ), ) -> None: """Permanently delete a VirtualModel. @@ -524,12 +529,22 @@ async def delete_virtual_model( this VirtualModel. IGW's model cache is refreshed on its next polling cycle. """ try: - await entity_client.delete(VirtualModel, name=name, workspace=workspace) + await entity_client.delete( + VirtualModel, + name=name, + workspace=workspace, + expected_db_version=expected_db_version, + ) except EntityNotFoundError as exc: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail=f"VirtualModel '{name}' not found in workspace '{workspace}'.", ) from exc + except EntityConflictError as exc: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, + detail="Concurrent modification - please retry.", + ) from exc except HTTPException: raise except Exception: diff --git a/services/core/inference-gateway/tests/integration/test_inference.py b/services/core/inference-gateway/tests/integration/test_inference.py index a9ea4bf54f..8b31494ed8 100644 --- a/services/core/inference-gateway/tests/integration/test_inference.py +++ b/services/core/inference-gateway/tests/integration/test_inference.py @@ -252,6 +252,7 @@ def _manually_add_provider_to_cache( workspace=ws, name=entity_name, parent=ws, + db_version=1, default_model_entity=f"{ws}/{entity_name}", autoprovisioned=True, created_at=now_iso, diff --git a/services/core/inference-gateway/tests/integration/test_middleware_pipeline.py b/services/core/inference-gateway/tests/integration/test_middleware_pipeline.py index 5ced30d534..3f63ca73d0 100644 --- a/services/core/inference-gateway/tests/integration/test_middleware_pipeline.py +++ b/services/core/inference-gateway/tests/integration/test_middleware_pipeline.py @@ -316,6 +316,7 @@ def _inject_vm_and_plugins( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=default_model_entity, diff --git a/services/core/inference-gateway/tests/unit/conftest.py b/services/core/inference-gateway/tests/unit/conftest.py index a81eea9334..3aeabfd212 100644 --- a/services/core/inference-gateway/tests/unit/conftest.py +++ b/services/core/inference-gateway/tests/unit/conftest.py @@ -114,6 +114,7 @@ def autoprovisioned_vms_for_cache(model_cache: ModelCache) -> list[VirtualModel] workspace=workspace, name=name, parent=workspace, + db_version=1, default_model_entity=f"{workspace}/{name}", autoprovisioned=True, created_at=now, diff --git a/services/core/inference-gateway/tests/unit/test_middleware_registry.py b/services/core/inference-gateway/tests/unit/test_middleware_registry.py index a69d537d1f..ab044fbc99 100644 --- a/services/core/inference-gateway/tests/unit/test_middleware_registry.py +++ b/services/core/inference-gateway/tests/unit/test_middleware_registry.py @@ -57,6 +57,7 @@ def _make_sdk_vm( name=name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at=updated_at, default_model_entity=f"{workspace}/{name}", @@ -514,6 +515,7 @@ async def test_vm_with_none_name_is_skipped(self): name=None, workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", ) @@ -846,6 +848,7 @@ def test_sdk_vm_to_plugin_vm_handles_none_name(): name=None, workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", ) diff --git a/services/core/inference-gateway/tests/unit/test_model_router.py b/services/core/inference-gateway/tests/unit/test_model_router.py index f873f06f8f..d5e81489d2 100644 --- a/services/core/inference-gateway/tests/unit/test_model_router.py +++ b/services/core/inference-gateway/tests/unit/test_model_router.py @@ -36,6 +36,7 @@ def _autoprovisioned_vms_for_cache(model_cache: ModelCache) -> list[SDKVirtualMo workspace=workspace, name=name, parent=workspace, + db_version=1, default_model_entity=f"{workspace}/{name}", autoprovisioned=True, created_at="2026-01-01T00:00:00Z", @@ -452,6 +453,7 @@ def _make_sdk_vm( name=name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=default_model_entity or f"{workspace}/{name}", @@ -500,6 +502,7 @@ def test_model_entity_proxy_virtual_model_no_default_model_entity_no_middleware_ name="mw-only", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=None, diff --git a/services/core/inference-gateway/tests/unit/test_openai_router.py b/services/core/inference-gateway/tests/unit/test_openai_router.py index 527c3233d2..c0e26e3c4c 100644 --- a/services/core/inference-gateway/tests/unit/test_openai_router.py +++ b/services/core/inference-gateway/tests/unit/test_openai_router.py @@ -44,6 +44,7 @@ def _autoprovisioned_vms_for_cache(model_cache: ModelCache) -> list[SDKVirtualMo workspace=workspace, name=name, parent=workspace, + db_version=1, default_model_entity=f"{workspace}/{name}", autoprovisioned=True, created_at="2026-01-01T00:00:00Z", @@ -62,6 +63,7 @@ def _custom_vm(workspace: str, name: str, default_model_entity: str | None = Non workspace=workspace, name=name, parent=workspace, + db_version=1, default_model_entity=default_model_entity, autoprovisioned=False, created_at="2026-01-01T00:00:00Z", @@ -445,6 +447,7 @@ def _make_lora_vm(workspace: str, name: str) -> SDKVirtualModel: name=name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=f"{workspace}/{name}", @@ -648,6 +651,7 @@ def test_proxy_no_providers_for_model_entity(app: FastAPI, client: TestClient): workspace="ns1", name="orphan-model", parent="ns1", + db_version=1, default_model_entity="ns1/orphan-model", autoprovisioned=True, created_at="2026-01-01T00:00:00Z", @@ -701,6 +705,7 @@ def test_proxy_with_unresolved_provider_secret(app: FastAPI, client: TestClient) workspace="ns1", name="secure-model", parent="ns1", + db_version=1, default_model_entity="ns1/secure-model", autoprovisioned=True, created_at="2026-01-01T00:00:00Z", @@ -805,6 +810,7 @@ def test_proxy_resolves_served_model_name_with_slashes(app: FastAPI, client: Tes workspace="ns1", name="my-model", parent="ns1", + db_version=1, default_model_entity="ns1/my-model", autoprovisioned=True, created_at="2026-01-01T00:00:00Z", @@ -883,6 +889,7 @@ def _make_sdk_vm( name=name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=default_model_entity or f"{workspace}/{name}", @@ -931,6 +938,7 @@ def test_openai_proxy_virtual_model_no_default_model_entity_no_middleware_return name="mw-only", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=None, diff --git a/services/core/inference-gateway/tests/unit/test_proxy.py b/services/core/inference-gateway/tests/unit/test_proxy.py index 922476f0a7..7b0fbce351 100644 --- a/services/core/inference-gateway/tests/unit/test_proxy.py +++ b/services/core/inference-gateway/tests/unit/test_proxy.py @@ -323,6 +323,7 @@ async def process_response( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -420,6 +421,7 @@ async def process_request( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -556,6 +558,7 @@ async def process_response( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -619,6 +622,7 @@ async def test_virtual_model_proxy_with_unresolved_provider_secret_returns_424(m name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -700,6 +704,7 @@ async def process_response( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", ), @@ -783,6 +788,7 @@ async def process_response( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", ), @@ -903,6 +909,7 @@ async def process_response( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -985,6 +992,7 @@ async def process_response(self, ctx, response, middleware_config): name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=f"{workspace}/missing", @@ -1802,6 +1810,7 @@ async def test_virtual_model_proxy_rewrites_response_model_non_streaming(mock_pr name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -1886,6 +1895,7 @@ async def _iter_any(): name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -1990,6 +2000,7 @@ async def _iter_any(): name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, @@ -2103,6 +2114,7 @@ async def process_request( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=None, @@ -2208,6 +2220,7 @@ async def process_response( name=vm_name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=model_entity_id, diff --git a/services/core/inference-gateway/tests/unit/test_virtual_model_cache.py b/services/core/inference-gateway/tests/unit/test_virtual_model_cache.py index 3cbdd3f0c2..80a2bb7a14 100644 --- a/services/core/inference-gateway/tests/unit/test_virtual_model_cache.py +++ b/services/core/inference-gateway/tests/unit/test_virtual_model_cache.py @@ -36,6 +36,7 @@ def _make_vm(workspace: str, name: str, default_model_entity: str | None = None) name=name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity=default_model_entity or f"{workspace}/{name}", @@ -244,6 +245,7 @@ def _make_vm_at(workspace: str, name: str, updated_at: str = "2026-01-01T00:00:0 name=name, workspace=workspace, parent=workspace, + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at=updated_at, default_model_entity=f"{workspace}/{name}", @@ -404,6 +406,7 @@ async def test_refresh_middleware_references_receives_deduped_config_refs(): name="a", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity="ws/m", @@ -415,6 +418,7 @@ async def test_refresh_middleware_references_receives_deduped_config_refs(): name="b", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity="ws/m", @@ -446,6 +450,7 @@ def __init__(self, ts): name="only", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity="ws/m", @@ -486,6 +491,7 @@ def __init__(self, ts): name="only", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity="ws/m", @@ -535,6 +541,7 @@ async def test_refresh_removed_vm_clears_broken_state(): name="only", workspace="ws", parent="ws", + db_version=1, created_at="2026-01-01T00:00:00Z", updated_at="2026-01-01T00:00:00Z", default_model_entity="ws/m", diff --git a/services/core/inference-gateway/tests/unit/test_virtual_models_router.py b/services/core/inference-gateway/tests/unit/test_virtual_models_router.py index 63d694055b..1d18a52bf9 100644 --- a/services/core/inference-gateway/tests/unit/test_virtual_models_router.py +++ b/services/core/inference-gateway/tests/unit/test_virtual_models_router.py @@ -538,3 +538,20 @@ def test_delete_not_found_returns_404(self, client: TestClient): """DELETE for a non-existent name → 404.""" resp = client.delete(f"{BASE}/no-such-vm") assert resp.status_code == 404 + + def test_delete_with_stale_expected_db_version_returns_409(self, client: TestClient): + """DELETE with a stale entity version is rejected and leaves the VM intact.""" + created = _create(client, "vm-stale-delete") + patch_resp = client.patch(f"{BASE}/vm-stale-delete", json={"autoprovisioned": True}) + assert patch_resp.status_code == 200 + assert patch_resp.json()["db_version"] != created["db_version"] + + resp = client.delete( + f"{BASE}/vm-stale-delete", + params={"expected_db_version": created["db_version"]}, + ) + + assert resp.status_code == 409 + + get_resp = client.get(f"{BASE}/vm-stale-delete") + assert get_resp.status_code == 200 diff --git a/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py b/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py index c67be6d441..9278971d8f 100644 --- a/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py +++ b/services/core/models/src/nmp/core/models/api/service/adapter_entity_service.py @@ -186,7 +186,12 @@ async def delete_adapter(self, adapter_workspace: str, parent_model_ref: str, ad ) return -2 - await self.entity_client.delete_by_id(Adapter, adapter.id) + await self.entity_client.delete( + Adapter, + adapter.name, + workspace=adapter_workspace, + parent=adapter.parent, + ) logger.info( f"Successfully deleted adapter {adapter_name} from model entity: {adapter_workspace}/{parent_model_ref}" ) @@ -243,7 +248,12 @@ async def delete_adapter_in_workspace(self, adapter_workspace: str, adapter_name raise ValueError(str(err)) from err if to_delete is None or not to_delete.parent: return -2 - await self.entity_client.delete(Adapter, to_delete.name, workspace=adapter_workspace, parent=to_delete.parent) + await self.entity_client.delete( + Adapter, + to_delete.name, + workspace=adapter_workspace, + parent=to_delete.parent, + ) logger.info(f"Successfully deleted adapter: {adapter_workspace}/{adapter_name}") return 0 diff --git a/services/core/models/src/nmp/core/models/api/service/model_deployment_config_service.py b/services/core/models/src/nmp/core/models/api/service/model_deployment_config_service.py index 95d599fb25..755b1c3241 100644 --- a/services/core/models/src/nmp/core/models/api/service/model_deployment_config_service.py +++ b/services/core/models/src/nmp/core/models/api/service/model_deployment_config_service.py @@ -407,7 +407,11 @@ async def delete_deployment_config(self, workspace: str, name: str, version: int logger.warning(f"Deployment config not found for deletion: {workspace}/{name} version {version}") return False - await self.entity_client.delete(ModelDeploymentConfigEntity, entity.name, workspace=workspace) + await self.entity_client.delete( + ModelDeploymentConfigEntity, + entity.name, + workspace=workspace, + ) logger.info(f"Successfully deleted deployment config: {workspace}/{name} version {version}") return True else: @@ -424,7 +428,11 @@ async def delete_deployment_config(self, workspace: str, name: str, version: int return False for entity in result.data: - await self.entity_client.delete(ModelDeploymentConfigEntity, entity.name, workspace=workspace) + await self.entity_client.delete( + ModelDeploymentConfigEntity, + entity.name, + workspace=workspace, + ) logger.info(f"Successfully deleted all versions of deployment config: {workspace}/{name}") return True diff --git a/services/core/models/src/nmp/core/models/api/service/model_deployment_service.py b/services/core/models/src/nmp/core/models/api/service/model_deployment_service.py index 5495f39d45..bea06b4176 100644 --- a/services/core/models/src/nmp/core/models/api/service/model_deployment_service.py +++ b/services/core/models/src/nmp/core/models/api/service/model_deployment_service.py @@ -483,7 +483,12 @@ async def delete_deployment(self, workspace: str, name: str, version: int | None # If already DELETED, hard delete from database if entity.status == ModelDeploymentStatus.DELETED: logger.info(f"Deployment already DELETED, removing from database: {workspace}/{name} version {version}") - await self.entity_client.delete(ModelDeploymentEntity, entity.name, workspace=workspace) + await self.entity_client.delete( + ModelDeploymentEntity, + entity.name, + workspace=workspace, + expected_db_version=entity.db_version, + ) return None # Otherwise, mark for deletion @@ -511,7 +516,12 @@ async def delete_deployment(self, workspace: str, name: str, version: int | None if all_deleted: logger.info(f"All versions already DELETED, removing from database: {workspace}/{name}") for entity in result.data: - await self.entity_client.delete(ModelDeploymentEntity, entity.name, workspace=workspace) + await self.entity_client.delete( + ModelDeploymentEntity, + entity.name, + workspace=workspace, + expected_db_version=entity.db_version, + ) return None # Mark all non-DELETED versions for deletion diff --git a/services/core/models/src/nmp/core/models/api/service/model_entity_service.py b/services/core/models/src/nmp/core/models/api/service/model_entity_service.py index c6e1ecb2c8..e1958ac756 100644 --- a/services/core/models/src/nmp/core/models/api/service/model_entity_service.py +++ b/services/core/models/src/nmp/core/models/api/service/model_entity_service.py @@ -541,8 +541,7 @@ async def delete_model_entity(self, workspace: str, name: str) -> bool: logger.debug(f"Deleting model entity: {workspace}/{name}") try: - model: Model = await self.entity_client.get(Model, workspace=workspace, name=name) - await self.entity_client.delete(Model, model.name, workspace=workspace) + await self.entity_client.delete(Model, name, workspace=workspace) logger.info(f"Successfully deleted model entity: {workspace}/{name}") return True except EntityNotFoundError: diff --git a/services/core/models/src/nmp/core/models/api/service/model_provider_service.py b/services/core/models/src/nmp/core/models/api/service/model_provider_service.py index f058e2cce6..a52ac4686d 100644 --- a/services/core/models/src/nmp/core/models/api/service/model_provider_service.py +++ b/services/core/models/src/nmp/core/models/api/service/model_provider_service.py @@ -294,9 +294,14 @@ async def delete_model_provider(self, request: DeleteModelProviderRequest) -> bo return False provider_id = f"{request.workspace}/{request.name}" - await self._cleanup_model_entity_references(provider, provider_id) + await self.entity_client.delete( + ModelProviderEntity, + provider.name, + workspace=request.workspace, + expected_db_version=provider.db_version, + ) - await self.entity_client.delete(ModelProviderEntity, request.name, workspace=request.workspace) + await self._cleanup_model_entity_references(provider, provider_id) logger.info("Model provider deleted", extra={"workspace": request.workspace, "provider_name": request.name}) return True diff --git a/services/core/models/src/nmp/core/models/api/service/prompt_service.py b/services/core/models/src/nmp/core/models/api/service/prompt_service.py index 9d068be2dd..7ef2058135 100644 --- a/services/core/models/src/nmp/core/models/api/service/prompt_service.py +++ b/services/core/models/src/nmp/core/models/api/service/prompt_service.py @@ -161,7 +161,7 @@ async def delete_prompt(self, request: DeletePromptRequest) -> bool: logger.debug("Deleting prompt", extra={"workspace": request.workspace, "prompt_name": request.name}) try: - await self.entity_client.get(PromptEntity, workspace=request.workspace, name=request.name) + await self.entity_client.delete(PromptEntity, request.name, workspace=request.workspace) except EntityNotFoundError: logger.warning( "Prompt not found for deletion", @@ -169,6 +169,5 @@ async def delete_prompt(self, request: DeletePromptRequest) -> bool: ) return False - await self.entity_client.delete(PromptEntity, request.name, workspace=request.workspace) logger.info("Prompt deleted", extra={"workspace": request.workspace, "prompt_name": request.name}) return True diff --git a/services/core/models/src/nmp/core/models/api/v2/adapters.py b/services/core/models/src/nmp/core/models/api/v2/adapters.py index 08c3b1d7d2..59a0a921b8 100644 --- a/services/core/models/src/nmp/core/models/api/v2/adapters.py +++ b/services/core/models/src/nmp/core/models/api/v2/adapters.py @@ -10,7 +10,7 @@ from nmp.common.api.common import Page from nmp.common.api.parsed_filter import ParsedFilter, make_filter_dep from nmp.common.api.utils import generate_openapi_extra_params -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.common.service.dependencies import get_sdk_client from nmp.core.models.api.dependencies import get_adapter_entity_service from nmp.core.models.api.permissions import check_fileset_access @@ -165,6 +165,10 @@ async def delete_adapter( deleted = await service.delete_adapter_in_workspace(workspace, name) except ValueError as e: raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e if deleted == -2: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, diff --git a/services/core/models/src/nmp/core/models/api/v2/deployment_configs.py b/services/core/models/src/nmp/core/models/api/v2/deployment_configs.py index c7b12b6fcc..30d716492f 100644 --- a/services/core/models/src/nmp/core/models/api/v2/deployment_configs.py +++ b/services/core/models/src/nmp/core/models/api/v2/deployment_configs.py @@ -9,7 +9,7 @@ from nmp.common.api.parsed_filter import ParsedFilter, make_filter_dep from nmp.common.api.utils import generate_openapi_extra_params from nmp.common.auth import AuthClient, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.dependencies import get_model_deployment_config_service from nmp.core.models.api.permissions import check_model_entity_access from nmp.core.models.api.service.model_deployment_config_service import ( @@ -322,6 +322,10 @@ async def delete_all_deployment_config_versions( raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=e.message) except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception("Unexpected error deleting deployment config") raise HTTPException( @@ -382,6 +386,10 @@ async def delete_deployment_config_version( raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail=e.message) except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception("Unexpected error deleting deployment config version") raise HTTPException( diff --git a/services/core/models/src/nmp/core/models/api/v2/deployments.py b/services/core/models/src/nmp/core/models/api/v2/deployments.py index 5398847e41..a0fddeefd2 100644 --- a/services/core/models/src/nmp/core/models/api/v2/deployments.py +++ b/services/core/models/src/nmp/core/models/api/v2/deployments.py @@ -9,7 +9,7 @@ from nmp.common.api.parsed_filter import ParsedFilter, make_filter_dep from nmp.common.api.utils import generate_openapi_extra_params from nmp.common.auth import AuthClient, AuthContext, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.dependencies import get_model_deployment_service from nmp.core.models.api.permissions import check_deployment_config_access from nmp.core.models.api.service.model_deployment_service import DeploymentStatusConflictError, ModelDeploymentService @@ -444,6 +444,10 @@ async def delete_all_deployment_versions( return None except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception("Unexpected error deleting deployment") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to delete deployment") @@ -519,6 +523,10 @@ async def delete_deployment_version( return None except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception("Unexpected error deleting deployment version") raise HTTPException( diff --git a/services/core/models/src/nmp/core/models/api/v2/models.py b/services/core/models/src/nmp/core/models/api/v2/models.py index 10ff5dcbca..553149a7b0 100644 --- a/services/core/models/src/nmp/core/models/api/v2/models.py +++ b/services/core/models/src/nmp/core/models/api/v2/models.py @@ -25,7 +25,7 @@ from nmp.common.api.parsed_filter import ParsedFilter, make_filter_dep from nmp.common.api.utils import generate_openapi_extra_params from nmp.common.auth import AuthClient, get_auth_client -from nmp.common.entities.client import EntityNotFoundError, EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityNotFoundError, EntityValidationError from nmp.common.sdk_factory import get_async_platform_sdk from nmp.common.service.dependencies import get_sdk_client from nmp.core.models.api.dependencies import get_adapter_entity_service, get_model_entity_service @@ -460,6 +460,10 @@ async def delete_model( except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception("Failed to delete model entity") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to delete model entity") @@ -570,6 +574,10 @@ async def delete_model_adapter( except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception as e: logger.exception(f"Failed to delete model entity - {e}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Failed to delete model entity") diff --git a/services/core/models/src/nmp/core/models/api/v2/prompts.py b/services/core/models/src/nmp/core/models/api/v2/prompts.py index 86dd87cdbc..71c6860da1 100644 --- a/services/core/models/src/nmp/core/models/api/v2/prompts.py +++ b/services/core/models/src/nmp/core/models/api/v2/prompts.py @@ -7,7 +7,7 @@ from nmp.common.api.common import Page from nmp.common.api.parsed_filter import ParsedFilter, make_filter_dep from nmp.common.api.utils import generate_openapi_extra_params -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.dependencies import get_prompt_service from nmp.core.models.api.service.prompt_service import PromptService from nmp.core.models.schemas import ( @@ -189,6 +189,10 @@ async def delete_prompt( return None except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception: logger.exception(f"Failed to delete prompt {_sanitize_for_log(workspace)}/{_sanitize_for_log(name)}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="Internal server error") diff --git a/services/core/models/src/nmp/core/models/api/v2/providers.py b/services/core/models/src/nmp/core/models/api/v2/providers.py index 9043f8a6a1..85208b3a22 100644 --- a/services/core/models/src/nmp/core/models/api/v2/providers.py +++ b/services/core/models/src/nmp/core/models/api/v2/providers.py @@ -9,7 +9,7 @@ from nmp.common.api.parsed_filter import ParsedFilter, make_filter_dep from nmp.common.api.utils import generate_openapi_extra_params from nmp.common.auth import AuthClient, AuthContext, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.common.service.dependencies import get_sdk_client from nmp.core.models.api.dependencies import get_model_provider_service from nmp.core.models.api.permissions import check_deployment_access, check_secret_access @@ -285,6 +285,10 @@ async def delete_provider( return None except HTTPException: raise + except EntityConflictError as e: + raise HTTPException( + status_code=status.HTTP_409_CONFLICT, detail="Concurrent modification - please retry." + ) from e except Exception as e: logger.exception(f"Failed to delete model provider {workspace}/{name}") raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=str(e)) diff --git a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py index 8b8a33cec6..0cae172d84 100644 --- a/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py +++ b/services/core/models/src/nmp/core/models/controllers/backends/deployments_plugin/backend.py @@ -9,7 +9,7 @@ from nemo_deployments_plugin.entities import Deployment, DeploymentConfig, Prerequisite, Volume from nemo_platform import AsyncNeMoPlatform from nemo_platform.resources.entities import AsyncEntitiesResource -from nemo_platform_plugin.entity_client import NemoEntitiesClient, NemoEntityNotFoundError +from nemo_platform_plugin.entity_client import NemoEntitiesClient, NemoEntityConflictError, NemoEntityNotFoundError from nemo_platform_plugin.sdk_provider import get_async_platform_sdk from nmp.common.config import Runtime from nmp.core.models.app.constants import MODEL_MANAGED_BY_LABEL, MODEL_MANAGED_BY_MODELS_CONTROLLER @@ -211,9 +211,16 @@ async def _complete_deployment_delete(self, workspace: str, deployment_name: str # Plugin reconciler gave up on substrate teardown; remove the stale # entity so models delete can finish config/volume cleanup. try: - await self._entity_client().delete(Deployment, name=deployment_name, workspace=workspace) + await self._entity_client().delete( + Deployment, + name=deployment.name, + workspace=workspace, + expected_db_version=deployment.db_version, + ) except NemoEntityNotFoundError: pass + except NemoEntityConflictError: + return False elif deployment.status != "DELETING" or deployment.desired_state != "STOPPED": deployment.status = "DELETING" deployment.desired_state = "STOPPED" diff --git a/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py b/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py index 88b15c3f92..120a720697 100644 --- a/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py +++ b/services/core/models/src/nmp/core/models/controllers/provider_reconciler.py @@ -63,6 +63,15 @@ def _has_backend_format(model_entity: ModelEntity) -> bool: return isinstance(value, str) and bool(value) +def _get_virtual_model_db_version(virtual_model: object) -> int | None: + db_version = getattr(virtual_model, "db_version", None) + if isinstance(db_version, bool): + return None + if isinstance(db_version, int): + return db_version + return None + + # --------------------------------------------------------------------------- # Discovery result types # --------------------------------------------------------------------------- @@ -1102,10 +1111,20 @@ async def _cleanup_orphaned_virtual_models( ): continue + expected_db_version = _get_virtual_model_db_version(virtual_model) + if expected_db_version is None: + logger.warning( + "Skipping orphaned autoprovisioned VirtualModel %s/%s because it has no database version for conditional deletion", + virtual_model.workspace, + virtual_model.name, + ) + continue + try: await self._models_sdk.inference.virtual_models.delete( name=virtual_model.name, workspace=virtual_model.workspace, + expected_db_version=expected_db_version, ) logger.info( "Deleted orphaned autoprovisioned VirtualModel %s/%s", diff --git a/services/core/models/tests/unit/api/test_deployment_configs_api.py b/services/core/models/tests/unit/api/test_deployment_configs_api.py index 03dc9a105a..dcb0366721 100644 --- a/services/core/models/tests/unit/api/test_deployment_configs_api.py +++ b/services/core/models/tests/unit/api/test_deployment_configs_api.py @@ -11,7 +11,7 @@ from fastapi.testclient import TestClient from nmp.common.api.common import Page, PaginationData from nmp.common.auth import AuthClient, Principal, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.service.model_deployment_config_service import ModelDeploymentConfigService from nmp.core.models.api.v2.deployment_configs import router from nmp.core.models.api.v2.utils import ERR_DEPLOYMENTS_NOT_ENABLED as _DEPLOYMENTS_NOT_ENABLED @@ -386,3 +386,27 @@ def test_delete_deployment_config_version_when_deployments_disabled_returns_422( assert response.status_code == 422 assert response.json()["detail"] == _DEPLOYMENTS_NOT_ENABLED assert not mock_deployment_config_service.delete_deployment_config.called + + +@patch("nmp.core.models.api.v2.deployment_configs.deployments_enabled", return_value=True) +def test_delete_deployment_config_conflict_returns_409( + _mock_deployments_enabled, client, mock_deployment_config_service +): + """Stale deployment config deletes return 409.""" + mock_deployment_config_service.delete_deployment_config.side_effect = EntityConflictError("stale version") + + response = client.delete("/v2/workspaces/default/deployment-configs/cfg") + + assert response.status_code == 409 + + +@patch("nmp.core.models.api.v2.deployment_configs.deployments_enabled", return_value=True) +def test_delete_deployment_config_version_conflict_returns_409( + _mock_deployments_enabled, client, mock_deployment_config_service +): + """Stale deployment config version deletes return 409.""" + mock_deployment_config_service.delete_deployment_config.side_effect = EntityConflictError("stale version") + + response = client.delete("/v2/workspaces/default/deployment-configs/cfg/versions/1") + + assert response.status_code == 409 diff --git a/services/core/models/tests/unit/api/test_deployments_api.py b/services/core/models/tests/unit/api/test_deployments_api.py index 4e1d2c34c0..3ba1edfa76 100644 --- a/services/core/models/tests/unit/api/test_deployments_api.py +++ b/services/core/models/tests/unit/api/test_deployments_api.py @@ -11,7 +11,7 @@ from fastapi.testclient import TestClient from nmp.common.api.common import Page, PaginationData from nmp.common.auth import AuthClient, Principal, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.service.model_deployment_service import DeploymentStatusConflictError, ModelDeploymentService from nmp.core.models.api.v2.deployments import router from nmp.core.models.api.v2.utils import ERR_DEPLOYMENTS_NOT_ENABLED as _DEPLOYMENTS_NOT_ENABLED @@ -450,3 +450,23 @@ def test_delete_deployment_version_when_deployments_disabled_returns_422( assert response.status_code == 422 assert response.json()["detail"] == _DEPLOYMENTS_NOT_ENABLED assert not mock_deployment_service.delete_deployment.called + + +@patch("nmp.core.models.api.v2.deployments.deployments_enabled", return_value=True) +def test_delete_deployment_conflict_returns_409(_mock_deployments_enabled, client, mock_deployment_service): + """Stale deployment deletes return 409.""" + mock_deployment_service.delete_deployment.side_effect = EntityConflictError("stale version") + + response = client.delete("/v2/workspaces/default/deployments/d") + + assert response.status_code == 409 + + +@patch("nmp.core.models.api.v2.deployments.deployments_enabled", return_value=True) +def test_delete_deployment_version_conflict_returns_409(_mock_deployments_enabled, client, mock_deployment_service): + """Stale deployment version deletes return 409.""" + mock_deployment_service.delete_deployment.side_effect = EntityConflictError("stale version") + + response = client.delete("/v2/workspaces/default/deployments/d/versions/1") + + assert response.status_code == 409 diff --git a/services/core/models/tests/unit/api/test_models_api.py b/services/core/models/tests/unit/api/test_models_api.py index 3b0a9c2a03..83997c3ced 100644 --- a/services/core/models/tests/unit/api/test_models_api.py +++ b/services/core/models/tests/unit/api/test_models_api.py @@ -13,7 +13,7 @@ from nemo_platform_plugin.client.errors import NemoTransportError from nmp.common.api.common import Page, PaginationData from nmp.common.auth import AuthClient, Principal, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.service.adapter_entity_service import AdapterEntityService from nmp.core.models.api.service.model_entity_service import ModelEntityService from nmp.core.models.api.v2.models import router, start_update_model_spec_job @@ -38,6 +38,7 @@ def mock_model_entity_service(): def mock_adapter_entity_service(): service = Mock(spec=AdapterEntityService) service.create_adapter = AsyncMock() + service.delete_adapter = AsyncMock() service.update_adapter = AsyncMock() return service @@ -474,6 +475,24 @@ def test_update_model_with_verbose_true(client, mock_model_entity_service, sampl assert call_args.kwargs["verbose"] is True +def test_delete_model_conflict_returns_409(client, mock_model_entity_service): + """Test stale model deletes return 409.""" + mock_model_entity_service.delete_model_entity.side_effect = EntityConflictError("stale version") + + response = client.delete("/apis/models/v2/workspaces/default/models/my-model") + + assert response.status_code == 409 + + +def test_delete_model_adapter_conflict_returns_409(client, mock_adapter_entity_service): + """Test stale model adapter deletes return 409.""" + mock_adapter_entity_service.delete_adapter.side_effect = EntityConflictError("stale version") + + response = client.delete("/apis/models/v2/workspaces/default/models/my-model/adapters/lora") + + assert response.status_code == 409 + + def test_create_model_adapter_entity_validation_error_returns_422(client, mock_adapter_entity_service): """Test that entity store validation errors during adapter creation return 422.""" mock_adapter_entity_service.create_adapter.side_effect = EntityValidationError("adapter name invalid") diff --git a/services/core/models/tests/unit/api/test_prompts_api.py b/services/core/models/tests/unit/api/test_prompts_api.py index da8f30cd8a..72fa690477 100644 --- a/services/core/models/tests/unit/api/test_prompts_api.py +++ b/services/core/models/tests/unit/api/test_prompts_api.py @@ -10,7 +10,7 @@ from fastapi import FastAPI from fastapi.testclient import TestClient from nmp.common.api.common import Page, PaginationData -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.service.prompt_service import PromptService from nmp.core.models.api.v2.prompts import router from nmp.core.models.schemas import Prompt, PromptMessage, PromptMessageRole @@ -261,3 +261,11 @@ def test_delete_prompt_not_found_returns_404(client, mock_prompt_service): response = client.delete("/apis/models/v2/workspaces/default/prompts/missing") assert response.status_code == 404 + + +def test_delete_prompt_conflict_returns_409(client, mock_prompt_service): + mock_prompt_service.delete_prompt.side_effect = EntityConflictError("stale version") + + response = client.delete("/apis/models/v2/workspaces/default/prompts/summarizer") + + assert response.status_code == 409 diff --git a/services/core/models/tests/unit/api/test_providers_api.py b/services/core/models/tests/unit/api/test_providers_api.py index fe882ad07a..95cd444684 100644 --- a/services/core/models/tests/unit/api/test_providers_api.py +++ b/services/core/models/tests/unit/api/test_providers_api.py @@ -11,7 +11,7 @@ from fastapi.testclient import TestClient from nmp.common.api.common import Page, PaginationData from nmp.common.auth import AuthClient, Principal, get_auth_client -from nmp.common.entities.client import EntityValidationError +from nmp.common.entities.client import EntityConflictError, EntityValidationError from nmp.core.models.api.service.model_provider_service import ModelProviderService, ModelProviderValidationError from nmp.core.models.api.v2.providers import router from nmp.core.models.schemas import ModelProvider, ModelProviderStatus @@ -503,3 +503,12 @@ def test_update_provider_status_entity_validation_error_returns_422(client, mock assert response.status_code == 422 assert "served_models invalid" in response.json()["detail"] + + +def test_delete_provider_conflict_returns_409(client, mock_model_provider_service): + """Test stale provider deletes return 409.""" + mock_model_provider_service.delete_model_provider.side_effect = EntityConflictError("stale version") + + response = client.delete("/apis/models/v2/workspaces/default/providers/test-provider") + + assert response.status_code == 409 diff --git a/services/core/models/tests/unit/controllers/backends/deployments_plugin/test_backend.py b/services/core/models/tests/unit/controllers/backends/deployments_plugin/test_backend.py index cb45f7fad2..ea8009a21a 100644 --- a/services/core/models/tests/unit/controllers/backends/deployments_plugin/test_backend.py +++ b/services/core/models/tests/unit/controllers/backends/deployments_plugin/test_backend.py @@ -394,7 +394,12 @@ async def test_delete_completes_when_plugin_deployment_failed() -> None: result = await backend.delete_model_deployment("default", "my-dep") assert result.status == "DELETED" - backend._entities.delete.assert_any_await(Deployment, name="my-dep-server", workspace="default") + backend._entities.delete.assert_any_await( + Deployment, + name="my-dep-server", + workspace="default", + expected_db_version=server.db_version, + ) @pytest.mark.asyncio diff --git a/services/core/models/tests/unit/controllers/test_provider_reconciler.py b/services/core/models/tests/unit/controllers/test_provider_reconciler.py index 8742af39ba..eaf867846e 100644 --- a/services/core/models/tests/unit/controllers/test_provider_reconciler.py +++ b/services/core/models/tests/unit/controllers/test_provider_reconciler.py @@ -1994,12 +1994,15 @@ def _virtual_model( workspace: str = "ws", default_model_entity: str | None = None, autoprovisioned: bool = True, + db_version: int | None = 1, ): vm = MagicMock() vm.name = name vm.workspace = workspace vm.default_model_entity = default_model_entity vm.autoprovisioned = autoprovisioned + if db_version is not None: + vm.db_version = db_version return vm @@ -2045,7 +2048,11 @@ async def test_reconcile_with_no_providers_deletes_orphaned_autoprovisioned_virt await reconciler.reconcile_model_providers([]) mock_models_sdk.inference.virtual_models.list.assert_called_once_with(workspace="-", page_size=200) - mock_models_sdk.inference.virtual_models.delete.assert_awaited_once_with(name="model-a", workspace="ws") + mock_models_sdk.inference.virtual_models.delete.assert_awaited_once_with( + name="model-a", + workspace="ws", + expected_db_version=1, + ) @pytest.mark.asyncio @@ -2119,7 +2126,11 @@ async def test_cleanup_lost_provider_does_not_protect_autoprovisioned_virtual_mo vm_snapshot, _ = await reconciler._load_virtual_models() await reconciler._cleanup_orphaned_virtual_models([ctx], vm_snapshot) - mock_models_sdk.inference.virtual_models.delete.assert_awaited_once_with(name="model-a", workspace="ws") + mock_models_sdk.inference.virtual_models.delete.assert_awaited_once_with( + name="model-a", + workspace="ws", + expected_db_version=1, + ) @pytest.mark.asyncio @@ -2163,10 +2174,41 @@ async def test_cleanup_delete_failure_is_logged_and_non_fatal(reconciler, mock_m vm_snapshot, _ = await reconciler._load_virtual_models() await reconciler._cleanup_orphaned_virtual_models([], vm_snapshot) - mock_models_sdk.inference.virtual_models.delete.assert_awaited_once_with(name="model-a", workspace="ws") + mock_models_sdk.inference.virtual_models.delete.assert_awaited_once_with( + name="model-a", + workspace="ws", + expected_db_version=1, + ) assert any("Failed to delete orphaned autoprovisioned VirtualModel ws/model-a" in r.message for r in caplog.records) +@pytest.mark.asyncio +async def test_cleanup_skips_orphaned_virtual_model_without_db_version(reconciler, mock_models_sdk, caplog): + """Cleanup must not fall back to an unconditional delete when the listed VM has no version.""" + mock_models_sdk.inference.virtual_models.list = MagicMock( + return_value=_AsyncPaginator( + [ + _virtual_model( + "model-a", + default_model_entity="ws/model-a", + autoprovisioned=True, + db_version=None, + ) + ] + ) + ) + + with caplog.at_level(logging.WARNING): + vm_snapshot, _ = await reconciler._load_virtual_models() + await reconciler._cleanup_orphaned_virtual_models([], vm_snapshot) + + mock_models_sdk.inference.virtual_models.delete.assert_not_awaited() + assert any( + "Skipping orphaned autoprovisioned VirtualModel ws/model-a because it has no database version" in r.message + for r in caplog.records + ) + + # ============================================================================= # Deployment-backed provider tests (no autocreate, served_models from id/root/parent) # ============================================================================= diff --git a/services/core/models/tests/unit/test_model_deployment_config_service_unit.py b/services/core/models/tests/unit/test_model_deployment_config_service_unit.py index 0e958a9df1..081379a952 100644 --- a/services/core/models/tests/unit/test_model_deployment_config_service_unit.py +++ b/services/core/models/tests/unit/test_model_deployment_config_service_unit.py @@ -534,7 +534,9 @@ async def test_delete_deployment_config_success(deployment_config_service, mock_ # Assert assert result is True mock_entity_client.delete.assert_called_once_with( - ModelDeploymentConfigEntity, sample_config_entity.name, workspace="default" + ModelDeploymentConfigEntity, + sample_config_entity.name, + workspace="default", ) @@ -556,7 +558,9 @@ async def test_delete_deployment_config_specific_version( # Assert assert result is True mock_entity_client.delete.assert_called_once_with( - ModelDeploymentConfigEntity, sample_config_entity.name, workspace="default" + ModelDeploymentConfigEntity, + sample_config_entity.name, + workspace="default", ) diff --git a/services/core/models/tests/unit/test_model_deployment_service_unit.py b/services/core/models/tests/unit/test_model_deployment_service_unit.py index 49b44ad554..aef88ff491 100644 --- a/services/core/models/tests/unit/test_model_deployment_service_unit.py +++ b/services/core/models/tests/unit/test_model_deployment_service_unit.py @@ -597,7 +597,12 @@ async def test_delete_deployment_already_deleted_hard_deletes(deployment_service # Assert assert result is None # Returns None for hard delete - mock_entity_client.delete.assert_called_once_with(ModelDeploymentEntity, deleted_entity.name, workspace="default") + mock_entity_client.delete.assert_called_once_with( + ModelDeploymentEntity, + deleted_entity.name, + workspace="default", + expected_db_version=deleted_entity.db_version, + ) @pytest.mark.asyncio diff --git a/services/core/models/tests/unit/test_model_entity_service_unit.py b/services/core/models/tests/unit/test_model_entity_service_unit.py index 405b8099d8..96593a0d8f 100644 --- a/services/core/models/tests/unit/test_model_entity_service_unit.py +++ b/services/core/models/tests/unit/test_model_entity_service_unit.py @@ -571,7 +571,6 @@ async def test_update_model_entity_non_verbose_keeps_adapters(model_entity_servi async def test_delete_model_entity_success(model_entity_service, mock_entity_client, sample_model): """Test successful model entity deletion.""" # Arrange - mock_entity_client.get.return_value = sample_model mock_entity_client.delete.return_value = None # Act @@ -579,23 +578,25 @@ async def test_delete_model_entity_success(model_entity_service, mock_entity_cli # Assert assert result is True - mock_entity_client.get.assert_called_once_with(Model, workspace="default", name="test-model") - mock_entity_client.delete.assert_called_once_with(Model, sample_model.name, workspace="default") + mock_entity_client.delete.assert_called_once_with( + Model, + "test-model", + workspace="default", + ) @pytest.mark.asyncio async def test_delete_model_entity_not_found(model_entity_service, mock_entity_client): """Test deleting a non-existent model entity.""" # Arrange - mock_entity_client.get.side_effect = EntityNotFoundError("Entity not found") + mock_entity_client.delete.side_effect = EntityNotFoundError("Entity not found") # Act result = await model_entity_service.delete_model_entity("default", "nonexistent") # Assert assert result is False - mock_entity_client.get.assert_called_once_with(Model, workspace="default", name="nonexistent") - mock_entity_client.delete.assert_not_called() + mock_entity_client.delete.assert_called_once_with(Model, "nonexistent", workspace="default") @pytest.mark.asyncio @@ -1095,7 +1096,12 @@ async def test_delete_model_adapter_success( result = await adapter_entity_service.delete_adapter("default", "test-model", "lora-adapter") assert result == 0 - mock_entity_client.delete_by_id.assert_called_once() + mock_entity_client.delete.assert_called_once_with( + Adapter, + adapter.name, + workspace="default", + parent=adapter.parent, + ) @pytest.mark.asyncio diff --git a/services/core/models/tests/unit/test_model_provider_service_unit.py b/services/core/models/tests/unit/test_model_provider_service_unit.py index 30aaf47038..f7c18e458b 100644 --- a/services/core/models/tests/unit/test_model_provider_service_unit.py +++ b/services/core/models/tests/unit/test_model_provider_service_unit.py @@ -8,7 +8,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest -from nmp.common.entities.client import EntityClient, EntityNotFoundError +from nmp.common.entities.client import EntityClient, EntityConflictError, EntityNotFoundError from nmp.core.models.api.service.model_provider_service import ( ModelProviderService, ModelProviderValidationError, @@ -277,7 +277,10 @@ async def test_delete_model_provider_success(model_provider_service, mock_entity # Assert assert result is True mock_entity_client.delete.assert_called_once_with( - ModelProviderEntity, sample_provider_entity.name, workspace="default" + ModelProviderEntity, + sample_provider_entity.name, + workspace="default", + expected_db_version=sample_provider_entity.db_version, ) @@ -569,7 +572,41 @@ async def test_delete_provider_cleans_up_model_entity_references(model_provider_ mock_entity_client.update.assert_called_once() updated_model = mock_entity_client.update.call_args[0][0] assert updated_model.model_providers == ["ws/other-provider"] - mock_entity_client.delete.assert_called_once_with(ModelProviderEntity, "provider-to-delete", workspace="ws") + mock_entity_client.delete.assert_called_once_with( + ModelProviderEntity, + "provider-to-delete", + workspace="ws", + expected_db_version=provider_entity.db_version, + ) + + +@pytest.mark.asyncio +async def test_delete_provider_conflict_leaves_model_entity_references(model_provider_service, mock_entity_client): + """A stale provider version must not clean up linked model references.""" + model_entity = _create_model_entity( + name="my-model", + workspace="ws", + model_providers=["ws/provider-to-delete", "ws/other-provider"], + ) + provider_entity = create_provider_entity( + name="provider-to-delete", + workspace="ws", + host_url="https://api.example.com/v1", + served_models=[ + ServedModelMapping(model_entity_id="ws/my-model", served_model_name="my-model"), + ], + status=ModelProviderStatus.READY, + ) + + mock_entity_client.get.side_effect = _make_entity_get_dispatcher(provider_entity, {"my-model": model_entity}) + mock_entity_client.delete.side_effect = EntityConflictError("stale version") + + request = DeleteModelProviderRequest(workspace="ws", name="provider-to-delete") + with pytest.raises(EntityConflictError): + await model_provider_service.delete_model_provider(request) + + mock_entity_client.update.assert_not_called() + assert model_entity.model_providers == ["ws/provider-to-delete", "ws/other-provider"] @pytest.mark.asyncio @@ -638,7 +675,12 @@ async def test_delete_provider_cleanup_provider_not_in_model_providers(model_pro assert result is True mock_entity_client.update.assert_not_called() - mock_entity_client.delete.assert_called_once_with(ModelProviderEntity, "my-provider", workspace="ws") + mock_entity_client.delete.assert_called_once_with( + ModelProviderEntity, + "my-provider", + workspace="ws", + expected_db_version=provider_entity.db_version, + ) @pytest.mark.asyncio @@ -662,7 +704,12 @@ async def test_delete_provider_cleanup_malformed_model_entity_id(model_provider_ assert result is True mock_entity_client.update.assert_not_called() - mock_entity_client.delete.assert_called_once_with(ModelProviderEntity, "my-provider", workspace="ws") + mock_entity_client.delete.assert_called_once_with( + ModelProviderEntity, + "my-provider", + workspace="ws", + expected_db_version=provider_entity.db_version, + ) @pytest.mark.asyncio @@ -695,7 +742,12 @@ async def _failing_model_get(entity_type, **kwargs): assert result is True assert call_count == 1 - mock_entity_client.delete.assert_called_once_with(ModelProviderEntity, "my-provider", workspace="ws") + mock_entity_client.delete.assert_called_once_with( + ModelProviderEntity, + "my-provider", + workspace="ws", + expected_db_version=provider_entity.db_version, + ) @pytest.mark.asyncio diff --git a/services/core/models/tests/unit/test_prompt_service_unit.py b/services/core/models/tests/unit/test_prompt_service_unit.py index 37b6415d31..a72f41066e 100644 --- a/services/core/models/tests/unit/test_prompt_service_unit.py +++ b/services/core/models/tests/unit/test_prompt_service_unit.py @@ -233,15 +233,19 @@ async def test_delete_prompt_success(prompt_service, mock_entity_client, sample_ result = await prompt_service.delete_prompt(DeletePromptRequest(workspace="default", name="summarizer")) assert result is True - mock_entity_client.delete.assert_called_once() + mock_entity_client.delete.assert_called_once_with( + PromptEntity, + "summarizer", + workspace="default", + ) @pytest.mark.asyncio async def test_delete_prompt_not_found(prompt_service, mock_entity_client): - """Test that deleting a missing prompt returns False and does not call delete.""" - mock_entity_client.get.side_effect = EntityNotFoundError("not found") + """Test that deleting a missing prompt returns False.""" + mock_entity_client.delete.side_effect = EntityNotFoundError("not found") result = await prompt_service.delete_prompt(DeletePromptRequest(workspace="default", name="missing")) assert result is False - mock_entity_client.delete.assert_not_called() + mock_entity_client.delete.assert_called_once_with(PromptEntity, "missing", workspace="default") diff --git a/services/guardrails/src/nmp/guardrails/api/v2/configs/endpoints.py b/services/guardrails/src/nmp/guardrails/api/v2/configs/endpoints.py index 53ccc44dc7..00aadfd6dd 100644 --- a/services/guardrails/src/nmp/guardrails/api/v2/configs/endpoints.py +++ b/services/guardrails/src/nmp/guardrails/api/v2/configs/endpoints.py @@ -197,7 +197,12 @@ async def delete_config( except EntityNotFoundError: raise HTTPException(status_code=404, detail="Guardrail config not found.") - await entities_client.delete(GuardrailConfig, config.name, workspace=workspace) + try: + await entities_client.delete(GuardrailConfig, config.name, workspace=workspace) + except EntityNotFoundError: + raise HTTPException(status_code=404, detail="Guardrail config not found.") + except EntityConflictError as exc: + raise HTTPException(status_code=409, detail="Concurrent modification - please retry.") from exc full_config_id = f"{workspace}/{config.name}" if workspace else config.name diff --git a/services/guardrails/tests/apis/test_configs_api.py b/services/guardrails/tests/apis/test_configs_api.py index 727feea205..f08b11c2f8 100644 --- a/services/guardrails/tests/apis/test_configs_api.py +++ b/services/guardrails/tests/apis/test_configs_api.py @@ -3,8 +3,13 @@ """Tests for the Guardrails config API endpoints.""" +from unittest.mock import AsyncMock + import pytest from fastapi.testclient import TestClient +from nmp.common.entities import EntityNotFoundError +from nmp.common.service.dependencies import get_entity_client +from nmp.guardrails.entities import GuardrailConfig class TestGuardrailConfigsAPI: @@ -92,6 +97,21 @@ def test_delete_config(self, client: TestClient): assert response.status_code == 200 assert response.json()["message"] == "Resource deleted successfully." + def test_delete_config_not_found_during_delete(self, client: TestClient): + """Test a config deleted after lookup returns 404.""" + entities_client = AsyncMock() + entities_client.get.return_value = GuardrailConfig(name="delete-race-config", workspace="default") + entities_client.delete.side_effect = EntityNotFoundError("not found") + dependency_overrides = getattr(client.app, "dependency_overrides") + dependency_overrides[get_entity_client] = lambda: entities_client + try: + response = client.delete("/apis/guardrails/v2/workspaces/default/configs/delete-race-config") + finally: + dependency_overrides.pop(get_entity_client, None) + + assert response.status_code == 404 + assert response.json()["detail"] == "Guardrail config not found." + def test_get_config_not_found(self, client: TestClient): """Test getting a non-existent config returns 404.""" response = client.get("/apis/guardrails/v2/workspaces/default/configs/nonexistent-config") diff --git a/services/hello-world/src/nmp/hello_world/api/v1/messages/endpoints.py b/services/hello-world/src/nmp/hello_world/api/v1/messages/endpoints.py index 7899ce83cf..c205c8ec82 100644 --- a/services/hello-world/src/nmp/hello_world/api/v1/messages/endpoints.py +++ b/services/hello-world/src/nmp/hello_world/api/v1/messages/endpoints.py @@ -6,7 +6,7 @@ from typing import List from fastapi import APIRouter, Depends, HTTPException -from nmp.common.entities.client import EntityClient, EntityConflictError +from nmp.common.entities.client import EntityClient, EntityConflictError, EntityNotFoundError from nmp.common.service.dependencies import get_entity_client from nmp.hello_world.api.v1.messages.schemas import CreateHelloWorldMessageRequest, UpdateHelloWorldMessageRequest from nmp.hello_world.entities import HelloWorldMessage @@ -107,18 +107,14 @@ async def delete_message( entity_store: EntityClient = Depends(get_entity_client), ) -> None: """Delete a HelloWorld message.""" - # Get the message to find its ID try: - existing = await entity_store.get(HelloWorldMessage, name=name, workspace=workspace) - except Exception as e: - if "not found" in str(e).lower() or "404" in str(e): - raise HTTPException( - status_code=404, - detail=f"Message '{name}' not found in workspace '{workspace}'", - ) from e - raise HTTPException(status_code=500, detail=str(e)) from e - - try: - await entity_store.delete(HelloWorldMessage, existing.name, workspace=workspace) + await entity_store.delete(HelloWorldMessage, name, workspace=workspace) + except EntityConflictError as e: + raise HTTPException(status_code=409, detail="Concurrent modification - please retry.") from e + except EntityNotFoundError as e: + raise HTTPException( + status_code=404, + detail=f"Message '{name}' not found in workspace '{workspace}'", + ) from e except Exception as e: raise HTTPException(status_code=500, detail=str(e)) from e diff --git a/services/hello-world/tests/unit/test_messages_endpoints.py b/services/hello-world/tests/unit/test_messages_endpoints.py new file mode 100644 index 0000000000..8772b310eb --- /dev/null +++ b/services/hello-world/tests/unit/test_messages_endpoints.py @@ -0,0 +1,37 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Unit tests for HelloWorld message endpoints.""" + +import pytest +from fastapi import HTTPException +from nemo_platform.types import DeleteResponse +from nmp.common.entities.client import EntityClient, EntityNotFoundError +from nmp.hello_world.api.v1.messages.endpoints import delete_message + + +class DeleteNotFoundEntityClient(EntityClient): + def __init__(self) -> None: + pass + + async def delete( + self, + entity_type: object, + name: str, + *, + workspace: str | None = None, + parent: str | None = None, + expected_db_version: int | None = None, + ) -> DeleteResponse: + raise EntityNotFoundError("not found") + + +@pytest.mark.asyncio +async def test_delete_message_maps_delete_not_found_to_404(): + entity_store = DeleteNotFoundEntityClient() + + with pytest.raises(HTTPException) as exc_info: + await delete_message(workspace="default", name="missing-message", entity_store=entity_store) + + assert exc_info.value.status_code == 404 + assert exc_info.value.detail == "Message 'missing-message' not found in workspace 'default'" diff --git a/web/packages/studio/src/mocks/handlers/guardrails.ts b/web/packages/studio/src/mocks/handlers/guardrails.ts index c1ffe47607..6f81e0e504 100644 --- a/web/packages/studio/src/mocks/handlers/guardrails.ts +++ b/web/packages/studio/src/mocks/handlers/guardrails.ts @@ -10,6 +10,7 @@ export const mockGuardrailConfigs: GuardrailConfig[] = [ id: 'cfg-1', entity_id: 'cfg-1', parent: 'ws-default', + db_version: 1, name: 'pii-filter', workspace: 'default', description: 'Blocks PII in user inputs and outputs', @@ -32,6 +33,7 @@ export const mockGuardrailConfigs: GuardrailConfig[] = [ id: 'cfg-2', entity_id: 'cfg-2', parent: 'ws-default', + db_version: 1, name: 'toxicity-guard', workspace: 'default', description: 'Detects and blocks toxic language', diff --git a/web/packages/studio/src/routes/VirtualModelsListRoute/VirtualModelDetailsSidePanel/index.test.tsx b/web/packages/studio/src/routes/VirtualModelsListRoute/VirtualModelDetailsSidePanel/index.test.tsx index f34a652d8e..0c80b27e7b 100644 --- a/web/packages/studio/src/routes/VirtualModelsListRoute/VirtualModelDetailsSidePanel/index.test.tsx +++ b/web/packages/studio/src/routes/VirtualModelsListRoute/VirtualModelDetailsSidePanel/index.test.tsx @@ -24,6 +24,7 @@ const vm: VirtualModel = { updated_by: null, entity_id: 'default/my-vm', parent: '', + db_version: 1, }; describe('VirtualModelDetailsSidePanel', () => {