diff --git a/solstone/think/callosum.py b/solstone/think/callosum.py index b413950bf..901956ca8 100644 --- a/solstone/think/callosum.py +++ b/solstone/think/callosum.py @@ -49,6 +49,23 @@ class CallosumServer: with self.lock: return len(self.clients) + def _close_server_socket(self) -> None: + with self.lock: + server_socket = self.server_socket + self.server_socket = None + + if server_socket is None: + return + + try: + server_socket.shutdown(socket.SHUT_RDWR) + except OSError: + pass + try: + server_socket.close() + except OSError: + pass + def start(self) -> None: """Start the broadcast server.""" # Ensure health directory exists @@ -75,7 +92,11 @@ class CallosumServer: try: while not self.stop_event.is_set(): try: - conn, _ = self.server_socket.accept() + with self.lock: + server_socket = self.server_socket + if server_socket is None: + break + conn, _ = server_socket.accept() # Handle client in background thread threading.Thread( target=self._handle_client, args=(conn,), daemon=True @@ -86,7 +107,7 @@ class CallosumServer: if not self.stop_event.is_set(): logger.error(f"Accept error: {e}") finally: - self.server_socket.close() + self._close_server_socket() if self.socket_path.exists(): self.socket_path.unlink() @@ -218,6 +239,7 @@ class CallosumServer: def stop(self) -> None: """Stop the server and writer thread.""" self.stop_event.set() + self._close_server_socket() # Wait for writer thread to finish if self.writer_thread and self.writer_thread.is_alive(): diff --git a/tests/test_api_baselines.py b/tests/test_api_baselines.py index 81b0fdfc3..ef46aa2f6 100644 --- a/tests/test_api_baselines.py +++ b/tests/test_api_baselines.py @@ -26,10 +26,12 @@ from tests.verify_api import ( ) FREEZEGUN_IGNORE = [ + "_pytest", "librosa", "numba", "pandas", "pyarrow", + "pytest", "scipy", "sentencepiece", "sklearn", diff --git a/tests/test_cortex_client.py b/tests/test_cortex_client.py index 181087bcf..62ac0337b 100644 --- a/tests/test_cortex_client.py +++ b/tests/test_cortex_client.py @@ -578,7 +578,7 @@ def test_wait_for_agents_partial(callosum_server): def wait_thread(): result["completed"], result["timed_out"] = wait_for_uses( - [completing_agent, timeout_agent], timeout=2 + [completing_agent, timeout_agent], timeout=1 ) waiter = threading.Thread(target=wait_thread) diff --git a/tests/test_deferred_deletes.py b/tests/test_deferred_deletes.py index 63a95a925..d5be4441c 100644 --- a/tests/test_deferred_deletes.py +++ b/tests/test_deferred_deletes.py @@ -62,6 +62,7 @@ def test_cancel_commit_race_runs_at_most_once(): for _ in range(iterations): commit_count = 0 commit_lock = threading.Lock() + commit_event = threading.Event() start = threading.Event() cancel_results = [] @@ -69,8 +70,11 @@ def test_cancel_commit_race_runs_at_most_once(): nonlocal commit_count with commit_lock: commit_count += 1 + commit_event.set() deferred_id = deferred_deletes.schedule(commit, ttl_seconds=0.05) + with deferred_deletes._LOCK: + timer = deferred_deletes._TIMERS[deferred_id] def attempt_cancel(): start.wait() @@ -81,11 +85,16 @@ def test_cancel_commit_race_runs_at_most_once(): thread.start() start.set() - time.sleep(0.15) for thread in threads: thread.join() true_cancels = sum(cancel_results) + if true_cancels: + timer.join(timeout=0.1) + assert not timer.is_alive() + assert not commit_event.is_set() + else: + assert commit_event.wait(0.1) assert commit_count in (0, 1) assert not (true_cancels and commit_count) assert (true_cancels == 1 and commit_count == 0) or ( diff --git a/tests/test_importer_stall_timeout.py b/tests/test_importer_stall_timeout.py index 2258f86b1..3212b8092 100644 --- a/tests/test_importer_stall_timeout.py +++ b/tests/test_importer_stall_timeout.py @@ -53,8 +53,8 @@ def test_status_only_stream_stalls(caplog): failed_segments, completed_count = _wait_for_segments( message_queue, pending, - segment_timeout=1.0, - poll_timeout=0.05, + segment_timeout=0.15, + poll_timeout=0.02, ) elapsed = time.monotonic() - started_at @@ -79,8 +79,8 @@ def test_mixed_observed_and_status_stalls_remaining_segments(caplog): failed_segments, completed_count = _wait_for_segments( message_queue, pending, - segment_timeout=1.0, - poll_timeout=0.05, + segment_timeout=0.15, + poll_timeout=0.02, ) elapsed = time.monotonic() - started_at diff --git a/tests/test_readiness.py b/tests/test_readiness.py index 5fe6f4af3..ceec95584 100644 --- a/tests/test_readiness.py +++ b/tests/test_readiness.py @@ -92,7 +92,7 @@ def test_wait_ready_times_out_without_marker(tmp_path, monkeypatch): monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) _write_identity(tmp_path) - assert readiness.wait_ready(timeout=1.0) is None + assert readiness.wait_ready(timeout=0.1) is None def test_wait_ready_ignores_malformed_marker(tmp_path, monkeypatch, caplog): @@ -136,9 +136,10 @@ def test_wait_ready_rejects_reused_pid(tmp_path, monkeypatch): def test_wait_ready_observes_marker_written_mid_flight(tmp_path, monkeypatch): monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) + monkeypatch.setattr(readiness, "_POLL_INTERVAL_S", 0.02) _write_identity(tmp_path) timer = threading.Timer( - 0.1, readiness.signal_ready, kwargs={"payload": {"stage": "ready"}} + 0.02, readiness.signal_ready, kwargs={"payload": {"stage": "ready"}} ) try: