290 lines
10 KiB
Python
290 lines
10 KiB
Python
from __future__ import annotations
|
|
|
|
import io
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
|
|
from arr_processing.errors import ProcessingTransportError
|
|
from arr_processing.remote_files import PrefixRemoteFilePort
|
|
from arr_storage.aliyun_oss_v2 import AliyunOssConfig, AliyunOssV2Client
|
|
from arr_storage.contracts import BackendObject, ObjectKeyPolicy
|
|
from arr_storage.exchange import (
|
|
OutputExchangeConfig,
|
|
OutputExchangePublisher,
|
|
)
|
|
from arr_storage.remote import CloudClientError, CloudObjectBackend
|
|
from arr_storage.store import ManagedObjectStore
|
|
|
|
|
|
class Request:
|
|
def __init__(self, **kwargs):
|
|
self.__dict__.update(kwargs)
|
|
|
|
|
|
class FakeSdk:
|
|
GetBucketInfoRequest = Request
|
|
PutObjectRequest = Request
|
|
HeadObjectRequest = Request
|
|
GetObjectRequest = Request
|
|
CopyObjectRequest = Request
|
|
DeleteObjectRequest = Request
|
|
|
|
|
|
class ServiceFailure(Exception):
|
|
def __init__(self, status_code: int, code: str):
|
|
self.status_code = status_code
|
|
self.code = code
|
|
|
|
|
|
class OperationFailure(Exception):
|
|
def __init__(self, service: ServiceFailure):
|
|
self.service = service
|
|
|
|
def unwrap(self):
|
|
return self.service
|
|
|
|
|
|
class SdkBody:
|
|
def __init__(self, value):
|
|
self._stream = io.BytesIO(value)
|
|
|
|
def read(self):
|
|
return self._stream.read()
|
|
|
|
def iter_bytes(self, *, block_size):
|
|
while True:
|
|
chunk = self._stream.read(block_size)
|
|
if not chunk:
|
|
return
|
|
yield chunk
|
|
|
|
def close(self):
|
|
self._stream.close()
|
|
|
|
|
|
class FakeOssClient:
|
|
def __init__(
|
|
self,
|
|
*,
|
|
version_status=None,
|
|
acl="private",
|
|
block_public_access=True,
|
|
location="oss-cn-hangzhou",
|
|
encrypted=True,
|
|
):
|
|
self.version_status = version_status
|
|
self.acl = acl
|
|
self.block_public_access = block_public_access
|
|
self.location = location
|
|
self.encrypted = encrypted
|
|
self.objects = {}
|
|
|
|
def get_bucket_info(self, request):
|
|
return SimpleNamespace(
|
|
bucket_info=SimpleNamespace(
|
|
versioning=self.version_status,
|
|
acl=self.acl,
|
|
block_public_access=self.block_public_access,
|
|
location=self.location,
|
|
sse_rule=object() if self.encrypted else None,
|
|
)
|
|
)
|
|
|
|
def put_object_from_file(self, request, source):
|
|
if request.key in self.objects:
|
|
raise OperationFailure(ServiceFailure(409, "FileAlreadyExists"))
|
|
value = Path(source).read_bytes()
|
|
self.objects[request.key] = {
|
|
"value": value,
|
|
"metadata": dict(request.metadata),
|
|
"content_type": request.content_type,
|
|
"acl": request.acl,
|
|
}
|
|
return SimpleNamespace(etag="put-etag", version_id=None)
|
|
|
|
def head_object(self, request):
|
|
item = self.objects.get(request.key)
|
|
if item is None:
|
|
raise OperationFailure(ServiceFailure(404, "NoSuchKey"))
|
|
return SimpleNamespace(
|
|
content_length=len(item["value"]),
|
|
content_type=item["content_type"],
|
|
metadata=dict(item["metadata"]),
|
|
etag="head-etag",
|
|
version_id=None,
|
|
)
|
|
|
|
def get_object(self, request):
|
|
item = self.objects.get(request.key)
|
|
if item is None:
|
|
raise OperationFailure(ServiceFailure(404, "NoSuchKey"))
|
|
return SimpleNamespace(body=SdkBody(item["value"]))
|
|
|
|
def copy_object(self, request):
|
|
if request.key in self.objects:
|
|
raise OperationFailure(ServiceFailure(409, "FileAlreadyExists"))
|
|
source = self.objects.get(request.source_key)
|
|
if source is None:
|
|
raise OperationFailure(ServiceFailure(404, "NoSuchKey"))
|
|
self.objects[request.key] = {
|
|
"value": source["value"],
|
|
"metadata": dict(request.metadata),
|
|
"content_type": request.content_type,
|
|
"acl": request.acl,
|
|
}
|
|
return SimpleNamespace(etag="copy-etag", version_id=None)
|
|
|
|
def delete_object(self, request):
|
|
self.objects.pop(request.key, None)
|
|
|
|
|
|
class AliyunOssV2Tests(unittest.TestCase):
|
|
def client(self, fake: FakeOssClient) -> AliyunOssV2Client:
|
|
return AliyunOssV2Client(
|
|
AliyunOssConfig("cn-hangzhou", "arr-private-test"),
|
|
sdk=FakeSdk,
|
|
client=fake,
|
|
)
|
|
|
|
def test_configuration_and_bucket_versioning_fail_closed(self) -> None:
|
|
config = AliyunOssConfig.from_environment(
|
|
{
|
|
"ARR_OSS_REGION": "cn-hangzhou",
|
|
"ARR_OSS_BUCKET": "arr-private-test",
|
|
"ARR_OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
|
|
}
|
|
)
|
|
self.assertEqual(config.bucket, "arr-private-test")
|
|
exchange = OutputExchangeConfig.from_environment(
|
|
{"ARR_AGENT_OUTPUT_PREFIX": "private-agent-results"}
|
|
)
|
|
self.assertEqual(
|
|
exchange.object_key("handle_001"),
|
|
"private-agent-results/handle_001",
|
|
)
|
|
self.client(FakeOssClient()).assert_immutable_writes_supported()
|
|
self.client(
|
|
FakeOssClient(acl="public-read", block_public_access=False)
|
|
).assert_immutable_writes_supported()
|
|
for status in ("Enabled", "Suspended"):
|
|
with self.subTest(status=status), self.assertRaises(CloudClientError) as caught:
|
|
self.client(FakeOssClient(version_status=status)).assert_immutable_writes_supported()
|
|
self.assertEqual(caught.exception.kind, "versioning_incompatible")
|
|
incompatible = (
|
|
(FakeOssClient(acl="public-read-write"), "public_access_incompatible"),
|
|
(FakeOssClient(encrypted=False), "encryption_incompatible"),
|
|
(FakeOssClient(location="oss-cn-beijing"), "region_mismatch"),
|
|
)
|
|
for fake, expected in incompatible:
|
|
with self.subTest(expected=expected), self.assertRaises(CloudClientError) as caught:
|
|
self.client(fake).assert_immutable_writes_supported()
|
|
self.assertEqual(caught.exception.kind, expected)
|
|
|
|
def test_constructed_sdk_client_does_not_trust_ambient_proxies(self) -> None:
|
|
created = {}
|
|
|
|
class Client:
|
|
def __init__(self, config):
|
|
created["config"] = config
|
|
|
|
def close(self):
|
|
created["closed"] = True
|
|
|
|
sdk = SimpleNamespace(
|
|
credentials=SimpleNamespace(
|
|
EnvironmentVariableCredentialsProvider=lambda: object()
|
|
),
|
|
config=SimpleNamespace(
|
|
load_default=lambda: SimpleNamespace(
|
|
credentials_provider=None,
|
|
region=None,
|
|
endpoint=None,
|
|
http_client=None,
|
|
)
|
|
),
|
|
transport=SimpleNamespace(
|
|
RequestsHttpClient=lambda **kwargs: SimpleNamespace(**kwargs)
|
|
),
|
|
Client=Client,
|
|
)
|
|
client = AliyunOssV2Client(
|
|
AliyunOssConfig("cn-guangzhou", "arr-private-test"),
|
|
sdk=sdk,
|
|
)
|
|
session = created["config"].http_client.session
|
|
self.assertFalse(session.trust_env)
|
|
client.close()
|
|
self.assertTrue(created["closed"])
|
|
|
|
def test_managed_store_uses_conditional_staged_commit_and_rechecks_bytes(self) -> None:
|
|
fake = FakeOssClient()
|
|
client = self.client(fake)
|
|
store = ManagedObjectStore(
|
|
CloudObjectBackend(client),
|
|
ObjectKeyPolicy("arr-test"),
|
|
)
|
|
with tempfile.TemporaryDirectory() as temporary:
|
|
source = Path(temporary) / "source.xml"
|
|
source.write_bytes(b"<root />")
|
|
committed = store.upload_committed(
|
|
job_id="arr-job-oss-001",
|
|
attempt_no=1,
|
|
role="source_xml",
|
|
source=source,
|
|
original_filename="source.xml",
|
|
)
|
|
self.assertEqual(committed.state, "committed")
|
|
self.assertIn("/committed/source_xml/", committed.object_key)
|
|
self.assertFalse(any("/staged/" in key for key in fake.objects))
|
|
self.assertTrue(all(item["acl"] == "private" for item in fake.objects.values()))
|
|
with client.stream_object(committed.object_key) as reader:
|
|
first = reader.read(3)
|
|
self.assertEqual(first + reader.read(), source.read_bytes())
|
|
self.assertEqual(
|
|
store.inspect_committed(committed.object_key, "source.xml").sha256,
|
|
committed.sha256,
|
|
)
|
|
|
|
def test_exchange_publisher_is_idempotent_and_remote_port_is_prefix_bound(self) -> None:
|
|
fake = FakeOssClient()
|
|
client = self.client(fake)
|
|
config = OutputExchangeConfig("arr-agent-output-test")
|
|
publisher = OutputExchangePublisher(client, config)
|
|
with tempfile.TemporaryDirectory() as temporary:
|
|
root = Path(temporary)
|
|
source = root / "result.json"
|
|
source.write_bytes(b'{"ok":true}')
|
|
metadata = {
|
|
"arr-sha256": "a" * 64,
|
|
"arr-byte-size": str(source.stat().st_size),
|
|
"arr-mime-type": "application/json",
|
|
}
|
|
first = publisher.publish(
|
|
file_handle="arrout_test-001",
|
|
source=source,
|
|
mime_type="application/json",
|
|
metadata=metadata,
|
|
)
|
|
second = publisher.publish(
|
|
file_handle="arrout_test-001",
|
|
source=source,
|
|
mime_type="application/json",
|
|
metadata=metadata,
|
|
)
|
|
self.assertEqual(first.object_key, second.object_key)
|
|
destination = root / "materialized.json"
|
|
PrefixRemoteFilePort(client, config).materialize(
|
|
"arrout_test-001", destination, 1024
|
|
)
|
|
self.assertEqual(destination.read_bytes(), source.read_bytes())
|
|
with self.assertRaises(ProcessingTransportError):
|
|
PrefixRemoteFilePort(client, config).materialize(
|
|
"../escape", root / "escape", 1024
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|