Skip to content

Commit 00939bb

Browse files
Allow concurrent workers in one shared runtime (#54)
1 parent 0e9f036 commit 00939bb

9 files changed

Lines changed: 271 additions & 43 deletions

File tree

‎CHANGELOG.md‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,8 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
1010
### Fixed
1111

1212
- `tilebox-datasets`: Allow `iter_datapoints()` to handle empty query results.
13+
- `tilebox-workflows`: Make built-in worker state safe for concurrent task execution in a shared Python runtime and
14+
document the concurrency contract for custom runner contexts, caches, and shared task state.
1315

1416
## [0.60.0] - 2026-08-25
1517

‎tilebox-workflows/README.md‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -71,6 +71,21 @@ runner = client.runner(tasks=[MyFirstTask])
7171
runner.run_all()
7272
```
7373

74+
## Concurrent worker execution
75+
76+
A worker runtime can execute multiple tasks concurrently in one Python process. Each execution receives a newly
77+
deserialized task instance and its own `ExecutionContext`, including task-local subtask and progress state. The
78+
`RunnerContext`, configured `JobCache`, and any class or module state are process-level resources shared by those
79+
executions.
80+
81+
Custom runner contexts, caches, and shared task state must therefore support concurrent access from multiple threads.
82+
Asynchronous task executions may also run on different event loops. Configure and register these resources before the
83+
worker starts; do not mutate runner configuration while tasks are executing. Compound cache operations are not atomic
84+
unless the cache implementation explicitly provides that guarantee.
85+
86+
Concurrency in one runtime avoids repeated process initialization and allows overlapping I/O or native code that
87+
releases Python's GIL. CPU-bound Python code still needs multiple runtime processes for parallel execution.
88+
7489
## Documentation
7590

7691
Check out the [Tilebox Workflows documentation](https://docs.tilebox.com/workflows/introduction) for more information.
Lines changed: 175 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,175 @@
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()

‎tilebox-workflows/tilebox/workflows/cache.py‎

Lines changed: 36 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
from io import BytesIO
66
from pathlib import Path
77
from pathlib import PurePosixPath as ObjectPath
8+
from threading import RLock
89
from typing import TYPE_CHECKING, Any
910

1011
if TYPE_CHECKING:
@@ -17,6 +18,13 @@
1718

1819

1920
class JobCache(ABC):
21+
"""Cache shared by tasks belonging to the same job.
22+
23+
Task executions may access a cache concurrently, including from different threads in a shared worker runtime.
24+
Implementations must therefore make individual cache operations and :meth:`group` thread-safe. Sequences of
25+
operations, such as checking for a key before setting it, are not atomic.
26+
"""
27+
2028
@abstractmethod
2129
def __contains__(self, key: str) -> bool: ...
2230
@abstractmethod
@@ -121,46 +129,50 @@ def group(self, key: str) -> "ObstoreCache":
121129

122130

123131
class InMemoryCache(JobCache):
124-
def __init__(self) -> None:
132+
def __init__(self, *, _lock: Any | None = None) -> None:
125133
"""A simple in-memory cache implementation.
126134
127135
Useful for testing and development. Provides no persistence, and
128136
no way of sharing data between multiple task runners.
129137
"""
130138
self.cache: dict[str, bytes | InMemoryCache] = {}
139+
self._lock = _lock or RLock()
131140

132141
def __contains__(self, key: str) -> bool:
133-
return key in self.cache
142+
with self._lock:
143+
return key in self.cache
134144

135145
def __setitem__(self, key: str, value: bytes) -> None:
136-
parent_group, key = self._resolve_slashes(key, create_missing=False)
137-
parent_group.cache[key] = value
146+
with self._lock:
147+
parent_group, key = self._resolve_slashes(key, create_missing=False)
148+
parent_group.cache[key] = value
138149

139150
def __getitem__(self, key: str) -> bytes:
140-
parent_group, key = self._resolve_slashes(key, create_missing=False)
141-
item = parent_group.cache[key]
142-
if not isinstance(item, bytes):
143-
# item is a directory
144-
raise KeyError(f"{key} is not cached!")
145-
return item
151+
with self._lock:
152+
parent_group, key = self._resolve_slashes(key, create_missing=False)
153+
item = parent_group.cache[key]
154+
if not isinstance(item, bytes):
155+
# item is a directory
156+
raise KeyError(f"{key} is not cached!")
157+
return item
146158

147159
def __iter__(self) -> Iterator[str]:
148-
for k, v in self.cache.items():
149-
if isinstance(v, bytes):
150-
yield k
160+
with self._lock:
161+
return iter([key for key, value in self.cache.items() if isinstance(value, bytes)])
151162

152163
def group(self, key: str) -> "InMemoryCache":
153-
parent_group, key = self._resolve_slashes(key, create_missing=True)
154-
try:
155-
group = parent_group.cache[key]
156-
except KeyError:
157-
group = InMemoryCache()
158-
parent_group.cache[key] = group
164+
with self._lock:
165+
parent_group, key = self._resolve_slashes(key, create_missing=True)
166+
try:
167+
group = parent_group.cache[key]
168+
except KeyError:
169+
group = InMemoryCache(_lock=self._lock)
170+
parent_group.cache[key] = group
159171

160-
if not isinstance(group, InMemoryCache):
161-
# if key is a file, we return an empty group
162-
return InMemoryCache()
163-
return group
172+
if not isinstance(group, InMemoryCache):
173+
# if key is a file, we return an empty group
174+
return InMemoryCache(_lock=self._lock)
175+
return group
164176

165177
def _resolve_slashes(self, key: str, create_missing: bool = False) -> tuple["InMemoryCache", str]:
166178
"""Resolve slashes in a given cache key, by converting them into nested groups.
@@ -187,7 +199,7 @@ def _resolve_slashes(self, key: str, create_missing: bool = False) -> tuple["InM
187199
except KeyError:
188200
if create_missing:
189201
# create a new group for this key if it doesn't exist
190-
sub_group = InMemoryCache()
202+
sub_group = InMemoryCache(_lock=self._lock)
191203
group.cache[part] = sub_group
192204
else:
193205
raise KeyError(f"{part} is not cached!") from None

‎tilebox-workflows/tilebox/workflows/data.py‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1136,6 +1136,12 @@ def to_message(self) -> automation_pb.AutomationPrototype:
11361136

11371137

11381138
class RunnerContext:
1139+
"""Process-level context shared by task executions in a runner.
1140+
1141+
A runner creates one context instance during initialization. Worker runtimes may access that instance concurrently
1142+
from multiple threads, so subclasses must synchronize mutable state and use clients that support concurrent access.
1143+
"""
1144+
11391145
def __init__(
11401146
self,
11411147
tracer: WorkflowTracer | None = None,

‎tilebox-workflows/tilebox/workflows/runner/executor.py‎

Lines changed: 25 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -8,6 +8,7 @@
88
from concurrent.futures import ThreadPoolExecutor
99
from contextlib import AbstractContextManager, contextmanager
1010
from contextvars import copy_context
11+
from threading import RLock
1112
from typing import TYPE_CHECKING
1213
from uuid import UUID
1314
from warnings import warn
@@ -241,35 +242,42 @@ def __init__(self, client: Client, runner_context: RunnerContext) -> None:
241242
self._runner_context = runner_context
242243
self._locations: dict[UUID, StorageLocation] = {}
243244
self._loaded = False
245+
self._lock = RLock()
244246

245247
def _load(self) -> None:
246-
if self._loaded:
247-
return
248-
self._locations = {
249-
location.id: location._with_runner_context(self._runner_context) # noqa: SLF001
250-
for location in self._client.automations().storage_locations()
251-
}
252-
self._loaded = True
248+
with self._lock:
249+
if self._loaded:
250+
return
251+
self._locations = {
252+
location.id: location._with_runner_context(self._runner_context) # noqa: SLF001
253+
for location in self._client.automations().storage_locations()
254+
}
255+
self._loaded = True
253256

254257
def __getitem__(self, key: UUID) -> StorageLocation:
255-
self._load()
256-
return self._locations[key]
258+
with self._lock:
259+
self._load()
260+
return self._locations[key]
257261

258262
def __setitem__(self, key: UUID, value: StorageLocation) -> None:
259-
self._load()
260-
self._locations[key] = value
263+
with self._lock:
264+
self._load()
265+
self._locations[key] = value
261266

262267
def __delitem__(self, key: UUID) -> None:
263-
self._load()
264-
del self._locations[key]
268+
with self._lock:
269+
self._load()
270+
del self._locations[key]
265271

266272
def __iter__(self) -> Iterator[UUID]:
267-
self._load()
268-
return iter(self._locations)
273+
with self._lock:
274+
self._load()
275+
return iter(tuple(self._locations))
269276

270277
def __len__(self) -> int:
271-
self._load()
272-
return len(self._locations)
278+
with self._lock:
279+
self._load()
280+
return len(self._locations)
273281

274282

275283
def _finalize_mutable_progress_trackers(

0 commit comments

Comments
 (0)