diff --git a/conftest.py b/conftest.py index 8194b08ea..858581c79 100644 --- a/conftest.py +++ b/conftest.py @@ -160,9 +160,11 @@ def _reset_local_slot_sizing_state(): local_server = sys.modules.get("solstone.think.providers.local_server") if local_server is not None: local_server.reset_parallel_slots_cache() + fanout_policy = sys.modules.get("solstone.think.providers.fanout_policy") + if fanout_policy is not None: + fanout_policy.reset_default_cap_log_state() thinking = sys.modules.get("solstone.think.thinking") if thinking is not None: - thinking.reset_default_cap_log_state() thinking.reset_dispatch_admission_state() _reset() diff --git a/solstone/observe/describe.py b/solstone/observe/describe.py index 1d31b2bfb..eb80db124 100644 --- a/solstone/observe/describe.py +++ b/solstone/observe/describe.py @@ -54,6 +54,7 @@ from solstone.think.callosum import callosum_send from solstone.think.journal_io import install_file from solstone.think.markdown import bound_extraction_markdown from solstone.think.prompts import load_prompt +from solstone.think.providers import fanout_policy from solstone.think.providers import state as provider_state from solstone.think.utils import ( day_from_path, @@ -688,7 +689,7 @@ class VideoProcessor: async def process_with_vision( self, - max_concurrent: int = 10, + max_concurrent: int, output_path: Optional[Path] = None, work_key: str | None = None, ) -> None: @@ -703,7 +704,7 @@ class VideoProcessor: Parameters ---------- max_concurrent : int - Maximum number of concurrent API requests (default: 10) + Maximum number of concurrent API requests. output_path : Optional[Path] Path to write JSONL output (when None, no output file is written) """ @@ -1257,8 +1258,8 @@ async def async_main(): "-j", "--jobs", type=int, - default=10, - help="Max concurrent vision API requests (default: 10)", + default=None, + help="Max concurrent vision API requests (default: provider policy)", ) parser.add_argument( "--frames-only", @@ -1319,8 +1320,13 @@ async def async_main(): output_qualified_frames(processor, qualified_frames) else: # New behavior: process with vision analysis + max_concurrent = ( + args.jobs + if args.jobs is not None + else fanout_policy.describe_per_proc_jobs(1) + ) await processor.process_with_vision( - max_concurrent=args.jobs, + max_concurrent=max_concurrent, output_path=output_path, work_key=work_key, ) diff --git a/solstone/observe/sense.py b/solstone/observe/sense.py index b195bcc91..3c9c30358 100644 --- a/solstone/observe/sense.py +++ b/solstone/observe/sense.py @@ -33,6 +33,7 @@ from solstone.observe.utils import ( from solstone.think import admission from solstone.think.callosum import CallosumConnection from solstone.think.processing import load_processing_settings +from solstone.think.providers import fanout_policy from solstone.think.runner import KILL_REAP_GRACE_S from solstone.think.runner import ManagedProcess as RunnerManagedProcess from solstone.think.utils import ( @@ -1374,8 +1375,19 @@ def main(): sensor.register(f"*{ext}", "transcribe", ["journal", "transcribe", "{file}"]) # Video files in segment directories + describe_configured = sensor._resolve_concurrency("describe") + describe_effective_procs = ( + max(args.jobs, describe_configured) if args.day else describe_configured + ) + describe_per_proc_jobs = fanout_policy.describe_per_proc_jobs( + describe_effective_procs + ) for ext in VIDEO_EXTENSIONS: - sensor.register(f"*{ext}", "describe", ["journal", "describe", "{file}"]) + sensor.register( + f"*{ext}", + "describe", + ["journal", "describe", "{file}", "-j", str(describe_per_proc_jobs)], + ) for ext in IMAGE_EXTENSIONS: sensor.register(f"*{ext}", "depict", ["journal", "depict", "{file}"]) diff --git a/solstone/think/providers/fanout_policy.py b/solstone/think/providers/fanout_policy.py new file mode 100644 index 000000000..6c2dcbfa6 --- /dev/null +++ b/solstone/think/providers/fanout_policy.py @@ -0,0 +1,122 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +"""Provider-aware fan-out sizing policy.""" + +from __future__ import annotations + +import logging +import os + +from solstone.think.models import is_local_provider_needed, resolve_provider +from solstone.think.providers.local_endpoint import resolve_local_endpoint +from solstone.think.providers.local_server import read_server_parallel_slots + +_DEFAULT_DESCRIBE_PER_PROC_JOBS = 10 +_CAP_LOGGED: set[str] = set() + + +def reset_default_cap_log_state() -> None: + """Clear the per-process record of which cap lines have been logged.""" + _CAP_LOGGED.clear() + + +def _segment_work_uses_local() -> bool: + """Return True when segment-pipeline work can resolve to the local provider.""" + return is_local_provider_needed() + + +def _describe_uses_local() -> bool: + """Return True when screen-describe resolves to the local provider.""" + provider, _ = resolve_provider("generate") + return provider == "local" + + +def _local_fanout_slots() -> int | None: + endpoint = resolve_local_endpoint() + if endpoint.is_bundled: + return read_server_parallel_slots() + return endpoint.parallel_slots + + +def cap_default_at_local_slots(formula: int, log_key: str) -> int: + """Clamp a CPU-derived default to the local provider's client-side slots.""" + slots = _local_fanout_slots() + if slots is None: + return formula + derived = min(formula, slots) + if derived < formula and log_key not in _CAP_LOGGED: + _CAP_LOGGED.add(log_key) + logging.info( + "%s capped provider=local slots=%d formula=%d derived=%d", + log_key, + slots, + formula, + derived, + ) + return derived + + +def default_segment_workers() -> int: + """Return the default segment-level worker count for repair mode.""" + cpu_count = os.cpu_count() or 2 + formula = max(1, min(8, cpu_count // 2)) + if not _segment_work_uses_local(): + return formula + return cap_default_at_local_slots(formula, "default_segment_workers") + + +def default_describe_jobs() -> int: + """Return the default screen-describe process count for repair mode.""" + formula = max(1, min(4, (os.cpu_count() or 2) // 4)) + if not _describe_uses_local(): + return formula + return cap_default_at_local_slots(formula, "default_describe_jobs") + + +def describe_per_proc_jobs(effective_procs: int) -> int: + """Return the per-process describe request concurrency for an effective process fan-out. + + For governed local providers, size per-process concurrency so + effective_procs * per_proc <= 2 * slots whenever integer flooring permits it. + If effective_procs already exceeds 2 * slots, per_proc floors at 1; this leaves + a residual product above the target, logs that residual once per key, and returns + 1 because per-invocation concurrency cannot be reduced further. + + This function sizes one describe invocation. It does not choose the number of + describe processes, enforce local admission, or account for work outside the + effective_procs supplied by the caller. Non-local providers and ungoverned + confidential BYO endpoints keep the historical per-process default of 10. + """ + if ( + not isinstance(effective_procs, int) + or isinstance(effective_procs, bool) + or effective_procs < 1 + ): + raise ValueError("effective_procs must be an integer >= 1") + + if not _describe_uses_local(): + return _DEFAULT_DESCRIBE_PER_PROC_JOBS + + slots = _local_fanout_slots() + if slots is None: + return _DEFAULT_DESCRIBE_PER_PROC_JOBS + + per_proc = max(1, (2 * slots) // effective_procs) + product = effective_procs * per_proc + limit = 2 * slots + if product > limit: + log_key = f"describe_per_proc_jobs_residual:{slots}:{effective_procs}" + if log_key not in _CAP_LOGGED: + _CAP_LOGGED.add(log_key) + logging.info( + "%s residual provider=local slots=%d effective_procs=%d " + "per_proc=%d product=%d limit=%d", + "describe_per_proc_jobs", + slots, + effective_procs, + per_proc, + product, + limit, + ) + return per_proc diff --git a/solstone/think/providers/local.py b/solstone/think/providers/local.py index 258513731..ace7d5323 100644 --- a/solstone/think/providers/local.py +++ b/solstone/think/providers/local.py @@ -26,6 +26,7 @@ from solstone.think.providers.local_endpoint import ( LOCAL_ENDPOINT_CONTRACT_COPY, LOCAL_ENDPOINT_UNREACHABLE_COPY, classify_byo_cogitate_error, + is_byo_capacity_error, is_byo_network_error, local_endpoint_reason_copy, redact_local_endpoint_credential, @@ -435,6 +436,8 @@ def _telemetry_record( def _classify_byo_generate_error(exc: BaseException) -> LocalProviderError: + if is_byo_capacity_error(exc): + return LocalCapacityExhausted() if is_byo_network_error(exc): return LocalProviderError( "local_endpoint_unreachable", diff --git a/solstone/think/providers/local_endpoint.py b/solstone/think/providers/local_endpoint.py index 560d3b7a2..16b136cfe 100644 --- a/solstone/think/providers/local_endpoint.py +++ b/solstone/think/providers/local_endpoint.py @@ -166,6 +166,16 @@ _BYO_NETWORK_EXC_NAMES = frozenset( ) +_BYO_CAPACITY_EXC_NAMES = frozenset( + { + "ReadTimeout", + "PoolTimeout", + "TimeoutException", + "WriteTimeout", + } +) + + def is_byo_network_error(exc: BaseException) -> bool: """True if the cause chain names a connection/timeout/network failure.""" @@ -173,6 +183,13 @@ def is_byo_network_error(exc: BaseException) -> bool: return bool(_BYO_NETWORK_EXC_NAMES & names) +def is_byo_capacity_error(exc: BaseException) -> bool: + """True if the cause chain names a serving-capacity timeout.""" + + names = {type(item).__name__ for item in _exception_chain(exc)} + return bool(_BYO_CAPACITY_EXC_NAMES & names) + + def classify_byo_cogitate_error(exc: BaseException) -> str | None: """Return a BYO local-endpoint reason code for known OpenHands/LiteLLM errors.""" @@ -305,6 +322,7 @@ __all__ = [ "LocalEndpoint", "classify_byo_cogitate_error", "confidential_provenance_block", + "is_byo_capacity_error", "is_byo_network_error", "local_endpoint_reason_copy", "normalize_local_endpoint_url", diff --git a/solstone/think/thinking.py b/solstone/think/thinking.py index 3cdcd15fa..5359af713 100644 --- a/solstone/think/thinking.py +++ b/solstone/think/thinking.py @@ -14,7 +14,6 @@ import argparse import fnmatch import json import logging -import os import subprocess import sys import threading @@ -58,7 +57,7 @@ from solstone.think.facets import ( load_segment_facets, ) from solstone.think.journal_io import atomic_replace -from solstone.think.models import is_local_provider_needed, resolve_provider +from solstone.think.models import is_local_provider_needed from solstone.think.pipeline_health import ( SEGMENT_FLOOR_TALENTS, DeterministicFailure, @@ -74,8 +73,7 @@ from solstone.think.pipeline_health import ( segment_fully_thought, segment_requires_processing, ) -from solstone.think.providers.local_endpoint import resolve_local_endpoint -from solstone.think.providers.local_server import read_server_parallel_slots +from solstone.think.providers import fanout_policy from solstone.think.runner import DEFAULT_TASK_MAX_RUNTIME, run_task from solstone.think.sense_splitter import ( write_change_detection, @@ -481,15 +479,9 @@ def _flush_batch_state_machines( ) -_CAP_LOGGED: set[str] = set() _LOCAL_PROVIDER_NEEDED: bool | None = None -def reset_default_cap_log_state() -> None: - """Clear the per-process record of which cap lines have been logged.""" - _CAP_LOGGED.clear() - - def reset_dispatch_admission_state() -> None: """Clear the per-process local-provider admission memo.""" global _LOCAL_PROVIDER_NEEDED @@ -505,67 +497,6 @@ def _dispatch_local_provider_needed() -> bool: return _LOCAL_PROVIDER_NEEDED -def _segment_work_uses_local() -> bool: - """Return True when segment-pipeline work can resolve to the local provider.""" - return is_local_provider_needed() - - -def _describe_uses_local() -> bool: - """Return True when screen-describe resolves to the local provider.""" - provider, _ = resolve_provider("generate") - return provider == "local" - - -def _local_fanout_slots() -> int | None: - endpoint = resolve_local_endpoint() - if endpoint.is_bundled: - return read_server_parallel_slots() - return endpoint.parallel_slots - - -def _cap_default_at_local_slots(formula: int, log_key: str) -> int: - """Clamp a CPU-derived default to the local provider's client-side slots.""" - slots = _local_fanout_slots() - if slots is None: - return formula - derived = min(formula, slots) - if derived < formula and log_key not in _CAP_LOGGED: - _CAP_LOGGED.add(log_key) - logging.info( - "%s capped provider=local slots=%d formula=%d derived=%d", - log_key, - slots, - formula, - derived, - ) - return derived - - -def _default_segment_workers() -> int: - """Return the default segment-level worker count for repair mode. - - Capped at the local provider's client-side slot count when any - segment-pipeline work can resolve to a governed local lane. - """ - cpu_count = os.cpu_count() or 2 - formula = max(1, min(8, cpu_count // 2)) - if not _segment_work_uses_local(): - return formula - return _cap_default_at_local_slots(formula, "default_segment_workers") - - -def _default_describe_jobs() -> int: - """Return the default screen-describe worker count for repair mode. - - Capped at the local provider's client-side slot count when the describe - path resolves to a governed local lane. - """ - formula = max(1, min(4, (os.cpu_count() or 2) // 4)) - if not _describe_uses_local(): - return formula - return _cap_default_at_local_slots(formula, "default_describe_jobs") - - def _select_segment_repair_targets( day: str, segments: list[dict], @@ -3817,7 +3748,7 @@ def dry_run( "Pre-phase: journal sense --day " + day + " -j " - + str(_default_describe_jobs()) + + str(fanout_policy.default_describe_jobs()) ) if not all_prompts: @@ -4226,7 +4157,9 @@ def main() -> None: ) if args.segments: - segment_workers = args.segment_workers or _default_segment_workers() + segment_workers = ( + args.segment_workers or fanout_policy.default_segment_workers() + ) if args.jobs == 0 and segment_workers > 1: parser.error( "--jobs 0 is incompatible with multi-worker --segments; " @@ -4335,7 +4268,9 @@ def main() -> None: sys.exit(1) selected = len(selected_segments) - segment_workers = args.segment_workers or _default_segment_workers() + segment_workers = ( + args.segment_workers or fanout_policy.default_segment_workers() + ) logging.info( "Segment repair targets for %s: selected=%d complete=%d " "raw_blocked=%d total=%d workers=%d jobs=%s", @@ -4556,7 +4491,7 @@ def main() -> None: "--day", day, "-j", - str(_default_describe_jobs()), + str(fanout_policy.default_describe_jobs()), ] if args.verbose: cmd.append("-v") diff --git a/tests/test_describe_capacity_retry.py b/tests/test_describe_capacity_retry.py new file mode 100644 index 000000000..4988942c1 --- /dev/null +++ b/tests/test_describe_capacity_retry.py @@ -0,0 +1,261 @@ +# SPDX-License-Identifier: AGPL-3.0-only +# Copyright (c) 2026 sol pbc + +from __future__ import annotations + +import io +import json +from pathlib import Path +from types import SimpleNamespace + +import pytest +from PIL import Image + +from solstone.observe import describe as describe_module + + +def _video_path(tmp_path: Path) -> Path: + segment_dir = tmp_path / "chronicle" / "20250101" / "default" / "143022_300" + segment_dir.mkdir(parents=True) + video_path = segment_dir / "screen.webm" + video_path.write_text("video", encoding="utf-8") + return video_path + + +def _png_bytes() -> bytes: + image_bytes = io.BytesIO() + Image.new("RGB", (8, 8), "white").save(image_bytes, format="PNG") + return image_bytes.getvalue() + + +def _frame(frame_id: int, frame_bytes: bytes) -> dict: + return { + "frame_id": frame_id, + "timestamp": float(frame_id), + "frame_bytes": frame_bytes, + "aruco": None, + } + + +def _processor(video_path: Path, frames: list[dict], monkeypatch) -> object: + processor = describe_module.VideoProcessor.__new__(describe_module.VideoProcessor) + processor.video_path = video_path + processor.first_hash = None + processor.last_hash = None + processor.qualified_count = len(frames) + processor.qualified_frames = [] + monkeypatch.setattr(processor, "process", lambda: frames) + return processor + + +def _jsonl_rows(path: Path) -> list[dict]: + return [ + json.loads(line) + for line in path.read_text(encoding="utf-8").splitlines() + if line + ] + + +def _assert_no_describe_temp(directory: Path) -> None: + names = [path.name for path in directory.iterdir()] + assert not any( + name.startswith(".describe_") or name.endswith(".tmp") for name in names + ) + + +def _capacity_error(): + from solstone.think.providers import local as local_provider + + inner = type("ReadTimeout", (Exception,), {})("capacity wait timed out") + outer = RuntimeError("outer") + outer.__cause__ = inner + return local_provider._classify_byo_generate_error(outer) + + +class RetryBatch: + mode = "phase1" + attempts: dict[tuple, int] = {} + + def __init__(self, max_concurrent=5, client=None): + self.max_concurrent = max_concurrent + self.client = client + self.pending_tasks = set() + self.queue = [] + + def create(self, **kwargs): + return SimpleNamespace( + **kwargs, + response=None, + error=None, + duration=0.01, + model_used="gemini-test", + provider=None, + reason_code=None, + reset_at_ms=None, + ) + + def add(self, request): + self.queue.append(request) + + def update(self, request, **kwargs): + for key, value in kwargs.items(): + setattr(request, key, value) + request.error = None + request.reason_code = None + self.add(request) + + async def drain_batch(self): + while self.queue: + pending = self.queue + self.queue = [] + for request in pending: + key = ( + request.request_type.value, + request.frame_id, + getattr(request, "extraction_category", None), + ) + attempt = self.attempts.get(key, 0) + self.attempts[key] = attempt + 1 + + should_fail = False + if self.mode == "phase1" and request.request_type.value == "describe": + should_fail = attempt == 0 + elif self.mode == "phase3" and request.request_type.value == "category": + should_fail = attempt == 0 + elif self.mode == "one_frame_exhausted" and request.frame_id == 1: + should_fail = request.request_type.value == "describe" + + if should_fail: + error = _capacity_error() + request.error = str(error) + request.reason_code = error.reason_code + else: + request.error = None + request.reason_code = None + if request.request_type.value == "describe": + request.response = json.dumps( + {"primary": "code", "secondary": "none", "overlap": True} + ) + else: + request.response = "extracted text" + yield request + + +def _install_fakes(monkeypatch, *, mode: str) -> None: + from solstone.think import batch as batch_module + from solstone.think import models + + RetryBatch.mode = mode + RetryBatch.attempts = {} + monkeypatch.setattr(batch_module, "Batch", RetryBatch) + monkeypatch.setattr(models, "resolve_provider", lambda _interface: ("google", "g")) + monkeypatch.setattr(describe_module, "callosum_send", lambda *args, **kwargs: True) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["phase1", "phase3"]) +async def test_capacity_class_describe_errors_retry_and_promote_output( + tmp_path, monkeypatch, mode +): + video_path = _video_path(tmp_path) + output_path = video_path.with_suffix(".jsonl") + frame_bytes = _png_bytes() + processor = _processor( + video_path, [_frame(1, frame_bytes), _frame(2, frame_bytes)], monkeypatch + ) + _install_fakes(monkeypatch, mode=mode) + monkeypatch.setattr( + describe_module, + "select_frames_for_extraction", + lambda *_args, **_kwargs: [1] if mode == "phase3" else [], + ) + + await processor.process_with_vision( + max_concurrent=1, + output_path=output_path, + work_key="20250101/143022_300/screen", + ) + + rows = _jsonl_rows(output_path) + assert rows[0]["_solstone_processing"]["state"] == "analyzed" + assert {row["frame_id"] for row in rows[1:]} == {1, 2} + if mode == "phase3": + enhanced = next(row for row in rows[1:] if row["frame_id"] == 1) + assert enhanced["enhanced"] is True + assert enhanced["content"]["code"] == "extracted text" + _assert_no_describe_temp(output_path.parent) + + +@pytest.mark.asyncio +async def test_capacity_class_exhausted_frame_with_successful_sibling_stays_analyzed( + tmp_path, monkeypatch +): + video_path = _video_path(tmp_path) + output_path = video_path.with_suffix(".jsonl") + frame_bytes = _png_bytes() + processor = _processor( + video_path, [_frame(1, frame_bytes), _frame(2, frame_bytes)], monkeypatch + ) + _install_fakes(monkeypatch, mode="one_frame_exhausted") + monkeypatch.setattr( + describe_module, + "select_frames_for_extraction", + lambda *_args, **_kwargs: [], + ) + + await processor.process_with_vision( + max_concurrent=1, + output_path=output_path, + work_key="20250101/143022_300/screen", + ) + + rows = _jsonl_rows(output_path) + assert rows[0]["_solstone_processing"]["state"] == "analyzed" + failed = next(row for row in rows[1:] if row["frame_id"] == 1) + succeeded = next(row for row in rows[1:] if row["frame_id"] == 2) + assert "error" in failed + assert succeeded["analysis"]["primary"] == "code" + _assert_no_describe_temp(output_path.parent) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("jobs", [1, 7]) +async def test_describe_explicit_jobs_wins_without_policy_resolution( + tmp_path, monkeypatch, jobs +): + video_path = _video_path(tmp_path) + observed = [] + + async def fake_process_with_vision( + self, + max_concurrent: int, + output_path: Path | None = None, + work_key: str | None = None, + ) -> None: + del self, output_path, work_key + observed.append(max_concurrent) + + def fail_policy(_effective_procs: int) -> int: + raise AssertionError("explicit -j must not resolve through policy") + + monkeypatch.setattr(describe_module, "require_solstone", lambda: None) + monkeypatch.setattr( + describe_module, "_preflight_provider_readiness", lambda *a, **k: None + ) + monkeypatch.setattr( + describe_module.VideoProcessor, + "process_with_vision", + fake_process_with_vision, + ) + monkeypatch.setattr( + "solstone.think.providers.fanout_policy.describe_per_proc_jobs", + fail_policy, + ) + monkeypatch.setattr( + "sys.argv", + ["journal describe", str(video_path), "-j", str(jobs)], + ) + + await describe_module.async_main() + + assert observed == [jobs] diff --git a/tests/test_local.py b/tests/test_local.py index 196c4c24c..228eb6945 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -1641,18 +1641,40 @@ def test_run_generate_byo_network_error_maps_to_unreachable(monkeypatch): @pytest.mark.parametrize( "exc_name", [ - "ConnectError", - "APIConnectionError", - "ConnectTimeout", "ReadTimeout", "PoolTimeout", + "WriteTimeout", "TimeoutException", + ], +) +def test_classify_byo_generate_error_capacity_names_are_non_blocking(exc_name): + provider = _provider() + from solstone.convey.provider_readiness import is_blocking_reason + + inner = type(exc_name, (Exception,), {})(f"{exc_name} failed") + exc = RuntimeError("outer") + exc.__cause__ = inner + + classified = provider._classify_byo_generate_error(exc) + + assert classified.reason_code == "local_capacity_exhausted" + assert is_blocking_reason(classified.reason_code) is False + + +@pytest.mark.parametrize( + "exc_name", + [ + "ConnectError", + "APIConnectionError", + "ConnectTimeout", "NetworkError", "RequestError", ], ) -def test_classify_byo_generate_error_uses_shared_network_predicate(exc_name): +def test_classify_byo_generate_error_unreachable_names_are_blocking(exc_name): provider = _provider() + from solstone.convey.provider_readiness import is_blocking_reason + inner = type(exc_name, (Exception,), {})(f"{exc_name} failed") exc = RuntimeError("outer") exc.__cause__ = inner @@ -1661,6 +1683,21 @@ def test_classify_byo_generate_error_uses_shared_network_predicate(exc_name): assert classified.reason_code == "local_endpoint_unreachable" assert str(classified) == provider.LOCAL_ENDPOINT_UNREACHABLE_COPY + assert is_blocking_reason(classified.reason_code) is True + + +def test_classify_byo_generate_error_capacity_wins_mixed_chain(): + provider = _provider() + + capacity = type("ReadTimeout", (Exception,), {})("read timeout") + unreachable = type("ConnectError", (Exception,), {})("connect failed") + outer = RuntimeError("outer") + outer.__cause__ = unreachable + unreachable.__cause__ = capacity + + classified = provider._classify_byo_generate_error(outer) + + assert classified.reason_code == "local_capacity_exhausted" def test_classify_byo_generate_error_500_stays_contract_failed(): diff --git a/tests/test_sense.py b/tests/test_sense.py index e4bb76f63..f00a10536 100644 --- a/tests/test_sense.py +++ b/tests/test_sense.py @@ -26,6 +26,7 @@ from solstone.think.processing import ( ProcessingSettings, TimeWindowSettings, ) +from solstone.think.providers import fanout_policy from solstone.think.runner import DailyLogWriter as ProcessLogWriter from solstone.think.runner import _format_log_line @@ -2277,6 +2278,64 @@ def test_main_rejects_invalid_stream_filter(tmp_path, monkeypatch): assert exc_info.value.code == 2 +def _registered_describe_commands(sensor: FileSensor) -> list[list[str]]: + return [ + command + for handler_name, command in sensor.handlers.values() + if handler_name == "describe" + ] + + +@pytest.mark.parametrize( + ("argv", "configured", "expected_effective"), + [ + (["journal sense"], 4, 4), + (["journal sense", "--day", "20250101", "-j", "2"], 4, 4), + (["journal sense", "--day", "20250101", "-j", "2"], 1, 2), + ], +) +def test_main_registers_describe_with_policy_per_proc_jobs( + tmp_path, monkeypatch, argv, configured, expected_effective +): + from solstone.observe import sense + + monkeypatch.setenv("SOLSTONE_JOURNAL", str(tmp_path)) + monkeypatch.setattr(sense, "require_solstone", lambda: None) + monkeypatch.setattr("sys.argv", argv) + effective_calls = [] + observed_commands = [] + + def fake_resolve_concurrency(self, handler_name): + del self + return configured if handler_name == "describe" else 1 + + def fake_per_proc(effective_procs): + effective_calls.append(effective_procs) + return 9 + + def record_watch(self): + observed_commands.extend(_registered_describe_commands(self)) + self.stop() + + def record_day(self, *args, **kwargs): + del args, kwargs + observed_commands.extend(_registered_describe_commands(self)) + self.stop() + + monkeypatch.setattr( + sense.FileSensor, "_resolve_concurrency", fake_resolve_concurrency + ) + monkeypatch.setattr(fanout_policy, "describe_per_proc_jobs", fake_per_proc) + monkeypatch.setattr(sense.FileSensor, "start", record_watch) + monkeypatch.setattr(sense.FileSensor, "process_day", record_day) + + sense.main() + + assert effective_calls == [expected_effective] + assert observed_commands + assert all(command[-2:] == ["-j", "9"] for command in observed_commands) + + def test_main_reprocess_screen_passes_stream_and_modality_filter(tmp_path, monkeypatch): from solstone.observe import sense @@ -2290,6 +2349,9 @@ def test_main_reprocess_screen_passes_stream_and_modality_filter(tmp_path, monke def register(self, *_args, **_kwargs): pass + def _resolve_concurrency(self, _handler_name): + return 1 + def process_day(self, *args, **kwargs): calls.append((args, kwargs)) @@ -2345,6 +2407,9 @@ def test_main_reprocess_all_keeps_modality_filter_unset(tmp_path, monkeypatch): def register(self, *_args, **_kwargs): pass + def _resolve_concurrency(self, _handler_name): + return 1 + def process_day(self, *args, **kwargs): calls.append((args, kwargs)) diff --git a/tests/test_think_segment_prephase.py b/tests/test_think_segment_prephase.py index cdca4bcc7..b382f8ee7 100644 --- a/tests/test_think_segment_prephase.py +++ b/tests/test_think_segment_prephase.py @@ -15,6 +15,7 @@ from unittest.mock import Mock import pytest +from solstone.think.providers import fanout_policy from tests.test_think_segment import _segment_configs, _write_sense_output DAY = "20240115" @@ -187,11 +188,12 @@ def test_daily_health_log_keeps_segment_events_out(journal_copy, monkeypatch): def _forbid_slot_discovery(monkeypatch, mod): """Assert the non-local path never probes the local server.""" + del mod def _unreachable() -> int: raise AssertionError("slot discovery must not run for non-local defaults") - monkeypatch.setattr(mod, "read_server_parallel_slots", _unreachable) + monkeypatch.setattr(fanout_policy, "read_server_parallel_slots", _unreachable) def _pin_describe_non_local(monkeypatch, mod): @@ -200,7 +202,7 @@ def _pin_describe_non_local(monkeypatch, mod): The fixture journal resolves observe.* to google, but these tests assert an exact -j value; stub the predicate so they do not silently depend on that. """ - monkeypatch.setattr(mod, "_describe_uses_local", lambda: False) + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: False) _forbid_slot_discovery(monkeypatch, mod) @@ -211,12 +213,12 @@ def test_sense_repair_prephase_uses_default_describe_jobs(journal_copy, monkeypa _pin_describe_non_local(monkeypatch, mod) - monkeypatch.setattr(mod.os, "cpu_count", lambda: 16) - assert mod._default_describe_jobs() == 4 - monkeypatch.setattr(mod.os, "cpu_count", lambda: 7) - assert mod._default_describe_jobs() == 1 - monkeypatch.setattr(mod.os, "cpu_count", lambda: 8) - assert mod._default_describe_jobs() == 2 + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: 16) + assert fanout_policy.default_describe_jobs() == 4 + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: 7) + assert fanout_policy.default_describe_jobs() == 1 + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: 8) + assert fanout_policy.default_describe_jobs() == 2 def fake_bounded(cmd, day, timeout=None): bounded_calls.append((cmd, day, timeout)) @@ -263,7 +265,7 @@ def test_daily_segment_prephase_timeout_is_nonfatal(journal_copy, monkeypatch): _patch_main_runtime(monkeypatch) _pin_describe_non_local(monkeypatch, mod) - monkeypatch.setattr(mod.os, "cpu_count", lambda: 16) + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: 16) monkeypatch.setattr(mod, "run_bounded_phase", fake_bounded) monkeypatch.setattr(mod, "run_command", fake_command) monkeypatch.setattr(mod, "run_queued_command", lambda cmd, day, timeout=600: True) @@ -1573,21 +1575,21 @@ def test_segments_default_local_slot_fallback_logs_once_across_call_sites( journal = tmp_path / "journal" monkeypatch.setenv("SOLSTONE_JOURNAL", str(journal)) _patch_segment_run(monkeypatch, think, journal) - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: True) - monkeypatch.setattr(think.os, "cpu_count", lambda: 12) + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: 12) monkeypatch.setattr("sys.argv", ["sol think", "--segments", "--day", DAY]) assert not (journal / "health" / "local.port").exists() derived = [] - original_default = think._default_segment_workers + original_default = fanout_policy.default_segment_workers def spy_default() -> int: value = original_default() derived.append(value) return value - monkeypatch.setattr(think, "_default_segment_workers", spy_default) + monkeypatch.setattr(fanout_policy, "default_segment_workers", spy_default) caplog.set_level(logging.INFO) with pytest.raises(SystemExit) as excinfo: @@ -1621,10 +1623,10 @@ def test_segments_explicit_segment_workers_bypasses_local_default_at_call_site( _patch_segment_run(monkeypatch, think, journal) # The derived default would be 1; the CLI asks for 6. - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: True) - monkeypatch.setattr(think, "read_server_parallel_slots", lambda: 1) - monkeypatch.setattr(think.os, "cpu_count", lambda: 12) - assert think._default_segment_workers() == 1 + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "read_server_parallel_slots", lambda: 1) + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: 12) + assert fanout_policy.default_segment_workers() == 1 observed = [] original_batch = think._run_segment_repair_batch diff --git a/tests/test_think_skip_talents.py b/tests/test_think_skip_talents.py index f728e7062..ed5041a13 100644 --- a/tests/test_think_skip_talents.py +++ b/tests/test_think_skip_talents.py @@ -8,6 +8,8 @@ from pathlib import Path import pytest +from solstone.think.providers import fanout_policy + DAY = "20240115" SEGMENT = "120000_300" STREAM = "default" @@ -359,7 +361,7 @@ def test_segments_batch_forwards_live_false( # A tmp journal carries no local artifacts, so the real predicate resolves # non-local and the default must never probe the server. - monkeypatch.setattr(think, "read_server_parallel_slots", _unreachable) + monkeypatch.setattr(fanout_policy, "read_server_parallel_slots", _unreachable) monkeypatch.setattr( "sys.argv", [ diff --git a/tests/test_thinking_defaults.py b/tests/test_thinking_defaults.py index 3efd9fc75..00aa29b55 100644 --- a/tests/test_thinking_defaults.py +++ b/tests/test_thinking_defaults.py @@ -16,7 +16,7 @@ from pathlib import Path import pytest from solstone.think import thinking as think -from solstone.think.providers import local_server +from solstone.think.providers import fanout_policy, local_server FLOOR_SLOTS = local_server._FLOOR_TIER.parallel_slots CAPABLE_SLOTS = local_server._CAPABLE_TIER.parallel_slots @@ -26,15 +26,15 @@ def _forbid_slot_discovery(monkeypatch: pytest.MonkeyPatch) -> None: def _unreachable() -> int: raise AssertionError("slot discovery must not run for non-local defaults") - monkeypatch.setattr(think, "read_server_parallel_slots", _unreachable) + monkeypatch.setattr(fanout_policy, "read_server_parallel_slots", _unreachable) def _pin_slots(monkeypatch: pytest.MonkeyPatch, slots: int) -> None: - monkeypatch.setattr(think, "read_server_parallel_slots", lambda: slots) + monkeypatch.setattr(fanout_policy, "read_server_parallel_slots", lambda: slots) def _pin_cpu_count(monkeypatch: pytest.MonkeyPatch, cpu_count: int) -> None: - monkeypatch.setattr(think.os, "cpu_count", lambda: cpu_count) + monkeypatch.setattr(fanout_policy.os, "cpu_count", lambda: cpu_count) def _write_journal_config( @@ -58,26 +58,26 @@ def _write_journal_config( def test_default_segment_workers_local_floor_slots_returns_one(monkeypatch): _pin_cpu_count(monkeypatch, 12) - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: True) _pin_slots(monkeypatch, FLOOR_SLOTS) - assert think._default_segment_workers() == FLOOR_SLOTS == 1 + assert fanout_policy.default_segment_workers() == FLOOR_SLOTS == 1 # --- AC2 ------------------------------------------------------------------- def test_default_segment_workers_local_capable_slots_respects_formula_min(monkeypatch): - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: True) _pin_slots(monkeypatch, CAPABLE_SLOTS) # Formula (8) exceeds slots (2): slots win. _pin_cpu_count(monkeypatch, 16) - assert think._default_segment_workers() == CAPABLE_SLOTS == 2 + assert fanout_policy.default_segment_workers() == CAPABLE_SLOTS == 2 # Formula (1) is below slots (2): the formula wins. min() holds both ways. _pin_cpu_count(monkeypatch, 2) - assert think._default_segment_workers() == 1 + assert fanout_policy.default_segment_workers() == 1 # --- AC3 ------------------------------------------------------------------- @@ -85,11 +85,11 @@ def test_default_segment_workers_local_capable_slots_respects_formula_min(monkey def test_default_segment_workers_nonlocal_cpu_formula(monkeypatch): _pin_cpu_count(monkeypatch, 12) - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: False) + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: False) _forbid_slot_discovery(monkeypatch) # Distinct from the local cases' 1 and 2. - assert think._default_segment_workers() == 6 + assert fanout_policy.default_segment_workers() == 6 # --- AC4 ------------------------------------------------------------------- @@ -108,7 +108,7 @@ def test_default_segment_workers_ignores_context_local_pin(monkeypatch, tmp_path _pin_cpu_count(monkeypatch, 12) _pin_slots(monkeypatch, FLOOR_SLOTS) - assert think._default_segment_workers() == 6 + assert fanout_policy.default_segment_workers() == 6 # --- AC6 ------------------------------------------------------------------- @@ -118,12 +118,12 @@ def test_default_segment_workers_cap_logs_once_with_provider_slots_formula_deriv monkeypatch, caplog ): _pin_cpu_count(monkeypatch, 12) - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: True) _pin_slots(monkeypatch, FLOOR_SLOTS) caplog.set_level(logging.INFO) - assert think._default_segment_workers() == 1 - assert think._default_segment_workers() == 1 + assert fanout_policy.default_segment_workers() == 1 + assert fanout_policy.default_segment_workers() == 1 lines = [r.getMessage() for r in caplog.records if "capped" in r.getMessage()] assert lines == [ @@ -135,11 +135,11 @@ def test_default_segment_workers_does_not_log_when_cap_changes_nothing( monkeypatch, caplog ): _pin_cpu_count(monkeypatch, 2) # formula == 1 == slots - monkeypatch.setattr(think, "_segment_work_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_segment_work_uses_local", lambda: True) _pin_slots(monkeypatch, FLOOR_SLOTS) caplog.set_level(logging.INFO) - assert think._default_segment_workers() == 1 + assert fanout_policy.default_segment_workers() == 1 assert not [r for r in caplog.records if "capped" in r.getMessage()] @@ -149,34 +149,34 @@ def test_default_segment_workers_does_not_log_when_cap_changes_nothing( def test_default_describe_jobs_local_caps_and_nonlocal_uses_formula(monkeypatch): _pin_cpu_count(monkeypatch, 16) - monkeypatch.setattr(think, "_describe_uses_local", lambda: False) + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: False) _forbid_slot_discovery(monkeypatch) - assert think._default_describe_jobs() == 4 + assert fanout_policy.default_describe_jobs() == 4 - monkeypatch.setattr(think, "_describe_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: True) _pin_slots(monkeypatch, FLOOR_SLOTS) # Distinct from the non-local 4. - assert think._default_describe_jobs() == 1 + assert fanout_policy.default_describe_jobs() == 1 def test_default_describe_jobs_capable_slots_cap(monkeypatch): _pin_cpu_count(monkeypatch, 16) - monkeypatch.setattr(think, "_describe_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: True) _pin_slots(monkeypatch, CAPABLE_SLOTS) - assert think._default_describe_jobs() == CAPABLE_SLOTS == 2 + assert fanout_policy.default_describe_jobs() == CAPABLE_SLOTS == 2 def test_default_describe_jobs_cap_logs_provider_slots_formula_derived( monkeypatch, caplog ): _pin_cpu_count(monkeypatch, 16) - monkeypatch.setattr(think, "_describe_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: True) _pin_slots(monkeypatch, FLOOR_SLOTS) caplog.set_level(logging.INFO) - assert think._default_describe_jobs() == 1 - assert think._default_describe_jobs() == 1 + assert fanout_policy.default_describe_jobs() == 1 + assert fanout_policy.default_describe_jobs() == 1 lines = [r.getMessage() for r in caplog.records if "capped" in r.getMessage()] assert lines == [ @@ -207,8 +207,8 @@ def test_segment_default_byo_endpoint_uses_configured_slot_cap_and_does_not_prob _forbid_slot_discovery(monkeypatch) assert think.is_local_provider_needed() is True - assert think._segment_work_uses_local() is True - assert think._default_segment_workers() == 3 + assert fanout_policy._segment_work_uses_local() is True + assert fanout_policy.default_segment_workers() == 3 def test_segment_default_confidential_endpoint_uses_cpu_formula_and_does_not_probe( @@ -232,7 +232,7 @@ def test_segment_default_confidential_endpoint_uses_cpu_formula_and_does_not_pro _forbid_slot_discovery(monkeypatch) assert think.is_local_provider_needed() is True - assert think._default_segment_workers() == 6 + assert fanout_policy.default_segment_workers() == 6 def test_describe_default_byo_endpoint_uses_configured_slot_cap_and_does_not_probe( @@ -253,8 +253,8 @@ def test_describe_default_byo_endpoint_uses_configured_slot_cap_and_does_not_pro _pin_cpu_count(monkeypatch, 16) _forbid_slot_discovery(monkeypatch) - assert think._describe_uses_local() is True - assert think._default_describe_jobs() == 2 + assert fanout_policy._describe_uses_local() is True + assert fanout_policy.default_describe_jobs() == 2 def test_describe_default_confidential_endpoint_uses_cpu_formula_and_does_not_probe( @@ -276,5 +276,65 @@ def test_describe_default_confidential_endpoint_uses_cpu_formula_and_does_not_pr _pin_cpu_count(monkeypatch, 16) _forbid_slot_discovery(monkeypatch) - assert think._describe_uses_local() is True - assert think._default_describe_jobs() == 4 + assert fanout_policy._describe_uses_local() is True + assert fanout_policy.default_describe_jobs() == 4 + + +@pytest.mark.parametrize( + ("mode", "configured", "args_jobs", "effective_procs"), + [ + ("watch", 1, None, 1), + ("watch", 4, None, 4), + ("day", 1, 1, 1), + ("day", 1, 2, 2), + ("day", 4, 1, 4), + ("day", 4, 2, 4), + ], +) +@pytest.mark.parametrize("slots", [1, 2, 4]) +def test_describe_per_proc_product_matrix_and_residual_log_once( + monkeypatch, caplog, mode, configured, args_jobs, effective_procs, slots +): + del mode, configured, args_jobs + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: True) + _pin_slots(monkeypatch, slots) + caplog.set_level(logging.INFO) + + per_proc = fanout_policy.describe_per_proc_jobs(effective_procs) + repeat = fanout_policy.describe_per_proc_jobs(effective_procs) + + assert repeat == per_proc + product = effective_procs * per_proc + assert product <= max(2 * slots, effective_procs) + + residual_lines = [ + record.getMessage() + for record in caplog.records + if "describe_per_proc_jobs residual" in record.getMessage() + ] + if effective_procs == 4 and slots == 1: + assert residual_lines == [ + "describe_per_proc_jobs residual provider=local slots=1 " + "effective_procs=4 per_proc=1 product=4 limit=2" + ] + else: + assert residual_lines == [] + + +def test_describe_per_proc_bundled_floor_slots_direct_call_returns_two(monkeypatch): + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: True) + _pin_slots(monkeypatch, FLOOR_SLOTS) + + assert fanout_policy.describe_per_proc_jobs(1) == 2 + + +def test_describe_per_proc_nonlocal_and_confidential_byo_keep_historical_default( + monkeypatch, +): + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: False) + _forbid_slot_discovery(monkeypatch) + assert fanout_policy.describe_per_proc_jobs(1) == 10 + + monkeypatch.setattr(fanout_policy, "_describe_uses_local", lambda: True) + monkeypatch.setattr(fanout_policy, "_local_fanout_slots", lambda: None) + assert fanout_policy.describe_per_proc_jobs(4) == 10