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:
@@ -340,6 +340,16 @@ ensuring data is stored with proper parameter values.
|
||||
a `_TextualPromptInput` composite (Horizontal + Static + Input) when Textual is
|
||||
available, and a `_FallbackPromptInput` otherwise. Zero `# type: ignore`
|
||||
suppressions — all typing uses Protocol definitions and `cast()`.
|
||||
- **InvariantService thread safety** (#7524): Added `threading.RLock` to protect
|
||||
`_invariants` (dict) and `_enforcement_records` (list) from concurrent access
|
||||
during parallel plan execution. Prevents `RuntimeError: dictionary changed size
|
||||
during iteration` and data corruption when multiple subplans simultaneously
|
||||
call `add_invariant()`, `list_invariants()`, `remove_invariant()`,
|
||||
`get_effective_invariants()`, or `enforce_invariants()`. Also adds three
|
||||
thread-safe helper read methods: `get_enforcement_records()`, `get_invariant()`,
|
||||
and `get_invariants_snapshot()`. Includes comprehensive BDD tests for concurrent
|
||||
access patterns (adds, lists, removes, enforcement, and mixed operations).
|
||||
|
||||
- **Actor CLI NAME argument made optional, derived from YAML config** (#4186): The
|
||||
`agents actor add` positional `NAME` argument is now optional (defaults to
|
||||
`None`). When omitted, the actor name is derived from the `name` field in
|
||||
|
||||
@@ -38,6 +38,7 @@ Below are some of the specific details of various contributions.
|
||||
* HAL 9000 has contributed the plan artifacts JSON completeness fix (#9084): ensured `validation_summary` and `apply_summary` are correctly included in `_build_artifacts_dict`, removing stale `@tdd_expected_fail` tags from Behave scenarios to enable full regression test coverage.
|
||||
* HAL 9000 has contributed the bug-hunt-pool-supervisor non-blocking tracking fix (#7875 / PR #7957): updated step 5 to be best-effort and added rule 9 to prevent the automation-tracking-manager call from blocking the main supervisor loop.
|
||||
* Jeffrey Phillips Freeman has contributed the complete AUTO-BUG-SUP tracking prefix fix across agent-system-specification.md, automation-tracking.md documentation, replacing the incorrect `AUTO-BUG-POL` prefix with the correct `AUTO-BUG-SUP` prefix used by the bug-hunt-pool-supervisor agent (#7875).
|
||||
* HAL 9000 has contributed thread safety to InvariantService (issue #7524): added `threading.RLock` protection for all shared mutable state (_invariants dict, _enforcement_records list) across add, list, remove, effective-set computation, and enforcement operations, preventing RuntimeError: dictionary changed size during iteration in multi-threaded parallel plan execution environments.
|
||||
* HAL 9000 has contributed the plugin entry point security hardening fix (#7476): enforced entry point allowlist validation before importing plugin modules to prevent malicious plugin loading.
|
||||
* HAL 9000 has contributed the benchmark workflow separation (#9040): moved the benchmark-regression job out of the default PR workflow into a dedicated scheduled workflow, reducing median PR CI turnaround time from 99-132 minutes to under 30 minutes.
|
||||
* HAL 9000 has contributed the plan tree JSON/YAML command envelope fix (#9163): wrapped `agents plan tree --format json/yaml` output in the spec-required command envelope structure, added summary statistics, decision_ids mapping, child_plans list, and accurate timing measurement.
|
||||
|
||||
@@ -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
|
||||
@@ -14,11 +14,19 @@ a dict keyed by invariant ID.
|
||||
Effective invariants are computed using plan > action > project > global
|
||||
order. See ``merge_invariants`` for de-duplication semantics.
|
||||
|
||||
## Thread Safety
|
||||
|
||||
A ``threading.RLock`` (re-entrant lock) guards all shared mutable state
|
||||
(``_invariants`` dict, ``_enforcement_records`` list). Every public method
|
||||
acquires the lock before reading or writing shared data; the lock is re-entrant
|
||||
so that nested calls from within a locked block do **not** deadlock.
|
||||
|
||||
Based on ``docs/specification.md`` and implementation plan Stage M3.5.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from threading import RLock
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
import structlog
|
||||
@@ -50,6 +58,10 @@ class InvariantService:
|
||||
|
||||
Provides add, list, remove (soft-delete), effective-set computation,
|
||||
and enforcement record creation. All storage is in-memory.
|
||||
|
||||
Thread safety is provided via a ``threading.RLock`` that protects all
|
||||
shared state mutations so that concurrent readers and writers cannot
|
||||
race or raise ``RuntimeError: dictionary changed size during iteration``.
|
||||
"""
|
||||
|
||||
def __init__(self, event_bus: EventBus | None = None) -> None:
|
||||
@@ -60,6 +72,7 @@ class InvariantService:
|
||||
"""
|
||||
self._invariants: dict[str, Invariant] = {}
|
||||
self._enforcement_records: list[InvariantEnforcementRecord] = []
|
||||
self._lock = RLock() # guards _invariants and _enforcement_records
|
||||
self._logger = logger.bind(service="invariant")
|
||||
self._sanitizer = PromptSanitizer()
|
||||
self._event_bus = event_bus
|
||||
@@ -99,7 +112,9 @@ class InvariantService:
|
||||
source_name=source_name.strip(),
|
||||
)
|
||||
|
||||
self._invariants[invariant.id] = invariant
|
||||
with self._lock:
|
||||
self._invariants[invariant.id] = invariant
|
||||
|
||||
self._logger.info(
|
||||
"Invariant added",
|
||||
invariant_id=invariant.id,
|
||||
@@ -132,7 +147,8 @@ class InvariantService:
|
||||
project_name=source_name if scope == InvariantScope.PROJECT else None,
|
||||
)
|
||||
|
||||
result = [inv for inv in self._invariants.values() if inv.active]
|
||||
with self._lock:
|
||||
result = [inv for inv in self._invariants.values() if inv.active]
|
||||
|
||||
if scope is not None:
|
||||
result = [inv for inv in result if inv.scope == scope]
|
||||
@@ -157,15 +173,17 @@ class InvariantService:
|
||||
if not invariant_id or not invariant_id.strip():
|
||||
raise ValidationError("Invariant ID must not be empty")
|
||||
|
||||
inv = self._invariants.get(invariant_id)
|
||||
if inv is None:
|
||||
raise NotFoundError(
|
||||
resource_type="invariant",
|
||||
resource_id=invariant_id,
|
||||
)
|
||||
# Invariant is frozen (immutable); create a new instance with active=False
|
||||
deactivated = inv.model_copy(update={"active": False})
|
||||
self._invariants[invariant_id] = deactivated
|
||||
with self._lock:
|
||||
inv = self._invariants.get(invariant_id)
|
||||
if inv is None:
|
||||
raise NotFoundError(
|
||||
resource_type="invariant",
|
||||
resource_id=invariant_id,
|
||||
)
|
||||
# Invariant is frozen (immutable); create a new instance with active=False
|
||||
deactivated = inv.model_copy(update={"active": False})
|
||||
self._invariants[invariant_id] = deactivated
|
||||
|
||||
self._logger.info("Invariant removed (soft-delete)", invariant_id=invariant_id)
|
||||
return deactivated
|
||||
|
||||
@@ -194,7 +212,8 @@ class InvariantService:
|
||||
Returns:
|
||||
Merged, de-duplicated list of effective invariants.
|
||||
"""
|
||||
active = [inv for inv in self._invariants.values() if inv.active]
|
||||
with self._lock:
|
||||
active = [inv for inv in self._invariants.values() if inv.active]
|
||||
|
||||
plan_invs = [
|
||||
inv
|
||||
@@ -467,7 +486,9 @@ class InvariantService:
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
self._enforcement_records.extend(records)
|
||||
with self._lock:
|
||||
self._enforcement_records.extend(records)
|
||||
|
||||
self._logger.info(
|
||||
"Invariants enforced",
|
||||
plan_id=plan_id,
|
||||
@@ -513,3 +534,37 @@ class InvariantService:
|
||||
exc_info=True,
|
||||
)
|
||||
return records
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Thread-safe read helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
def get_enforcement_records(self) -> list[InvariantEnforcementRecord]:
|
||||
"""Return enforcement records (thread-safe read).
|
||||
|
||||
Returns:
|
||||
List of ``InvariantEnforcementRecord`` objects collected so far.
|
||||
"""
|
||||
with self._lock:
|
||||
return list(self._enforcement_records)
|
||||
|
||||
def get_invariant(self, invariant_id: str) -> Invariant | None:
|
||||
"""Return a single invariant by ID (thread-safe read).
|
||||
|
||||
Args:
|
||||
invariant_id: The ULID of the desired invariant.
|
||||
|
||||
Returns:
|
||||
The ``Invariant`` if found, else ``None``.
|
||||
"""
|
||||
with self._lock:
|
||||
return self._invariants.get(invariant_id)
|
||||
|
||||
def get_invariants_snapshot(self) -> dict[str, Invariant]:
|
||||
"""Return a snapshot copy of all invariants (thread-safe read).
|
||||
|
||||
Returns:
|
||||
A shallow copy of the internal ``_invariants`` dict.
|
||||
"""
|
||||
with self._lock:
|
||||
return dict(self._invariants)
|
||||
|
||||
Reference in New Issue
Block a user