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:
2026-05-08 13:25:52 +00:00
committed by Forgejo
parent a09873e10c
commit 7470f98155
5 changed files with 589 additions and 13 deletions
@@ -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