diff --git a/mock_tests/test_batch.py b/mock_tests/test_batch.py index c00bd08db..f8b114534 100644 --- a/mock_tests/test_batch.py +++ b/mock_tests/test_batch.py @@ -60,3 +60,44 @@ def test_ssb_canceled_stream( for i in range(HOW_MANY): batch.add_object({"name": f"Object {i}"}) assert len(service.uuids) == HOW_MANY + + +class MockFailedObjectWeaviateService(weaviate_pb2_grpc.WeaviateServicer): + def BatchStream( + self, + request_iterator: Generator[batch_pb2.BatchStreamRequest, None, None], + context: grpc.ServicerContext, + ) -> Generator[batch_pb2.BatchStreamReply, None, None]: + yield batch_pb2.BatchStreamReply(started=batch_pb2.BatchStreamReply.Started()) + for request in request_iterator: + if request.HasField("data"): + uuids = [obj.uuid for obj in request.data.objects.values] + yield batch_pb2.BatchStreamReply( + results=batch_pb2.BatchStreamReply.Results( + errors=[ + batch_pb2.BatchStreamReply.Results.Error( + uuid=uuid, error="mock failure" + ) + for uuid in uuids + ] + ) + ) + if request.HasField("stop"): + return + + +@pytest.fixture(scope="function") +def failed_object_stream( + canceled_stream_client: weaviate.WeaviateClient, start_grpc_server: grpc.Server +): + service = MockFailedObjectWeaviateService() + weaviate_pb2_grpc.add_WeaviateServicer_to_server(service, start_grpc_server) + return canceled_stream_client.collections.use(mock_class["class"]) + + +def test_ingest_has_errors_on_failed_object( + failed_object_stream: weaviate.collections.Collection, +): + result = failed_object_stream.data.ingest([{"name": "Object 1"}]) + assert result.has_errors is True + assert len(result.errors) == 1 diff --git a/weaviate/collections/batch/async_.py b/weaviate/collections/batch/async_.py index 8b997586c..2e7aa9279 100644 --- a/weaviate/collections/batch/async_.py +++ b/weaviate/collections/batch/async_.py @@ -407,6 +407,7 @@ async def __recv(self) -> None: result_objs += BatchObjectReturn( _all_responses=[err], errors={cached.index: err}, + has_errors=True, ) failed_objs.append(err) logger.warning( @@ -428,6 +429,7 @@ async def __recv(self) -> None: ) result_refs += BatchReferenceReturn( errors={cached.index: err}, + has_errors=True, ) failed_refs.append(err) logger.warning( diff --git a/weaviate/collections/batch/sync.py b/weaviate/collections/batch/sync.py index 6627f8911..cff631e0e 100644 --- a/weaviate/collections/batch/sync.py +++ b/weaviate/collections/batch/sync.py @@ -365,6 +365,7 @@ def __recv(self) -> None: result_objs += BatchObjectReturn( _all_responses=[err], errors={cached.index: err}, + has_errors=True, ) failed_objs.append(err) logger.warning( @@ -387,6 +388,7 @@ def __recv(self) -> None: failed_refs.append(err) result_refs += BatchReferenceReturn( errors={cached.index: err}, + has_errors=True, ) logger.warning( {