fix(concurrency): add thread safety to InvariantService
Add threading.RLock to InvariantService to protect shared state (_invariants dict, _enforcement_records list) from concurrent access by multiple threads during parallel plan execution. Prevents RuntimeError: dictionary changed size during iteration and data corruption in multi-threaded environments. Changes: - Added self._lock = RLock() in __init__ - Wrapped all public methods (add_invariant, list_invariants, remove_invariant, get_effective_invariants, enforce_invariants) with lock acquisition via context managers - Added helper read methods: get_enforcement_records(), get_invariant(), get_invariants_snapshot() -- all thread-safe - Added BDD tests for concurrent access patterns - Updated CHANGELOG.md and CONTRIBUTORS.md ISSUES CLOSED: #7524
This commit is contained in:
@@ -0,0 +1,75 @@
|
||||
# This test captures the concurrency bug discovered in issue #7524.
|
||||
#
|
||||
# InvariantService stored invariants (_invariants dict) and enforcement records
|
||||
# (_enforcement_records list) without any threading.Lock. Concurrent calls from
|
||||
# parallel plan execution could raise RuntimeError: dictionary changed size during
|
||||
# iteration or silently corrupt data.
|
||||
#
|
||||
# The fix adds a threading.RLock to InvariantService guarding all shared mutable
|
||||
# state (see commit that implements the fix).
|
||||
|
||||
@tdd_issue @tdd_issue_7524
|
||||
Feature: TDD Issue #7524 - InvariantService Thread Safety
|
||||
InvariantService must be thread-safe so that concurrent calls from parallel
|
||||
plan execution threads cannot race on _invariants or _enforcement_records.
|
||||
|
||||
Scenario Outline: Concurrent add_invariant() must not corrupt the dict
|
||||
Given an invariant service
|
||||
When <n> threads concurrently add <count>d invariants through a barrier
|
||||
And the service has no RuntimeError raised during concurrent access
|
||||
Then the service should have exactly <total>d active invariants
|
||||
And all added invariants should be retrievable via get_invariant_snapshot
|
||||
|
||||
Examples:
|
||||
| n | count | total |
|
||||
| 2 | 5 | 10 |
|
||||
| 4 | 5 | 20 |
|
||||
| 8 | 3 | 24 |
|
||||
| 10| 2 | 20 |
|
||||
|
||||
Scenario Outline: Concurrent list_invariants() must not raise during iteration
|
||||
Given an invariant service with <count>d pre-existing invariants
|
||||
When 5 threads concurrently list all invariants through a barrier
|
||||
And no thread raised RuntimeError during concurrent reads
|
||||
Then the caller should receive a valid list of at least <min>d active invariants
|
||||
|
||||
Examples:
|
||||
| count | min |
|
||||
| 10 | 10 |
|
||||
| 20 | 20 |
|
||||
| 5 | 5 |
|
||||
|
||||
Scenario Outline: Concurrent remove_invariant() must not race with add
|
||||
Given an invariant service
|
||||
When <n> threads concurrently add invariants while <m> concurrent threads remove different ones through a barrier
|
||||
And no thread raised RuntimeError during mixed access
|
||||
Then the service should have a consistent count of active invariants
|
||||
|
||||
Examples:
|
||||
| n | m |
|
||||
| 5 | 2 |
|
||||
| 4 | 3 |
|
||||
| 8 | 4 |
|
||||
|
||||
Scenario Outline: Concurrent enforce_invariants() must not corrupt the record list
|
||||
Given an invariant service with <count>d invariants to enforce
|
||||
When <n> threads concurrently call enforce_invariants on the same set through a barrier
|
||||
And no thread raised RuntimeError during enforcement
|
||||
Then each thread should have received exactly <count>d enforcement records
|
||||
|
||||
Examples:
|
||||
| count | n |
|
||||
| 3 | 4 |
|
||||
| 5 | 8 |
|
||||
| 2 | 10|
|
||||
|
||||
Scenario Outline: Concurrent enforce and list must both succeed
|
||||
Given an invariant service with <count>d invariants already enforced by a previous thread
|
||||
When <n> threads concurrently enforce new invariants while other threads list them through a barrier
|
||||
And no thread raised RuntimeError during mixed enforcement and listing
|
||||
Then the enforcement record count should be consistent across all threads
|
||||
|
||||
Examples:
|
||||
| count | n |
|
||||
| 3 | 4 |
|
||||
| 5 | 6 |
|
||||
@@ -0,0 +1,435 @@
|
||||
"""Step definitions for invariant_service_thread_safety.feature.
|
||||
|
||||
Tests thread safety of InvariantService by launching multiple threads that
|
||||
perform concurrent operations on a shared InvariantService instance, verifying
|
||||
that no RuntimeError (dictionary changed size during iteration) is raised and
|
||||
that data remains consistent across all concurrent accesses.
|
||||
|
||||
See: https://git.cleverthis.com/cleveragents/cleveragents-core/issues/7524
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from typing import cast
|
||||
|
||||
from behave import given, then, when # type: ignore[import-untyped]
|
||||
from behave.runner import Context
|
||||
|
||||
from cleveragents.application.services.invariant_service import InvariantService
|
||||
from cleveragents.domain.models.core.invariant import InvariantScope
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Setup helpers
|
||||
# ================================================================
|
||||
|
||||
|
||||
def _create_barrier_and_threads(n: int) -> tuple[threading.Barrier, list[threading.Thread]]:
|
||||
"""Create a barrier for n threads and an empty thread list."""
|
||||
barrier = threading.Barrier(n)
|
||||
threads: list[threading.Thread] = []
|
||||
return barrier, threads
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Thread-safety state tracking per-thread
|
||||
# ================================================================
|
||||
|
||||
|
||||
@given("an invariant service")
|
||||
def step_service(context: Context) -> None:
|
||||
"""Create a fresh InvariantService with empty concurrency tracking."""
|
||||
context.service = InvariantService()
|
||||
context.concurrent_errors: list[Exception] = []
|
||||
context.concurrent_results: dict[str, object] = {}
|
||||
|
||||
|
||||
@given("an invariant service with {count:d} pre-existing invariants")
|
||||
def step_service_initialized(context: Context, count: int) -> None:
|
||||
"""Create an InvariantService already populated with {count} invariants."""
|
||||
context.service = InvariantService()
|
||||
for i in range(count):
|
||||
context.service.add_invariant(
|
||||
text=f"Constraint {i + 1}",
|
||||
scope=InvariantScope.GLOBAL,
|
||||
source_name="test-context",
|
||||
)
|
||||
context.concurrent_errors = []
|
||||
context.concurrent_results = {}
|
||||
|
||||
|
||||
@given("an invariant service with {count:d} invariants to enforce")
|
||||
def step_service_with_invariants(context: Context, count: int) -> None:
|
||||
"""Create an InvariantService populated with specific invariants for enforcement."""
|
||||
context.service = InvariantService()
|
||||
for i in range(count):
|
||||
inv = context.service.add_invariant(
|
||||
text=f"Enforced constraint {i + 1}",
|
||||
scope=InvariantScope.GLOBAL,
|
||||
source_name="enforcement-test",
|
||||
)
|
||||
# Store the invariant IDs so enforcement scenarios can reference them
|
||||
context.enforcement_inv_ids = list(context.concurrent_results) if hasattr(context, "concurrent_results") else []
|
||||
context.concurrent_errors = []
|
||||
context.concurrent_results = {}
|
||||
|
||||
|
||||
@given("an invariant service with {count:d} invariants already enforced by a previous thread")
|
||||
def step_service_with_enforced_invariants(context: Context, count: int) -> None:
|
||||
"""Create an InvariantService where enforcement records already exist."""
|
||||
context.service = InvariantService()
|
||||
invariants = []
|
||||
for i in range(count):
|
||||
inv = context.service.add_invariant(
|
||||
text=f"Already enforced {i + 1}",
|
||||
scope=InvariantScope.GLOBAL,
|
||||
source_name="pre-enforced",
|
||||
)
|
||||
invariants.append(inv)
|
||||
context.service.enforce_invariants(
|
||||
plan_id="previous-plan",
|
||||
invariants=invariants,
|
||||
)
|
||||
context.concurrent_errors = []
|
||||
context.concurrent_results = {}
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Concurrent add_invariant() — Scenario: concurrent adds
|
||||
# ================================================================
|
||||
|
||||
|
||||
@when(
|
||||
"<n> threads concurrently add <count>d invariants through a barrier",
|
||||
)
|
||||
def step_concurrent_add(context: Context, n: int, count: int) -> None:
|
||||
"""Launch n thread groups, each adding {count} invariants at once."""
|
||||
total = n * count # actual total added by this scenario
|
||||
barrier, threads = _create_barrier_and_threads(n)
|
||||
|
||||
def worker(thread_index: int) -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
for i in range(count):
|
||||
context.service.add_invariant(
|
||||
text=f"Concurrent add T{thread_index}-{i}",
|
||||
scope=InvariantScope.GLOBAL,
|
||||
source_name="concurrent-test",
|
||||
)
|
||||
except Exception as exc:
|
||||
with context._error_mutex if hasattr(context, "_error_mutex") else threading.Lock():
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
for idx in range(n):
|
||||
t = threading.Thread(target=worker, args=(idx,), daemon=True)
|
||||
threads.append(t)
|
||||
t.start()
|
||||
|
||||
# Also create a mutex for the first call since it may not exist yet
|
||||
if not hasattr(context, "_error_mutex"):
|
||||
context._error_mutex = threading.Lock()
|
||||
|
||||
for t in threads:
|
||||
t.join(timeout=30)
|
||||
|
||||
|
||||
@then("the service has no RuntimeError raised during concurrent access")
|
||||
def step_no_runtime_error_adds(context: Context) -> None:
|
||||
"""Assert none of the worker threads hit a RuntimeError."""
|
||||
runtime_errors = [
|
||||
e for e in context.concurrent_errors if isinstance(e, RuntimeError)
|
||||
]
|
||||
assert not runtime_errors, (
|
||||
f"RuntimeError raised during concurrent adds: {runtime_errors}. "
|
||||
"The RLock in InvariantService is likely missing or ineffective."
|
||||
)
|
||||
|
||||
|
||||
@then("the service should have exactly {total:d} active invariants")
|
||||
def step_concurrent_add_count(context: Context, total: int) -> None:
|
||||
"""Verify the expected number of invariants was added (10 from 2x5 etc.)."""
|
||||
actual = context.service.list_invariants()
|
||||
assert len(actual) == total, (
|
||||
f"Expected {total} active invariants after concurrent adds, "
|
||||
f"but got {len(actual)}. Data loss or corruption detected."
|
||||
)
|
||||
|
||||
|
||||
@then("all added invariants should be retrievable via get_invariant_snapshot")
|
||||
def step_snapshot_contains_all(context: Context) -> None:
|
||||
"""Confirm get_invariants_snapshot returns consistent data."""
|
||||
snapshot = context.service.get_invariants_snapshot()
|
||||
listed = context.service.list_invariants()
|
||||
assert len(snapshot) == len(listed), (
|
||||
f"Snapshot has {len(snapshot)} invariants but list_invariants found "
|
||||
f"{len(listed)}. Inconsistent reads detected."
|
||||
)
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Concurrent list_invariants() — Scenario: concurrent reads
|
||||
# ================================================================
|
||||
|
||||
|
||||
@when("5 threads concurrently list all invariants through a barrier")
|
||||
def step_concurrent_list(context: Context) -> None:
|
||||
"""Launch 5 threads that each call list_invariants simultaneously."""
|
||||
barrier = threading.Barrier(5)
|
||||
context.concurrent_errors = []
|
||||
context._list_results: list[list] = []
|
||||
|
||||
def worker() -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
results = context.service.list_invariants()
|
||||
with getattr(context, "_list_lock", threading.Lock()):
|
||||
context._list_results.append(results)
|
||||
except Exception as exc:
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
# Initialize lock if not set
|
||||
if not hasattr(context, "_list_lock"):
|
||||
context._list_lock = threading.Lock()
|
||||
|
||||
threads = [threading.Thread(target=worker, daemon=True) for _ in range(5)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=30)
|
||||
|
||||
|
||||
@then("no thread raised RuntimeError during concurrent reads")
|
||||
def step_no_runtime_error_lists(context: Context) -> None:
|
||||
"""Verify no thread hit a RuntimeError during list_invariants()."""
|
||||
runtime_errors = [e for e in context.concurrent_errors if isinstance(e, RuntimeError)]
|
||||
assert not runtime_errors, (
|
||||
f"RuntimeError raised during concurrent reads: {runtime_errors}. "
|
||||
"The RLock is missing in list_invariants()."
|
||||
)
|
||||
|
||||
|
||||
@then("the caller should receive a valid list of at least {min:d} active invariants")
|
||||
def step_list_produces_valid_result(context: Context, min: int) -> None:
|
||||
"""Each thread's list result should be consistent."""
|
||||
# All threads should have seen the same count (snapshot consistency)
|
||||
for results in context._list_results:
|
||||
assert len(results) >= min, (
|
||||
f"Thread saw only {len(results)} active invariants "
|
||||
f"(expected at least {min}). "
|
||||
"A reader may have seen a partially-written dict."
|
||||
)
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Concurrent remove + add — Scenario: mixed access
|
||||
# ================================================================
|
||||
|
||||
|
||||
@when(
|
||||
"<n> threads concurrently add invariants while <m> concurrent "
|
||||
"threads remove different ones through a barrier",
|
||||
)
|
||||
def step_concurrent_mixed(context: Context, n: int, m: int) -> None:
|
||||
"""Launch {n} add-threads and {m} remove-threads simultaneously."""
|
||||
total_workers = n + m
|
||||
barrier, threads = _create_barrier_and_threads(total_workers)
|
||||
|
||||
def add_worker(idx: int) -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
context.service.add_invariant(
|
||||
text=f"Concurrent add mixed-{idx}",
|
||||
scope=InvariantScope.GLOBAL,
|
||||
source_name="mixed-test",
|
||||
)
|
||||
except Exception as exc:
|
||||
with getattr(context, "_mix_lock", threading.Lock()):
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
def remove_worker(idx: int) -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
# Get some invariants to remove. Use snapshot to ensure list is stable.
|
||||
invs = context.service.list_invariants()
|
||||
if len(invs) > 0:
|
||||
target_id = invs[idx % len(invs)].id
|
||||
try:
|
||||
context.service.remove_invariant(target_id)
|
||||
except Exception:
|
||||
pass # already removed by another worker - expected race
|
||||
except Exception as exc:
|
||||
with getattr(context, "_mix_lock", threading.Lock()):
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
if not hasattr(context, "_mix_lock"):
|
||||
context._mix_lock = threading.Lock()
|
||||
context.concurrent_errors = []
|
||||
|
||||
for i in range(n):
|
||||
threads.append(threading.Thread(target=add_worker, args=(i,), daemon=True))
|
||||
remove_workers = n # index offset into the list; add workers occupy first `n` slots
|
||||
for i in range(m):
|
||||
threads.append(threading.Thread(target=remove_worker, args=(i,), daemon=True))
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=30)
|
||||
|
||||
|
||||
@then("no thread raised RuntimeError during mixed access")
|
||||
def step_no_runtime_error_mixed(context: Context) -> None:
|
||||
"""Verify no RuntimeError was thrown during concurrent add + remove."""
|
||||
runtime_errors = [e for e in context.concurrent_errors if isinstance(e, RuntimeError)]
|
||||
assert not runtime_errors, (
|
||||
f"RuntimeError raised during mixed access: {runtime_errors}. "
|
||||
"The RLock is missing or ineffective for concurrent read+write."
|
||||
)
|
||||
|
||||
|
||||
@then("the service should have a consistent count of active invariants")
|
||||
def step_consistent_mixed_count(context: Context) -> None:
|
||||
"""Verify we can consistently query the service after mixed operations."""
|
||||
try:
|
||||
_ = context.service.list_invariants()
|
||||
_ = context.service.get_invariants_snapshot()
|
||||
except RuntimeError as exc:
|
||||
raise AssertionError(
|
||||
f"Inconsistent state detected after mixed concurrent access: {exc}"
|
||||
) from exc
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Concurrent enforce_invariants() — Scenario: concurrent enforcement
|
||||
# ================================================================
|
||||
|
||||
|
||||
@when(
|
||||
"<n> threads concurrently call enforce_invariants on the same set "
|
||||
"through a barrier",
|
||||
)
|
||||
def step_concurrent_enforce(context: Context, n: int) -> None:
|
||||
"""Launch {n} threads that all enforce the same invariants simultaneously."""
|
||||
barrier = threading.Barrier(n)
|
||||
context.concurrent_errors = []
|
||||
context._enforcement_results: dict[int, list] = {}
|
||||
|
||||
# Re-read invariants fresh inside each thread to get stable snapshot.
|
||||
invariants = context.service.list_invariants()
|
||||
|
||||
def worker(thread_idx: int) -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
# Each thread re-reads from snapshot under lock (via the service)
|
||||
invs = context.service.list_invariants()
|
||||
records = context.service.enforce_invariants(
|
||||
plan_id=f"plan-thread-{thread_idx}",
|
||||
invariants=invs,
|
||||
)
|
||||
with getattr(context, "_enf_lock", threading.Lock()):
|
||||
context._enforcement_results[thread_idx] = records
|
||||
except Exception as exc:
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
if not hasattr(context, "_enf_lock"):
|
||||
context._enf_lock = threading.Lock()
|
||||
|
||||
threads = [threading.Thread(target=worker, args=(i,), daemon=True) for i in range(n)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=30)
|
||||
|
||||
|
||||
@then("each thread should have received exactly {count:d} enforcement records")
|
||||
def step_enforce_record_count(context: Context, count: int) -> None:
|
||||
"""Every thread's enforcement result list length must match expected count."""
|
||||
for idx, records in context._enforcement_results.items():
|
||||
assert len(records) == count, (
|
||||
f"Thread {idx} received {len(records)} enforcement records, "
|
||||
f"expected {count}. Concurrent enforcement corrupted the result."
|
||||
)
|
||||
|
||||
|
||||
@then("no thread raised RuntimeError during enforcement")
|
||||
def step_no_runtime_error_enforce(context: Context) -> None:
|
||||
"""Verify no RuntimeError during concurrent enforce_invariants()."""
|
||||
runtime_errors = [e for e in context.concurrent_errors if isinstance(e, RuntimeError)]
|
||||
assert not runtime_errors, (
|
||||
f"RuntimeError raised during enforcement: {runtime_errors}. "
|
||||
"_enforcement_records.extend() is likely unprotected."
|
||||
)
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Mixed enforce + list — Scenario: concurrent enforcement and listing
|
||||
# ================================================================
|
||||
|
||||
|
||||
@when(
|
||||
"<n> threads concurrently enforce new invariants while other threads "
|
||||
"list them through a barrier",
|
||||
)
|
||||
def step_concurrent_enforce_list(context: Context, n: int) -> None:
|
||||
"""Launch {n} enforcers and {n} lister workers simultaneously."""
|
||||
total_workers = n + n # enforcers + listers
|
||||
barrier, threads = _create_barrier_and_threads(total_workers)
|
||||
context.concurrent_errors = []
|
||||
|
||||
invariants = context.service.list_invariants()
|
||||
|
||||
def enforce_worker(idx: int) -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
invs = context.service.list_invariants()
|
||||
context.service.enforce_invariants(
|
||||
plan_id=f"enforce-worker-{idx}",
|
||||
invariants=invs,
|
||||
)
|
||||
except Exception as exc:
|
||||
with getattr(context, "_mix2_lock", threading.Lock()):
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
def list_worker(idx: int) -> None:
|
||||
try:
|
||||
barrier.wait(timeout=10)
|
||||
_ = context.service.list_invariants()
|
||||
except Exception as exc:
|
||||
with getattr(context, "_mix2_lock", threading.Lock()):
|
||||
context.concurrent_errors.append(exc)
|
||||
|
||||
if not hasattr(context, "_mix2_lock"):
|
||||
context._mix2_lock = threading.Lock()
|
||||
|
||||
for i in range(n):
|
||||
threads.append(threading.Thread(target=enforce_worker, args=(i,), daemon=True))
|
||||
for i in range(n):
|
||||
threads.append(threading.Thread(target=list_worker, args=(i,), daemon=True))
|
||||
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join(timeout=30)
|
||||
|
||||
|
||||
@then("no thread raised RuntimeError during mixed enforcement and listing")
|
||||
def step_no_runtime_error_mixed_enforce(context: Context) -> None:
|
||||
"""Verify no RuntimeError during concurrent enforce + list."""
|
||||
runtime_errors = [e for e in context.concurrent_errors if isinstance(e, RuntimeError)]
|
||||
assert not runtime_errors, (
|
||||
f"RuntimeError raised during mixed enforcement and listing: {runtime_errors}"
|
||||
)
|
||||
|
||||
|
||||
@then("the enforcement record count should be consistent across all threads")
|
||||
def step_enforcement_records_consistent(context: Context) -> None:
|
||||
"""Verify the final state is readable without errors."""
|
||||
try:
|
||||
_ = context.service.get_enforcement_records()
|
||||
_ = context.service.list_invariants()
|
||||
_ = context.service.get_invariants_snapshot()
|
||||
except RuntimeError as exc:
|
||||
raise AssertionError(
|
||||
f"Inconsistent state after concurrent enforce + list: {exc}"
|
||||
) from exc
|
||||
Reference in New Issue
Block a user