Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
80 changes: 80 additions & 0 deletions src/mistralai/extra/tests/test_workflow_encoding.py
Original file line number Diff line number Diff line change
Expand Up @@ -760,3 +760,83 @@ async def test_workflow_encoding_hook_handles_gzipped_response():
assert isinstance(result, httpx.Response)
response_body = json.loads(result.content)
assert response_body["result"] == original_data


@pytest.mark.asyncio
async def test_encode_payload_content_does_not_offload_below_threshold(monkeypatch):
storage = InMemoryBlobStorage()
monkeypatch.setattr(
"mistralai.extra.workflows.encoding.payload_encoder.get_blob_storage",
lambda _: storage,
)
config = WorkflowEncodingConfig(
payload_offloading=PayloadOffloadingConfig(
min_size_bytes=1024,
storage_config=BlobStorageConfig(
storage_provider=StorageProvider.S3,
bucket_name="test-bucket",
),
),
)
encoder = PayloadEncoder(encoding_config=config)
context = WorkflowContext(namespace="test", execution_id="exec")

data, options = await encoder.encode_payload_content(b"tiny", context)

assert options == []
assert storage.blobs == {}
assert data == b"tiny"


@pytest.mark.asyncio
async def test_encode_payload_content_force_offload_bypasses_threshold(monkeypatch):
storage = InMemoryBlobStorage()
monkeypatch.setattr(
"mistralai.extra.workflows.encoding.payload_encoder.get_blob_storage",
lambda _: storage,
)
config = WorkflowEncodingConfig(
payload_offloading=PayloadOffloadingConfig(
min_size_bytes=1024,
storage_config=BlobStorageConfig(
storage_provider=StorageProvider.S3,
bucket_name="test-bucket",
),
),
)
encoder = PayloadEncoder(encoding_config=config)
context = WorkflowContext(namespace="test", execution_id="exec")

data, options = await encoder.encode_payload_content(
b"tiny", context, force_offload=True
)

assert EncodedPayloadOptions.OFFLOADED in options
assert len(storage.blobs) == 1
ref = json.loads(data)
assert storage.blobs[ref["key"]] == b"tiny"


@pytest.mark.asyncio
async def test_encode_payload_content_force_offload_is_idempotent(monkeypatch):
storage = InMemoryBlobStorage()
monkeypatch.setattr(
"mistralai.extra.workflows.encoding.payload_encoder.get_blob_storage",
lambda _: storage,
)
config = WorkflowEncodingConfig(
payload_offloading=PayloadOffloadingConfig(
min_size_bytes=1024,
storage_config=BlobStorageConfig(
storage_provider=StorageProvider.S3,
bucket_name="test-bucket",
),
),
)
encoder = PayloadEncoder(encoding_config=config)
context = WorkflowContext(namespace="test", execution_id="exec")

await encoder.encode_payload_content(b"same", context, force_offload=True)
await encoder.encode_payload_content(b"same", context, force_offload=True)

assert len(storage.blobs) == 1
12 changes: 9 additions & 3 deletions src/mistralai/extra/workflows/encoding/payload_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,7 +213,10 @@ def _decrypt(self, data: bytes) -> bytes:
) from main_exc

async def _handle_offloading(
self, data: bytes, context: Optional[WorkflowContext]
self,
data: bytes,
context: Optional[WorkflowContext],
force: bool = False,
) -> tuple[bytes, bool]:
if (
self.offloading_config is None
Expand All @@ -223,7 +226,7 @@ async def _handle_offloading(
"You must configure payload offloading storage"
)

if len(data) < self.offloading_config.min_size_bytes:
if not force and len(data) < self.offloading_config.min_size_bytes:
return data, False

if not context:
Expand Down Expand Up @@ -325,6 +328,7 @@ async def encode_payload_content(
context: Optional[WorkflowContext] = None,
*,
allow_offloading: bool = True,
force_offload: bool = False,
) -> tuple[bytes, list[EncodedPayloadOptions]]:
"""Handle payload encoding.

Expand Down Expand Up @@ -353,7 +357,9 @@ async def encode_payload_content(
encoding_options.append(EncodedPayloadOptions.COMPRESSED)

if allow_offloading and self.offloading_config is not None:
data, offloaded = await self._handle_offloading(data, context)
data, offloaded = await self._handle_offloading(
data, context, force=force_offload
)
if offloaded:
encoding_options.append(EncodedPayloadOptions.OFFLOADED)

Expand Down
Loading