|
| 1 | +import asyncio |
| 2 | +import logging |
| 3 | +import socket |
| 4 | +import threading |
| 5 | +from datetime import datetime, timedelta, timezone |
| 6 | +from typing import ClassVar |
| 7 | +from unittest.mock import MagicMock, patch |
| 8 | +from uuid import UUID, uuid4 |
| 9 | + |
| 10 | +import grpc |
| 11 | +import pytest |
| 12 | +from google.protobuf.empty_pb2 import Empty |
| 13 | + |
| 14 | +from tilebox.datasets.uuid import must_uuid_to_uuid_message |
| 15 | +from tilebox.workflows import ExecutionContext, Runner, Task |
| 16 | +from tilebox.workflows.cache import InMemoryCache, JobCache |
| 17 | +from tilebox.workflows.data import ExecutionStats, Job, JobState, RunnerContext, TaskState |
| 18 | +from tilebox.workflows.data import Task as TaskData |
| 19 | +from tilebox.workflows.observability.tracing import NoopWorkflowTracer |
| 20 | +from tilebox.workflows.runner.executor import LazyStorageLocations |
| 21 | +from tilebox.workflows.runner.worker_server import serve_runner |
| 22 | +from tilebox.workflows.task import TaskMeta |
| 23 | +from tilebox.workflows.workflows.v1 import core_pb2, worker_pb2, worker_pb2_grpc |
| 24 | + |
| 25 | + |
| 26 | +def test_worker_executes_tasks_concurrently_with_isolated_execution_state( |
| 27 | + caplog: pytest.LogCaptureFixture, |
| 28 | +) -> None: |
| 29 | + class SharedRunnerContext(RunnerContext): |
| 30 | + instances: ClassVar[list["SharedRunnerContext"]] = [] |
| 31 | + |
| 32 | + def __init__(self, tracer: NoopWorkflowTracer) -> None: |
| 33 | + super().__init__(tracer) |
| 34 | + self.instances.append(self) |
| 35 | + |
| 36 | + class ConcurrentTask(Task): |
| 37 | + label: str |
| 38 | + |
| 39 | + barrier: ClassVar[threading.Barrier] = threading.Barrier(2) |
| 40 | + observations: ClassVar[list[tuple[str, int, int, int, int]]] = [] |
| 41 | + observations_lock = threading.Lock() |
| 42 | + |
| 43 | + async def execute(self, context: ExecutionContext) -> None: |
| 44 | + await asyncio.sleep(0) |
| 45 | + with self.observations_lock: |
| 46 | + self.observations.append( |
| 47 | + (self.label, id(self), id(context), id(context.runner_context), id(asyncio.get_running_loop())) |
| 48 | + ) |
| 49 | + |
| 50 | + cache: JobCache = context.job_cache # ty: ignore[unresolved-attribute] |
| 51 | + cache[self.label] = self.label.encode() |
| 52 | + context.logger.info("Concurrent task executing", label=self.label) |
| 53 | + context.progress(self.label).add(1) |
| 54 | + self.barrier.wait(timeout=5) |
| 55 | + context.progress(self.label).done(1) |
| 56 | + |
| 57 | + cache = InMemoryCache() |
| 58 | + runner = Runner(tasks=[ConcurrentTask], cache=cache, context=SharedRunnerContext) |
| 59 | + fake_client = MagicMock() |
| 60 | + fake_client._tracer = NoopWorkflowTracer() |
| 61 | + fake_client._task_logger = logging.getLogger("tilebox.workflows.tests.shared-worker") |
| 62 | + caplog.set_level(logging.INFO, logger=fake_client._task_logger.name) |
| 63 | + |
| 64 | + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as free_socket: |
| 65 | + free_socket.bind(("127.0.0.1", 0)) |
| 66 | + address = f"127.0.0.1:{free_socket.getsockname()[1]}" |
| 67 | + |
| 68 | + server_thread = threading.Thread(target=serve_runner, args=(runner, address), daemon=True) |
| 69 | + |
| 70 | + with patch("tilebox.workflows.runner.worker_service.Client", return_value=fake_client): |
| 71 | + server_thread.start() |
| 72 | + channel = grpc.insecure_channel(address) |
| 73 | + grpc.channel_ready_future(channel).result(timeout=5) |
| 74 | + stub = worker_pb2_grpc.WorkerServiceStub(channel) |
| 75 | + stub.InitializeWorker( |
| 76 | + worker_pb2.InitializeRunnerRequest(runner_id=must_uuid_to_uuid_message(uuid4())), |
| 77 | + timeout=5, |
| 78 | + ) |
| 79 | + |
| 80 | + job = _job() |
| 81 | + tasks = [_task_message(ConcurrentTask(label), job) for label in ("first", "second")] |
| 82 | + responses = [stub.ExecuteTask.future(task, timeout=5) for task in tasks] |
| 83 | + |
| 84 | + try: |
| 85 | + results = [response.result() for response in responses] |
| 86 | + finally: |
| 87 | + stub.ShutdownWorker(Empty(), timeout=5) |
| 88 | + channel.close() |
| 89 | + server_thread.join(timeout=10) |
| 90 | + |
| 91 | + assert not server_thread.is_alive() |
| 92 | + assert len(SharedRunnerContext.instances) == 1 |
| 93 | + assert all(result.HasField("computed_task") for result in results) |
| 94 | + assert [result.computed_task.progress_updates[0].label for result in results] == ["first", "second"] |
| 95 | + assert all(result.computed_task.progress_updates[0].total == 1 for result in results) |
| 96 | + assert all(result.computed_task.progress_updates[0].done == 1 for result in results) |
| 97 | + |
| 98 | + assert {observation[0] for observation in ConcurrentTask.observations} == {"first", "second"} |
| 99 | + assert len({observation[1] for observation in ConcurrentTask.observations}) == 2 |
| 100 | + assert len({observation[2] for observation in ConcurrentTask.observations}) == 2 |
| 101 | + assert {observation[3] for observation in ConcurrentTask.observations} == {id(SharedRunnerContext.instances[0])} |
| 102 | + assert len({observation[4] for observation in ConcurrentTask.observations}) == 2 |
| 103 | + assert sorted(cache.group(str(job.id)).items()) == [("first", b"first"), ("second", b"second")] |
| 104 | + |
| 105 | + log_attributes = [ |
| 106 | + record.tilebox_structured_log_attributes # ty: ignore[unresolved-attribute] |
| 107 | + for record in caplog.records |
| 108 | + if record.message == "Concurrent task executing" |
| 109 | + ] |
| 110 | + assert {attributes["label"] for attributes in log_attributes} == {"first", "second"} |
| 111 | + assert {attributes["task_id"] for attributes in log_attributes} == {str(UUID(bytes=task.id.uuid)) for task in tasks} |
| 112 | + |
| 113 | + |
| 114 | +def test_lazy_storage_locations_are_loaded_once_during_concurrent_access() -> None: |
| 115 | + storage_location = MagicMock() |
| 116 | + storage_location.id = UUID(int=1) |
| 117 | + storage_location._with_runner_context.return_value = storage_location |
| 118 | + |
| 119 | + load_started = threading.Event() |
| 120 | + release_load = threading.Event() |
| 121 | + |
| 122 | + def storage_locations() -> list[MagicMock]: |
| 123 | + load_started.set() |
| 124 | + release_load.wait(timeout=5) |
| 125 | + return [storage_location] |
| 126 | + |
| 127 | + client = MagicMock() |
| 128 | + client.automations.return_value.storage_locations.side_effect = storage_locations |
| 129 | + locations = LazyStorageLocations(client, RunnerContext()) |
| 130 | + |
| 131 | + first = threading.Thread(target=len, args=(locations,)) |
| 132 | + first.start() |
| 133 | + assert load_started.wait(timeout=5) |
| 134 | + |
| 135 | + second_access_started = threading.Event() |
| 136 | + |
| 137 | + def read_locations() -> None: |
| 138 | + second_access_started.set() |
| 139 | + len(locations) |
| 140 | + |
| 141 | + second = threading.Thread(target=read_locations) |
| 142 | + second.start() |
| 143 | + assert second_access_started.wait(timeout=5) |
| 144 | + release_load.set() |
| 145 | + first.join(timeout=5) |
| 146 | + second.join(timeout=5) |
| 147 | + |
| 148 | + assert not first.is_alive() |
| 149 | + assert not second.is_alive() |
| 150 | + client.automations.return_value.storage_locations.assert_called_once_with() |
| 151 | + assert list(locations) == [storage_location.id] |
| 152 | + |
| 153 | + |
| 154 | +def _job() -> Job: |
| 155 | + return Job( |
| 156 | + id=uuid4(), |
| 157 | + name="concurrent worker test", |
| 158 | + trace_parent="00-0123456789abcdef0123456789abcdef-0123456789abcdef-01", |
| 159 | + state=JobState.RUNNING, |
| 160 | + submitted_at=datetime.now(tz=timezone.utc), |
| 161 | + progress=[], |
| 162 | + execution_stats=ExecutionStats(None, None, timedelta(), timedelta(), 0, 1, {}), |
| 163 | + ) |
| 164 | + |
| 165 | + |
| 166 | +def _task_message(task: Task, job: Job) -> core_pb2.Task: |
| 167 | + identifier = TaskMeta.for_task(task).identifier |
| 168 | + return TaskData( |
| 169 | + id=uuid4(), |
| 170 | + identifier=identifier, |
| 171 | + state=TaskState.RUNNING, |
| 172 | + input=task._serialize(), |
| 173 | + display=type(task).__name__, |
| 174 | + job=job, |
| 175 | + ).to_message() |
0 commit comments