-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathtest_workspace_path_thread_safety.py
More file actions
124 lines (103 loc) · 4.38 KB
/
Copy pathtest_workspace_path_thread_safety.py
File metadata and controls
124 lines (103 loc) · 4.38 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
"""
Regression tests for issue #43 — thread-safe _workspace_path_override.
Run:
python -m unittest tests.test_workspace_path_thread_safety -v
"""
from __future__ import annotations
import os
import shutil
import sys
import tempfile
import threading
import unittest
from concurrent.futures import ThreadPoolExecutor, as_completed
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
sys.path.insert(0, REPO_ROOT)
from utils.workspace_path import (
get_workspace_path_override,
resolve_workspace_path,
set_workspace_path_override,
)
class TestWorkspacePathThreadSafety(unittest.TestCase):
"""Concurrent set-workspace + resolve must not observe torn global state."""
def setUp(self):
self.tmp = tempfile.mkdtemp(prefix="cursor-ws-thread-test-")
self.addCleanup(shutil.rmtree, self.tmp, ignore_errors=True)
self.path_a = os.path.join(self.tmp, "storage-a")
self.path_b = os.path.join(self.tmp, "storage-b")
os.makedirs(self.path_a)
os.makedirs(self.path_b)
self._prior_workspace_env = os.environ.pop("WORKSPACE_PATH", None)
self.addCleanup(self._restore_workspace_env)
self.addCleanup(set_workspace_path_override, None)
def _restore_workspace_env(self):
if self._prior_workspace_env is None:
os.environ.pop("WORKSPACE_PATH", None)
else:
os.environ["WORKSPACE_PATH"] = self._prior_workspace_env
def test_concurrent_set_and_resolve_never_returns_mixed_paths(self):
iterations = 500
errors: list[str] = []
start = threading.Barrier(9) # 1 writer + 8 readers
# Seed before workers start so readers never observe the unset default path.
set_workspace_path_override(self.path_a)
def writer() -> None:
start.wait()
for i in range(iterations):
set_workspace_path_override(self.path_a if i % 2 == 0 else self.path_b)
def reader() -> None:
start.wait()
for _ in range(iterations):
override = get_workspace_path_override()
if override is None:
errors.append("override was unexpectedly cleared during run")
continue
if override not in (self.path_a, self.path_b):
errors.append(f"override returned unexpected value: {override!r}")
continue
resolved = resolve_workspace_path()
expected = os.path.realpath(override)
if resolved != expected:
errors.append(
f"resolve {resolved!r} != realpath(override) {expected!r}"
)
with ThreadPoolExecutor(max_workers=9) as pool:
futures = [pool.submit(writer)]
futures.extend(pool.submit(reader) for _ in range(8))
for fut in as_completed(futures):
fut.result()
self.assertEqual(errors, [], "\n".join(errors[:20]))
def test_concurrent_clear_and_set_stays_consistent(self):
iterations = 200
errors: list[str] = []
start = threading.Barrier(5)
def toggler() -> None:
start.wait()
for i in range(iterations):
if i % 3 == 0:
set_workspace_path_override(None)
else:
set_workspace_path_override(
self.path_a if i % 2 == 0 else self.path_b
)
def reader() -> None:
start.wait()
for _ in range(iterations):
override = get_workspace_path_override()
resolved = resolve_workspace_path()
if override is None:
continue
if override not in (self.path_a, self.path_b):
errors.append(f"unexpected override: {override!r}")
elif resolved != os.path.realpath(override):
errors.append(
f"resolve {resolved!r} != realpath({override!r})"
)
with ThreadPoolExecutor(max_workers=5) as pool:
futures = [pool.submit(toggler)]
futures.extend(pool.submit(reader) for _ in range(4))
for fut in as_completed(futures):
fut.result()
self.assertEqual(errors, [], "\n".join(errors[:20]))
if __name__ == "__main__":
unittest.main()