diff --git a/.github/assets/badge-pypi-dark.svg b/.github/assets/badge-pypi-dark.svg index a592454..86cc8ac 100644 --- a/.github/assets/badge-pypi-dark.svg +++ b/.github/assets/badge-pypi-dark.svg @@ -5,5 +5,5 @@ v0.2.1 + font-size="13" font-weight="700" fill="#ffffff">v0.2.2 \ No newline at end of file diff --git a/.github/assets/badge-pypi-light.svg b/.github/assets/badge-pypi-light.svg index b9096dd..ed6019a 100644 --- a/.github/assets/badge-pypi-light.svg +++ b/.github/assets/badge-pypi-light.svg @@ -5,5 +5,5 @@ v0.2.1 + font-size="13" font-weight="700" fill="#ffffff">v0.2.2 \ No newline at end of file diff --git a/.github/assets/wechat_group.jpeg b/.github/assets/wechat_group.jpeg index 820fb0d..3ac7a38 100644 Binary files a/.github/assets/wechat_group.jpeg and b/.github/assets/wechat_group.jpeg differ diff --git a/EvoScientist/channels/middleware.py b/EvoScientist/channels/middleware.py index 47de4fc..ba8b60e 100644 --- a/EvoScientist/channels/middleware.py +++ b/EvoScientist/channels/middleware.py @@ -75,11 +75,13 @@ class DedupCache: max_size: int = _DEDUP_MAX, trim_to: int = _DEDUP_TRIM, ttl_seconds: float = _DEDUP_TTL, + clock: Callable[[], float] | None = None, ) -> None: self._seen: OrderedDict[str, float] = OrderedDict() self._max = max_size self._trim = trim_to self._ttl = ttl_seconds + self._clock = clock or time.monotonic # ── public API ────────────────────────────────────────────────── @@ -93,15 +95,16 @@ class DedupCache: if not msg_id: return False - self._prune() + now = self._clock() + self._prune(now) if msg_id in self._seen: # LRU: refresh position and timestamp self._seen.move_to_end(msg_id) - self._seen[msg_id] = time.monotonic() + self._seen[msg_id] = now return True - self._seen[msg_id] = time.monotonic() + self._seen[msg_id] = now if len(self._seen) > self._max: while len(self._seen) > self._trim: self._seen.popitem(last=False) @@ -118,9 +121,9 @@ class DedupCache: # ── internal ──────────────────────────────────────────────────── - def _prune(self) -> None: + def _prune(self, now: float | None = None) -> None: """Remove entries older than *ttl_seconds*.""" - cutoff = time.monotonic() - self._ttl + cutoff = (self._clock() if now is None else now) - self._ttl # OrderedDict is insertion-ordered; oldest entries are first. while self._seen: _key, ts = next(iter(self._seen.items())) @@ -428,11 +431,13 @@ class DedupMiddleware(InboundMiddleware): max_size: int = 1000, trim_to: int = 500, ttl_seconds: float = 3600.0, + clock: Callable[[], float] | None = None, ) -> None: self._cache = DedupCache( max_size=max_size, trim_to=trim_to, ttl_seconds=ttl_seconds, + clock=clock, ) async def process_inbound( diff --git a/EvoScientist/config/onboard/steps.py b/EvoScientist/config/onboard/steps.py index 9ef7738..184e0bd 100644 --- a/EvoScientist/config/onboard/steps.py +++ b/EvoScientist/config/onboard/steps.py @@ -371,14 +371,31 @@ def _step_minimax_region(config: EvoScientistConfig) -> str: return _MINIMAX_REGIONS[region] -def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: - """Step 2a: Select Anthropic authentication mode (API key vs OAuth). +def _step_oauth_auth_mode( + config: EvoScientistConfig, + *, + provider_label: str, + ccproxy_provider: str, + config_attr: str, + prompt_login_label: str, + oauth_choice_label: str | None = None, + status_label: str | None = None, + question_label: str | None = None, +) -> str: + """Select API-key vs ccproxy OAuth authentication for a provider. Args: config: Current configuration. + provider_label: Provider display name for direct API-key access. + ccproxy_provider: ccproxy auth provider name. + config_attr: Config attribute storing this provider's auth mode. + prompt_login_label: Label used in "Log in to ..." prompts. + oauth_choice_label: Optional display label for the OAuth choice. + status_label: Optional display label for status messages. + question_label: Optional prompt label override. Returns: - Selected auth mode: "api_key", "oauth", or "auto". + Selected auth mode: "api_key" or "oauth". """ from ...ccproxy_manager import check_ccproxy_auth, is_ccproxy_available @@ -386,10 +403,14 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: from .prompter import BACK_SENTINEL, GoBack, install_navigation_keys + oauth_label = oauth_choice_label or f"{prompt_login_label} OAuth" + auth_status_label = status_label or oauth_label + auth_question_label = question_label or f"{provider_label} authentication mode" + choices = [ - Choice(title="API Key (direct Anthropic access)", value="api_key"), + Choice(title=f"API Key (direct {provider_label} access)", value="api_key"), Choice( - title="Claude Code OAuth (via ccproxy — no API key needed)" + title=f"{oauth_label} (via ccproxy — no API key needed)" + ( "" if ccproxy_available @@ -401,12 +422,12 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: Choice(title="← Back (re-pick provider)", value=BACK_SENTINEL), ] - current = config.anthropic_auth_mode + current = getattr(config, config_attr) if current not in ("api_key", "oauth"): current = "api_key" question = questionary.select( - "Authentication mode [Esc/← to go back]:", + f"{auth_question_label} [Esc/← to go back]:", choices=choices, default=current, style=WIZARD_STYLE, @@ -448,11 +469,9 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: if auth_mode == "oauth": _prompt_ccproxy_port(config) - # If OAuth selected, check auth status and offer login - if auth_mode in ("oauth", "auto"): - authed, msg = check_ccproxy_auth() + authed, msg = check_ccproxy_auth(ccproxy_provider) if authed: - console.print(f" [green]✓ OAuth: {msg}[/green]") + console.print(f" [green]✓ {auth_status_label}: {msg}[/green]") relogin = questionary.confirm( "Re-authenticate to refresh credentials?", default=False, @@ -462,11 +481,13 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: if relogin is None: raise KeyboardInterrupt() if relogin: - _run_ccproxy_login("claude_api", "OAuth") + _run_ccproxy_login(ccproxy_provider, auth_status_label) else: - console.print(f" [yellow]OAuth not authenticated: {msg}[/yellow]") + console.print( + f" [yellow]{auth_status_label} not authenticated: {msg}[/yellow]" + ) login = questionary.confirm( - "Log in to Claude now?", + f"Log in to {prompt_login_label} now?", default=True, style=CONFIRM_STYLE, qmark=QMARK, @@ -474,11 +495,32 @@ def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: if login is None: raise KeyboardInterrupt() if login: - _run_ccproxy_login("claude_api", "OAuth") + _run_ccproxy_login(ccproxy_provider, auth_status_label) return auth_mode +def _step_anthropic_auth_mode(config: EvoScientistConfig) -> str: + """Step 2a: Select Anthropic authentication mode (API key vs OAuth). + + Args: + config: Current configuration. + + Returns: + Selected auth mode: "api_key" or "oauth". + """ + return _step_oauth_auth_mode( + config, + provider_label="Anthropic", + ccproxy_provider="claude_api", + config_attr="anthropic_auth_mode", + prompt_login_label="Claude", + oauth_choice_label="Claude Code OAuth", + status_label="OAuth", + question_label="Authentication mode", + ) + + def _step_openai_auth_mode(config: EvoScientistConfig) -> str: """Step 2b: Select OpenAI authentication mode (API key vs Codex OAuth). @@ -488,101 +530,16 @@ def _step_openai_auth_mode(config: EvoScientistConfig) -> str: Returns: Selected auth mode: "api_key" or "oauth". """ - from ...ccproxy_manager import check_ccproxy_auth, is_ccproxy_available - - ccproxy_available = is_ccproxy_available() - - from .prompter import BACK_SENTINEL, GoBack, install_navigation_keys - - choices = [ - Choice(title="API Key (direct OpenAI access)", value="api_key"), - Choice( - title="Codex OAuth (via ccproxy — no API key needed)" - + ( - "" - if ccproxy_available - else " [requires: pip install evoscientist[oauth]]" - ), - value="oauth", - ), - questionary.Separator(), - Choice(title="← Back (re-pick provider)", value=BACK_SENTINEL), - ] - - current = config.openai_auth_mode - if current not in ("api_key", "oauth"): - current = "api_key" - - question = questionary.select( - "OpenAI authentication mode [Esc/← to go back]:", - choices=choices, - default=current, - style=WIZARD_STYLE, - qmark=QMARK, - use_indicator=True, + return _step_oauth_auth_mode( + config, + provider_label="OpenAI", + ccproxy_provider="codex", + config_attr="openai_auth_mode", + prompt_login_label="Codex", + oauth_choice_label="Codex OAuth", + status_label="Codex OAuth", + question_label="OpenAI authentication mode", ) - install_navigation_keys(question, with_back=True) - auth_mode = question.ask() - - if auth_mode is None: - raise KeyboardInterrupt() - if auth_mode == BACK_SENTINEL: - raise GoBack() - - if auth_mode == "oauth" and not ccproxy_available: - console.print(" [yellow]✗ ccproxy not installed[/yellow]") - console.print() - install = questionary.confirm( - 'Install ccproxy now? (pip install "evoscientist[oauth]")', - default=True, - style=WIZARD_STYLE, - qmark=f" {QMARK}", - ).ask() - if install is None: - raise KeyboardInterrupt() - if install: - console.print() - if _install_ccproxy(): - console.print(" [green]✓ ccproxy installed successfully.[/green]") - else: - console.print(" [yellow]Falling back to API key mode.[/yellow]") - return "api_key" - else: - console.print( - ' [dim]Skipped. Install manually: pip install "evoscientist[oauth]"[/dim]' - ) - return "api_key" - - # If OAuth selected, prompt for port and check auth status - if auth_mode == "oauth": - _prompt_ccproxy_port(config) - authed, msg = check_ccproxy_auth("codex") - if authed: - console.print(f" [green]✓ Codex OAuth: {msg}[/green]") - relogin = questionary.confirm( - "Re-authenticate to refresh credentials?", - default=False, - style=CONFIRM_STYLE, - qmark=QMARK, - ).ask() - if relogin is None: - raise KeyboardInterrupt() - if relogin: - _run_ccproxy_login("codex", "Codex OAuth") - else: - console.print(f" [yellow]Codex OAuth not authenticated: {msg}[/yellow]") - login = questionary.confirm( - "Log in to Codex now?", - default=True, - style=CONFIRM_STYLE, - qmark=QMARK, - ).ask() - if login is None: - raise KeyboardInterrupt() - if login: - _run_ccproxy_login("codex", "Codex OAuth") - - return auth_mode def _step_provider_api_key( diff --git a/EvoScientist/config/onboard/wizard.py b/EvoScientist/config/onboard/wizard.py index 21008dc..3994db3 100644 --- a/EvoScientist/config/onboard/wizard.py +++ b/EvoScientist/config/onboard/wizard.py @@ -129,6 +129,12 @@ _PROVIDER_KEY_ATTR = { "custom-anthropic": "custom_anthropic_api_key", } +_MINIMAX_GLOBAL_BASE_URL = "https://api.minimax.io/anthropic" +_CUSTOM_PROVIDER_BASE_URL = { + "custom-openai": ("custom_openai_base_url", "CUSTOM_OPENAI_BASE_URL"), + "custom-anthropic": ("custom_anthropic_base_url", "CUSTOM_ANTHROPIC_BASE_URL"), +} + def _autosave(config: EvoScientistConfig) -> None: """Persist current config to disk between phases. @@ -142,6 +148,201 @@ def _autosave(config: EvoScientistConfig) -> None: pass +def _configure_provider_base_url( + config: EvoScientistConfig, + provider: str, + *, + strict: bool, +) -> list[str]: + """Configure provider-specific base URL/region and return Ollama models.""" + if provider in _CUSTOM_PROVIDER_BASE_URL: + attr_name, env_name = _CUSTOM_PROVIDER_BASE_URL[provider] + current_base_url = getattr(config, attr_name) or os.environ.get(env_name, "") + if strict: + if not current_base_url: + raise RuntimeError( + f"--non-interactive: {provider} provider needs a base URL. " + f"Set the {env_name} env var or run without --non-interactive." + ) + setattr(config, attr_name, current_base_url) + else: + setattr( + config, + attr_name, + _step_base_url(config, current_value=current_base_url), + ) + elif provider == "minimax": + if strict: + config.minimax_base_url = ( + config.minimax_base_url or _MINIMAX_GLOBAL_BASE_URL + ) + else: + config.minimax_base_url = _step_minimax_region(config) + elif provider == "ollama": + if strict: + config.ollama_base_url = ( + config.ollama_base_url + or os.environ.get("OLLAMA_BASE_URL", "") + or "http://localhost:11434" + ) + else: + ollama_url, ollama_detected_models = _step_ollama_base_url(config) + config.ollama_base_url = ollama_url + return ollama_detected_models + return [] + + +def _configure_provider_auth_mode( + config: EvoScientistConfig, + provider: str, + *, + strict: bool, +) -> None: + """Configure Anthropic/OpenAI auth mode for the selected provider.""" + if provider == "anthropic": + if strict: + config.anthropic_auth_mode = "api_key" + else: + config.anthropic_auth_mode = _step_anthropic_auth_mode(config) + elif provider == "openai": + if strict: + config.openai_auth_mode = "api_key" + else: + config.openai_auth_mode = _step_openai_auth_mode(config) + + +def _active_llm_providers(config: EvoScientistConfig) -> set[str]: + """Return providers currently selected by the main and auxiliary models.""" + providers = {config.provider} + if config.auxiliary_provider: + providers.add(config.auxiliary_provider) + return providers + + +def _reconcile_oauth_modes(config: EvoScientistConfig) -> None: + """Clear OAuth flags for providers no selected model uses.""" + active_providers = _active_llm_providers(config) + if "anthropic" not in active_providers: + config.anthropic_auth_mode = "api_key" + if "openai" not in active_providers: + config.openai_auth_mode = "api_key" + + +def _provider_uses_oauth(config: EvoScientistConfig, provider: str) -> bool: + return (provider == "anthropic" and config.anthropic_auth_mode == "oauth") or ( + provider == "openai" and config.openai_auth_mode == "oauth" + ) + + +def _apply_preset_provider_api_key( + config: EvoScientistConfig, + provider: str, + preset_api_key: str, + *, + skip_validation: bool, +) -> None: + """Validate and store a CLI-supplied provider API key.""" + if not skip_validation: + from .helpers import _provider_key_info + + _info = _provider_key_info(config, provider) + validate_fn = _info[2] if _info else None + if validate_fn is not None: + console.print(" [dim]Validating preset API key...[/dim]", end="") + valid, msg = validate_fn(preset_api_key) + if valid: + console.print(f"\r [green]✓ {msg}[/green] ") + else: + console.print(f"\r [red]✗ {msg}[/red] ") + raise RuntimeError( + f"--api-key rejected by {provider} validator: {msg}. " + "Pass --skip-validation to override." + ) + + key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key") + setattr(config, key_attr, preset_api_key) + console.print( + f" [green]✓ API key: ***{preset_api_key[-4:]}[/green] [dim](--api-key)[/dim]" + ) + + +def _configure_provider_api_key( + config: EvoScientistConfig, + provider: str, + *, + skip_validation: bool, + preset_api_key: str | None = None, + require_api_key=None, +) -> None: + """Configure provider API key unless the provider does not need one.""" + if provider == "ollama" or _provider_uses_oauth(config, provider): + return + + key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key") + if preset_api_key is not None: + _apply_preset_provider_api_key( + config, + provider, + preset_api_key, + skip_validation=skip_validation, + ) + return + + if require_api_key is not None: + require_api_key() + new_key = _step_provider_api_key(config, provider, skip_validation) + if new_key is not None: + setattr(config, key_attr, new_key) + elif not getattr(config, key_attr): + _print_step_skipped("API Key", "not set") + + +def _provider_connection_configured(config: EvoScientistConfig, provider: str) -> bool: + """Return True when provider-level setup can be safely reused.""" + if provider == "ollama": + return bool(config.ollama_base_url) + if provider == "custom-openai" and not config.custom_openai_base_url: + return False + if provider == "custom-anthropic" and not config.custom_anthropic_base_url: + return False + if provider == "minimax" and not config.minimax_base_url: + return False + if _provider_uses_oauth(config, provider): + return True + key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key") + return bool(getattr(config, key_attr)) + + +def _configure_provider_connection( + config: EvoScientistConfig, + provider: str, + *, + strict: bool, + skip_validation: bool, + preset_api_key: str | None = None, + require_api_key=None, +) -> list[str]: + """Configure provider base URL/region, auth mode, and API key.""" + ollama_detected_models = _configure_provider_base_url( + config, + provider, + strict=strict, + ) + _configure_provider_auth_mode( + config, + provider, + strict=strict, + ) + _configure_provider_api_key( + config, + provider, + skip_validation=skip_validation, + preset_api_key=preset_api_key, + require_api_key=require_api_key, + ) + return ollama_detected_models + + # Sections offered in Keep/Modify/Reset → which step labels they enable. _SECTION_LABELS: list[tuple[str, str]] = [ ("ui", "UI backend"), @@ -479,102 +680,17 @@ def run_onboard( provider = _step_provider(config) config.provider = provider - # Step 2a: Base URL (custom-openai, custom-anthropic, - # minimax, ollama). In strict non-interactive mode we - # never call the interactive _step_base_url / - # _step_minimax_region / _step_ollama_base_url helpers — - # fall back to the existing config value or the - # CUSTOM_*_BASE_URL / OLLAMA_BASE_URL env var instead. - # If neither is set for a provider that needs it, raise - # so the user sees the same "missing required answer" - # error as for other required prompts. - if provider == "custom-openai": - current_base_url = ( - config.custom_openai_base_url - or os.environ.get("CUSTOM_OPENAI_BASE_URL", "") - ) - if strict: - if not current_base_url: - raise RuntimeError( - "--non-interactive: custom-openai provider " - "needs a base URL. Set the " - "CUSTOM_OPENAI_BASE_URL env var or run " - "without --non-interactive." - ) - config.custom_openai_base_url = current_base_url - else: - config.custom_openai_base_url = _step_base_url( - config, current_value=current_base_url - ) - elif provider == "custom-anthropic": - current_base_url = ( - config.custom_anthropic_base_url - or os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", "") - ) - if strict: - if not current_base_url: - raise RuntimeError( - "--non-interactive: custom-anthropic " - "provider needs a base URL. Set the " - "CUSTOM_ANTHROPIC_BASE_URL env var or run " - "without --non-interactive." - ) - config.custom_anthropic_base_url = current_base_url - else: - config.custom_anthropic_base_url = _step_base_url( - config, current_value=current_base_url - ) - elif provider == "minimax": - if strict: - # MiniMax has 2 region URLs; default to whatever - # is already in config, else the Global endpoint. - config.minimax_base_url = ( - config.minimax_base_url - or "https://api.minimax.io/anthropic" - ) - else: - config.minimax_base_url = _step_minimax_region(config) - elif provider == "ollama": - if strict: - # Ollama: existing config value > env var > - # localhost default. Skip the live connection - # validation under strict — model discovery - # happens at runtime anyway. - config.ollama_base_url = ( - config.ollama_base_url - or os.environ.get("OLLAMA_BASE_URL", "") - or "http://localhost:11434" - ) - # ollama_detected_models stays [] — model picker - # will fall back to free-text or the preset. - else: - ollama_url, ollama_detected_models = _step_ollama_base_url( - config - ) - config.ollama_base_url = ollama_url - - # Step 2b: Auth mode (Anthropic or OpenAI — API key vs OAuth). - # In strict non-interactive mode we assume "api_key". - # The prompt offers a `← Back` choice that raises GoBack so - # the user can re-pick the provider without exiting the wizard. try: - if provider == "anthropic": - if strict: - config.anthropic_auth_mode = "api_key" - else: - config.anthropic_auth_mode = _step_anthropic_auth_mode( - config - ) - elif provider == "openai": - if strict: - config.openai_auth_mode = "api_key" - else: - config.openai_auth_mode = _step_openai_auth_mode(config) - else: - # Non-Anthropic/OpenAI provider: reset OAuth modes to - # avoid stale oauth config triggering ccproxy at startup. - config.anthropic_auth_mode = "api_key" - config.openai_auth_mode = "api_key" + ollama_detected_models = _configure_provider_connection( + config, + provider, + strict=strict, + skip_validation=skip_validation, + preset_api_key=_preset("api_key"), + require_api_key=lambda provider=provider: _require( + "api_key", f"{provider} API key" + ), + ) except GoBack: # User picked "← Back" — restore config to its state at the # top of this iteration (drops any base_url / region / @@ -594,60 +710,9 @@ def run_onboard( ollama_detected_models = [] console.print(" [dim]↩ Returning to provider selection.[/dim]") continue - break # auth_mode succeeded — exit sub-loop + break # Provider setup succeeded — exit sub-loop - # Step 2c: Provider API Key (skip for Ollama and pure OAuth) - _skip_api_key = ( - provider == "ollama" - or ( - provider == "anthropic" - and config.anthropic_auth_mode == "oauth" - ) - or (provider == "openai" and config.openai_auth_mode == "oauth") - ) - if not _skip_api_key: - key_attr = _PROVIDER_KEY_ATTR.get(provider, "openai_api_key") - preset_api_key = _preset("api_key") - if preset_api_key is not None: - # Validate the preset key against the same validator - # the interactive path uses, unless --skip-validation - # was passed. Interactive flow shows a "Save anyway?" - # confirm on failure; the non-interactive path has no - # way to ask, so a failed validation is fatal. - if not skip_validation: - from .helpers import _provider_key_info - - _info = _provider_key_info(config, provider) - validate_fn = _info[2] if _info else None - if validate_fn is not None: - console.print( - " [dim]Validating preset API key...[/dim]", - end="", - ) - valid, msg = validate_fn(preset_api_key) - if valid: - console.print(f"\r [green]✓ {msg}[/green] ") - else: - console.print(f"\r [red]✗ {msg}[/red] ") - raise RuntimeError( - f"--api-key rejected by {provider} " - f"validator: {msg}. Pass " - "--skip-validation to override." - ) - setattr(config, key_attr, preset_api_key) - console.print( - f" [green]✓ API key: ***{preset_api_key[-4:]}[/green]" - " [dim](--api-key)[/dim]" - ) - else: - _require("api_key", f"{provider} API key") - new_key = _step_provider_api_key( - config, provider, skip_validation - ) - if new_key is not None: - setattr(config, key_attr, new_key) - elif not getattr(config, key_attr): - _print_step_skipped("API Key", "not set") + _reconcile_oauth_modes(config) _autosave(config) else: # Provider section skipped — keep prior provider value to drive @@ -680,44 +745,55 @@ def run_onboard( "kept current" if config.auxiliary_model else "not set", ) elif _step_auxiliary_enable(config): - # Assemble: pick provider -> base URL (custom) -> key -> model, - # mirroring the main flow's order. Keys/base URLs are stored - # per provider, so when the auxiliary provider matches the main - # one they're already set and the user just keeps them (Enter). - # Ollama needs no key. Re-runs default to the saved auxiliary - # provider/model rather than the main ones. - aux_provider = _step_provider( - config, - label="co-pilot", - default_value=config.auxiliary_provider, - ) - config.auxiliary_provider = aux_provider - if aux_provider == "custom-openai": - config.custom_openai_base_url = _step_base_url( + from .prompter import GoBack + + aux_ollama_detected_models: list[str] = [] + while True: + loop_snapshot = copy.deepcopy(config) + aux_provider = _step_provider( config, - current_value=config.custom_openai_base_url - or os.environ.get("CUSTOM_OPENAI_BASE_URL", ""), + label="co-pilot", + default_value=config.auxiliary_provider, ) - elif aux_provider == "custom-anthropic": - config.custom_anthropic_base_url = _step_base_url( - config, - current_value=config.custom_anthropic_base_url - or os.environ.get("CUSTOM_ANTHROPIC_BASE_URL", ""), - ) - elif aux_provider == "minimax": - config.minimax_base_url = _step_minimax_region(config) - if aux_provider != "ollama": - aux_key_attr = _PROVIDER_KEY_ATTR.get( - aux_provider, "openai_api_key" - ) - new_aux_key = _step_provider_api_key( - config, aux_provider, skip_validation - ) - if new_aux_key is not None: - setattr(config, aux_key_attr, new_aux_key) + config.auxiliary_provider = aux_provider + if ( + aux_provider == config.provider + and _provider_connection_configured(config, aux_provider) + ): + if aux_provider == "ollama": + aux_ollama_detected_models = ollama_detected_models + _print_step_skipped( + "Co-pilot credentials", + "reusing main provider settings", + ) + else: + try: + aux_ollama_detected_models = ( + _configure_provider_connection( + config, + aux_provider, + strict=False, + skip_validation=skip_validation, + ) + ) + except GoBack: + for field_name in vars(loop_snapshot): + setattr( + config, + field_name, + getattr(loop_snapshot, field_name), + ) + aux_ollama_detected_models = [] + console.print( + " [dim]↩ Returning to co-pilot provider " + "selection.[/dim]" + ) + continue + break config.auxiliary_model = _step_model( config, aux_provider, + ollama_detected_models=aux_ollama_detected_models, label="co-pilot", default_value=config.auxiliary_model, ) @@ -725,6 +801,7 @@ def run_onboard( # Skip: single driver — clear any prior auxiliary config. config.auxiliary_provider = "" config.auxiliary_model = "" + _reconcile_oauth_modes(config) _autosave(config) if "tavily" in sections_to_run: diff --git a/EvoScientist/config/settings.py b/EvoScientist/config/settings.py index fb0088e..0793a9c 100644 --- a/EvoScientist/config/settings.py +++ b/EvoScientist/config/settings.py @@ -548,6 +548,10 @@ def save_config(config: EvoScientistConfig) -> None: """ config_path = get_config_path() config_path.parent.mkdir(parents=True, exist_ok=True) + try: + config_path.parent.chmod(0o700) + except OSError: + pass data = _config_to_dict(config) @@ -560,6 +564,10 @@ def save_config(config: EvoScientistConfig) -> None: sort_keys=False, allow_unicode=True, ) + try: + config_path.chmod(0o600) + except OSError: + pass def reset_config() -> None: diff --git a/EvoScientist/langgraph_dev/manager.py b/EvoScientist/langgraph_dev/manager.py index b7539a4..69e5068 100644 --- a/EvoScientist/langgraph_dev/manager.py +++ b/EvoScientist/langgraph_dev/manager.py @@ -306,14 +306,19 @@ def is_async_subagents_available() -> bool: def _langgraph_exe() -> str | None: """Return the path to the langgraph CLI binary, or None if not found.""" + import sys + + executable_dir = os.path.dirname(sys.executable) + candidate_names = ( + ["langgraph.exe", "langgraph"] if os.name == "nt" else ["langgraph"] + ) + for candidate_name in candidate_names: + candidate = os.path.join(executable_dir, candidate_name) + if os.path.isfile(candidate) and os.access(candidate, os.X_OK): + return candidate found = shutil.which("langgraph") if found: return found - import sys as _sys - - candidate = os.path.join(os.path.dirname(_sys.executable), "langgraph") - if os.path.isfile(candidate) and os.access(candidate, os.X_OK): - return candidate return None @@ -709,6 +714,7 @@ def start_langgraph_dev( sub_env["EVOSCIENTIST_DEPLOY_MODE"] = "full" if deploy_mode else "stripped" try: + logger.info("Starting langgraph dev with CLI: %s", exe) proc = subprocess.Popen( [ exe, diff --git a/EvoScientist/llm/context_window.py b/EvoScientist/llm/context_window.py index e1f3c95..c8bca7c 100644 --- a/EvoScientist/llm/context_window.py +++ b/EvoScientist/llm/context_window.py @@ -20,9 +20,9 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = { # Qwen 3.7 closed-source tiers — Max flagship and Plus (1M). "qwen3.7-max": 1_000_000, "qwen3.7-plus": 1_000_000, - # xAI Grok — per-model windows (build-0.1: 256K, 4.3: 1M). + # xAI Grok — per-model windows (build-0.1: 256K, 4.5: 500K). "grok-build-0.1": 256_000, - "grok-4.3": 1_000_000, + "grok-4.5": 500_000, # Claude Haiku 4.5 — exception to the ``claude-`` family (200K, not 1M). "claude-haiku-4-5": 200_000, # MiniMax M3 — 1M context (M2.x variants stay at provider default ~204K). @@ -32,6 +32,8 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = { # Zhipu GLM-5.2 — 1M context, an exception to the ``glm-5`` family (203K). # Matches OpenRouter ``z-ai/glm-5.2`` via split('/')[-1]. "glm-5.2": 1_000_000, + # Tencent Hunyuan HY3 — 262K context (OpenRouter ``tencent/hy3``). + "hy3": 262_000, } # Family-level fallbacks: tried only after exact-name lookup misses. @@ -40,6 +42,8 @@ _KNOWN_MODEL_CONTEXT_WINDOWS: dict[str, int] = { _KNOWN_MODEL_FAMILIES: list[tuple[str, int]] = [ # All Claude — 1M via the ``context-1m-2025-08-07`` beta header. ("claude-", 1_000_000), + # OpenAI GPT-5.6 family — sol, terra, luna variants + ("gpt-5.6", 1_050_000), # OpenAI GPT-5.5 family — base, pro, future variants ("gpt-5.5", 1_050_000), # Google Gemini 3.x family — flash, flash-lite, pro (1.05M). Excludes 2.5. diff --git a/EvoScientist/llm/models.py b/EvoScientist/llm/models.py index c63735e..efa61d2 100644 --- a/EvoScientist/llm/models.py +++ b/EvoScientist/llm/models.py @@ -89,6 +89,9 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [ ("claude-sonnet-4-6", "claude-sonnet-4-6", "anthropic"), ("claude-haiku-4-5", "claude-haiku-4-5", "anthropic"), # OpenAI + ("gpt-5.6-sol", "gpt-5.6-sol", "openai"), + ("gpt-5.6-terra", "gpt-5.6-terra", "openai"), + ("gpt-5.6-luna", "gpt-5.6-luna", "openai"), ("gpt-5.5-pro", "gpt-5.5-pro", "openai"), ("gpt-5.5", "gpt-5.5", "openai"), ("gpt-5.4", "gpt-5.4", "openai"), @@ -145,6 +148,9 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [ ("claude-opus-4.8-fast", "anthropic/claude-opus-4.8-fast", "openrouter"), ("claude-sonnet-5", "anthropic/claude-sonnet-5", "openrouter"), ("claude-sonnet-4.6", "anthropic/claude-sonnet-4.6", "openrouter"), + ("gpt-5.6-sol", "openai/gpt-5.6-sol", "openrouter"), + ("gpt-5.6-terra", "openai/gpt-5.6-terra", "openrouter"), + ("gpt-5.6-luna", "openai/gpt-5.6-luna", "openrouter"), ("gpt-5.5-pro", "openai/gpt-5.5-pro", "openrouter"), ("gpt-5.5", "openai/gpt-5.5", "openrouter"), ("gpt-5.4", "openai/gpt-5.4", "openrouter"), @@ -159,7 +165,8 @@ _MODEL_ENTRIES: list[tuple[str, str, str]] = [ ("mimo-v2.5-pro", "xiaomi/mimo-v2.5-pro", "openrouter"), ("mimo-v2.5", "xiaomi/mimo-v2.5", "openrouter"), ("grok-build-0.1", "x-ai/grok-build-0.1", "openrouter"), - ("grok-4.3", "x-ai/grok-4.3", "openrouter"), + ("grok-4.5", "x-ai/grok-4.5", "openrouter"), + ("hy3", "tencent/hy3", "openrouter"), ("qwen3.7-max", "qwen/qwen3.7-max", "openrouter"), ("qwen3.7-plus", "qwen/qwen3.7-plus", "openrouter"), ("qwen3.6-flash", "qwen/qwen3.6-flash", "openrouter"), diff --git a/README.md b/README.md index d9f9910..11a1884 100644 --- a/README.md +++ b/README.md @@ -179,6 +179,7 @@ only a transport event, not proof that the LangGraph run has finished.
📦 Release Highlights — version changelog +- **[11 Jul 2026]** **[v0.2.2](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.2)** — New models selectable in onboarding and `/model`: GPT-5.6 (sol, terra, luna) for OpenAI and OpenRouter, plus Grok 4.5 and Tencent Hunyuan HY3 on OpenRouter; tighter config-file permissions and a reworked onboarding OAuth flow for auxiliary models. - **[07 Jul 2026]** **[v0.2.2](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.2)** — WebUI streaming resilience hotfix: SSE disconnects no longer imply run completion; the composer stays on **Stop** until backend thread state is terminal; final-answer checkpoint recovery fills dropped response tails; tool-selection JSON payloads are filtered from live transcripts; `EVOSCIENTIST_WEBUI_PACKAGE` lets the launcher run a local patched WebUI package for validation before npm publication. - **[05 Jul 2026]** **[v0.2.1](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.1)** — AutoSkills: EvoMemory drafts reusable skills from its own observation clusters for you to review via `/autoskills`; a new `--output-format stream-json` for headless / SDK clients; richer slash-command completions; Windows UTF-8 config reads; a TUI welcome-banner fix; langchain-openrouter 0.2.5. - **[26 Jun 2026]** **[v0.2.0](https://github.com/EvoScientist/EvoScientist/releases/tag/v0.2.0)** — Scheduled tasks: cron-style recurring runs via `/schedule` or natural language, run unattended with shell-access gating; self-linking memory that connects observations into a knowledge graph (complements / contradicts / supersedes); a read-only `GET /api/models` endpoint for the WebUI model picker. @@ -463,10 +464,7 @@ EvoSci deploy # standalone LangGraph server for external UIs EvoSci -p "query" --output-format stream-json --auto-mode # JSONL event stream on stdout (for programmatic clients) ``` -`--output-format stream-json` makes a single-shot (`-p`) run emit its native -events as line-delimited JSON on stdout (one object per line), with all human -output on stderr — the integration surface for headless clients (e.g. an agent -runtime). See [docs/stream-json.md](docs/stream-json.md) for the event schema. +`--output-format stream-json` makes a single-shot (`-p`) run emit its native events as line-delimited JSON on stdout (one object per line), with all human output on stderr — the integration surface for headless clients (e.g. an agent runtime). See [docs/guides/stream-json.md](docs/guides/stream-json.md) for the event schema.
diff --git a/docs/README.md b/docs/README.md index b40b6ab..e241fed 100644 --- a/docs/README.md +++ b/docs/README.md @@ -11,6 +11,11 @@ |------------------------------------------------------------|---------------------------------------------------------------------------------| | [macOS 24/7 Deployment](https://github.com/EvoScientist/EvoScientist/blob/main/docs/recipes/deployment-macos-24h.md#running-evoscientist-247-on-macos-telegram-bot--stt--ccproxy) | Run EvoScientist as an always-on service on macOS with OAuth + Telegram + STT | + +| Guide | Description | +|------------------------------------------------------------|---------------------------------------------------------------------------------| +| [`stream-json` output protocol](https://github.com/EvoScientist/EvoScientist/blob/main/docs/guides/stream-json.md#stream-json-output-protocol) | Line-delimited JSON event stream (`--output-format stream-json`) for driving EvoScientist headlessly from SDK / programmatic clients | + ## Contributing a Recipe See the [Contributing Guide](../CONTRIBUTING.md) for general guidelines. When adding a new recipe: diff --git a/docs/stream-json.md b/docs/guides/stream-json.md similarity index 100% rename from docs/stream-json.md rename to docs/guides/stream-json.md diff --git a/pyproject.toml b/pyproject.toml index e73a31b..ccd4890 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -48,6 +48,7 @@ dependencies = [ [dependency-groups] dev = [ "pytest>=8.0", + "pytest-asyncio>=1.0", "pytest-cov>=5.0", "pytest-timeout>=2.4", "ruff>=0.5", @@ -58,6 +59,7 @@ dev = [ [project.optional-dependencies] dev = [ "pytest>=8.0", + "pytest-asyncio>=1.0", "pytest-cov>=5.0", "pytest-timeout>=2.4", "ruff>=0.5", @@ -117,6 +119,8 @@ EvoScientist = [ [tool.pytest.ini_options] testpaths = ["tests"] +asyncio_mode = "auto" +asyncio_default_fixture_loop_scope = "function" filterwarnings = [ "ignore::UserWarning:langchain_nvidia_ai_endpoints", ] diff --git a/tests/conftest.py b/tests/conftest.py index 4fa7a29..a6458d7 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,36 +1,8 @@ """Shared fixtures for EvoScientist tests.""" -import asyncio - import pytest -def run_async(coro): - """Run an async coroutine safely, cancelling pending tasks before closing. - - This prevents 'Event loop is closed' errors from asyncio.Queue cleanup - when tasks are still waiting on Queue.get() at teardown time. - """ - loop = asyncio.new_event_loop() - try: - return loop.run_until_complete(coro) - finally: - # Cancel all pending tasks so Queue getters don't raise on close - pending = asyncio.all_tasks(loop) - for task in pending: - task.cancel() - if pending: - loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True)) - loop.run_until_complete(loop.shutdown_asyncgens()) - loop.close() - - -@pytest.fixture(name="run_async") -def run_async_fixture(): - """Pytest fixture that exposes run_async as a callable for test functions.""" - return run_async - - @pytest.fixture(autouse=True) def _reset_tool_selection_state(): """Isolate the process-global tool-selection state around every test. diff --git a/tests/stream_v3_fakes.py b/tests/stream_v3_fakes.py index d518caa..a3e747f 100644 --- a/tests/stream_v3_fakes.py +++ b/tests/stream_v3_fakes.py @@ -9,7 +9,6 @@ from typing import Any from unittest.mock import MagicMock from EvoScientist.stream.events import stream_agent_events -from tests.conftest import run_async async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]: @@ -17,24 +16,20 @@ async def async_iter(items: Iterable[Any]) -> AsyncIterator[Any]: yield item -def collect_events( +async def collect_events( agent, message: str = "hi", thread_id: str = "t1", ): - """Collect stream_agent_events output for synchronous tests.""" - - async def _run(): - events = [] - async for ev in stream_agent_events( - agent, - message, - thread_id, - ): - events.append(ev) - return events - - return run_async(_run()) + """Collect stream_agent_events output for tests.""" + events = [] + async for ev in stream_agent_events( + agent, + message, + thread_id, + ): + events.append(ev) + return events def protocol_event( diff --git a/tests/test_additional_channel_smoke.py b/tests/test_additional_channel_smoke.py index c443c73..17e8025 100644 --- a/tests/test_additional_channel_smoke.py +++ b/tests/test_additional_channel_smoke.py @@ -10,18 +10,17 @@ from EvoScientist.channels.imessage.channel_rpc import ( ) from EvoScientist.channels.qq.channel import QQChannel, QQConfig from EvoScientist.channels.signal.channel import SignalChannel, SignalConfig -from tests.conftest import run_async as _run class TestEmailChannelSmoke: - def test_start_raises_without_required_imap_settings(self): + async def test_start_raises_without_required_imap_settings(self): channel = EmailChannel(EmailConfig()) with pytest.raises( ChannelError, match="imap_host and imap_username are required" ): - _run(channel.start()) + await channel.start() - def test_send_returns_false_when_smtp_not_ready(self): + async def test_send_returns_false_when_smtp_not_ready(self): channel = EmailChannel(EmailConfig()) msg = OutboundMessage( channel="email", @@ -29,16 +28,16 @@ class TestEmailChannelSmoke: content="hello", metadata={"chat_id": "user@example.com"}, ) - assert _run(channel.send(msg)) is False + assert await channel.send(msg) is False class TestSignalChannelSmoke: - def test_start_raises_without_phone_number(self): + async def test_start_raises_without_phone_number(self): channel = SignalChannel(SignalConfig()) with pytest.raises(ChannelError, match="phone_number is required"): - _run(channel.start()) + await channel.start() - def test_send_returns_false_when_not_connected(self): + async def test_send_returns_false_when_not_connected(self): channel = SignalChannel(SignalConfig(phone_number="+123456789")) msg = OutboundMessage( channel="signal", @@ -46,27 +45,29 @@ class TestSignalChannelSmoke: content="hello", metadata={"chat_id": "+123456789"}, ) - assert _run(channel.send(msg)) is False + assert await channel.send(msg) is False class TestQQChannelSmoke: - def test_start_raises_when_sdk_missing(self, monkeypatch): + async def test_start_raises_when_sdk_missing(self, monkeypatch): from EvoScientist.channels.qq import channel as qq_module monkeypatch.setattr(qq_module, "QQ_AVAILABLE", False) channel = QQChannel(QQConfig(app_id="id", app_secret="secret")) with pytest.raises(ChannelError, match="SDK not installed"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_credentials_when_sdk_available(self, monkeypatch): + async def test_start_raises_without_credentials_when_sdk_available( + self, monkeypatch + ): from EvoScientist.channels.qq import channel as qq_module monkeypatch.setattr(qq_module, "QQ_AVAILABLE", True) channel = QQChannel(QQConfig(app_id="", app_secret="")) with pytest.raises(ChannelError, match="app_id and app_secret are required"): - _run(channel.start()) + await channel.start() - def test_send_returns_false_without_client(self): + async def test_send_returns_false_without_client(self): channel = QQChannel(QQConfig(app_id="id", app_secret="secret")) msg = OutboundMessage( channel="qq", @@ -74,11 +75,11 @@ class TestQQChannelSmoke: content="hello", metadata={"chat_id": "openid"}, ) - assert _run(channel.send(msg)) is False + assert await channel.send(msg) is False class TestIMessageChannelSmoke: - def test_start_wraps_rpc_bootstrap_error(self, monkeypatch): + async def test_start_wraps_rpc_bootstrap_error(self, monkeypatch): async def _broken_start(self): raise RuntimeError("imsg not found") @@ -87,9 +88,9 @@ class TestIMessageChannelSmoke: monkeypatch.setattr(imessage_module.ImsgRpcClient, "start", _broken_start) channel = IMessageChannelRpc(IMessageConfig()) with pytest.raises(ChannelError, match="Failed to start imsg"): - _run(channel.start()) + await channel.start() - def test_send_returns_false_without_rpc_client(self): + async def test_send_returns_false_without_rpc_client(self): channel = IMessageChannelRpc(IMessageConfig()) msg = OutboundMessage( channel="imessage", @@ -97,4 +98,4 @@ class TestIMessageChannelSmoke: content="hello", metadata={"chat_id": "+123456789"}, ) - assert _run(channel.send(msg)) is False + assert await channel.send(msg) is False diff --git a/tests/test_agent_loader.py b/tests/test_agent_loader.py index fdb84fd..fa2d6e3 100644 --- a/tests/test_agent_loader.py +++ b/tests/test_agent_loader.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import threading import pytest @@ -93,88 +94,95 @@ def _make_loader_fn(agent_value="AGENT", fail_with=None, capture=None): return _loader -def _run(coro): - return asyncio.run(coro) +class _GatedThreadLoader: + """Callable loader that blocks until tests explicitly release it.""" + + def __init__(self, agent_value="AGENT", progress_events=()): + self.agent_value = agent_value + self.progress_events = tuple(progress_events) + self.started = threading.Event() + self.release = threading.Event() + self.finished = threading.Event() + + def __call__(self, *, on_mcp_progress=None): + self.started.set() + self.release.wait(timeout=1) + try: + if on_mcp_progress is not None: + for event in self.progress_events: + on_mcp_progress(*event) + return self.agent_value + finally: + self.finished.set() + + +async def _wait_for_event(event, timeout=1): + return await asyncio.to_thread(event.wait, timeout) class TestBackgroundAgentLoaderStart: - def test_start_creates_task_and_forwards_kwargs(self): + async def test_start_creates_task_and_forwards_kwargs(self): captured: dict = {} loader = BackgroundAgentLoader(_make_loader_fn(capture=captured)) - async def _go(): - loader.start(workspace_dir="/ws", checkpointer="CK") - assert loader.task is not None - assert loader.is_pending - await loader.await_ready() + loader.start(workspace_dir="/ws", checkpointer="CK") + assert loader.task is not None + assert loader.is_pending + await loader.await_ready() - _run(_go()) assert captured["kwargs"][0] == {"workspace_dir": "/ws", "checkpointer": "CK"} - def test_start_bumps_load_id(self): + async def test_start_bumps_load_id(self): loader = BackgroundAgentLoader(_make_loader_fn()) - async def _go(): - assert loader._load_id == 0 - loader.start() - assert loader._load_id == 1 - loader.start() - assert loader._load_id == 2 - await loader.await_ready() + assert loader._load_id == 0 + loader.start() + assert loader._load_id == 1 + loader.start() + assert loader._load_id == 2 + await loader.await_ready() - _run(_go()) + async def test_start_cancels_in_flight_prior_task(self): + blocking = _GatedThreadLoader("LATE") - def test_start_cancels_in_flight_prior_task(self): - import time - - def _blocking(*, on_mcp_progress=None): - time.sleep(0.05) - return "LATE" - - async def _go(): - loader = BackgroundAgentLoader(_blocking) - loader.start() - first_task = loader.task - # Supersede immediately; asyncio.to_thread wrapper gets cancelled. - loader._loader_fn = _make_loader_fn("FRESH") - loader.start() - agent = await loader.await_ready() - assert agent == "FRESH" - # Let the first thread drain so its done callback (gated) fires. - await asyncio.sleep(0.1) - assert first_task.cancelled() or first_task.done() - - _run(_go()) + loader = BackgroundAgentLoader(blocking) + loader.start() + first_task = loader.task + assert first_task is not None + assert await _wait_for_event(blocking.started) + # Supersede immediately; asyncio.to_thread wrapper gets cancelled. + loader._loader_fn = _make_loader_fn("FRESH") + loader.start() + agent = await loader.await_ready() + assert agent == "FRESH" + blocking.release.set() + try: + await first_task + except asyncio.CancelledError: + pass + assert first_task.cancelled() or first_task.done() class TestBackgroundAgentLoaderCallbacks: - def test_progress_hook_sees_events_in_order(self): + async def test_progress_hook_sees_events_in_order(self): events: list[tuple[str, str, str]] = [] loader = BackgroundAgentLoader( _make_loader_fn(capture={}), on_progress=lambda e, s, d: events.append((e, s, d)), ) - async def _go(): - loader.start() - await loader.await_ready() + loader.start() + await loader.await_ready() - _run(_go()) assert events == [("start", "srv", ""), ("success", "srv", "1")] - def test_stale_progress_events_are_dropped(self): + async def test_stale_progress_events_are_dropped(self): """A progress event fired after a newer `start` must not reach the hook.""" - import time - + slow_loader = _GatedThreadLoader( + "slow-agent", progress_events=[("success", "from-slow", "1")] + ) seen: list[str] = [] - # Loader 1 sleeps so its progress event fires AFTER load 2 starts. - def slow_loader(*, on_mcp_progress=None): - time.sleep(0.08) - if on_mcp_progress is not None: - on_mcp_progress("success", "from-slow", "1") - return "slow-agent" - def fast_loader(*, on_mcp_progress=None): if on_mcp_progress is not None: on_mcp_progress("success", "from-fast", "1") @@ -184,36 +192,32 @@ class TestBackgroundAgentLoaderCallbacks: slow_loader, on_progress=lambda e, s, d: seen.append(s) ) - async def _go(): - loader.start() - # Supersede before the slow thread's event fires. - await asyncio.sleep(0.01) - loader._loader_fn = fast_loader - loader.start() - await loader.await_ready() - # Let the superseded thread finish (its event is gated out). - await asyncio.sleep(0.1) + loader.start() + assert await _wait_for_event(slow_loader.started) + # Loader 1 waits so its progress event fires AFTER load 2 starts. + loader._loader_fn = fast_loader + loader.start() + await loader.await_ready() + slow_loader.release.set() + assert await _wait_for_event(slow_loader.finished) - _run(_go()) assert "from-fast" in seen assert "from-slow" not in seen - def test_success_callback_fires_on_completion(self): + async def test_success_callback_fires_on_completion(self): got = [] loader = BackgroundAgentLoader( _make_loader_fn("MY_AGENT"), on_success=lambda a: got.append(a), ) - async def _go(): - loader.start() - await loader.await_ready() - await asyncio.sleep(0) # let done-callback run + loader.start() + await loader.await_ready() + await asyncio.sleep(0) # let done-callback run - _run(_go()) assert got == ["MY_AGENT"] - def test_failure_callback_fires_on_error(self): + async def test_failure_callback_fires_on_error(self): err = RuntimeError("load failed") got_failures = [] got_successes = [] @@ -223,40 +227,33 @@ class TestBackgroundAgentLoaderCallbacks: on_failure=lambda e: got_failures.append(e), ) - async def _go(): - loader.start() - with pytest.raises(RuntimeError, match="load failed"): - await loader.await_ready() - await asyncio.sleep(0) + loader.start() + with pytest.raises(RuntimeError, match="load failed"): + await loader.await_ready() + await asyncio.sleep(0) - _run(_go()) assert got_failures == [err] assert got_successes == [] class TestBackgroundAgentLoaderAwaitReady: - def test_returns_cached_agent_without_reawaiting(self): + async def test_returns_cached_agent_without_reawaiting(self): captured: dict = {} loader = BackgroundAgentLoader(_make_loader_fn("A", capture=captured)) - async def _go(): - loader.start() - assert await loader.await_ready() == "A" - assert await loader.await_ready() == "A" + loader.start() + assert await loader.await_ready() == "A" + assert await loader.await_ready() == "A" - _run(_go()) assert len(captured["kwargs"]) == 1 - def test_raises_if_started_not_called(self): + async def test_raises_if_started_not_called(self): loader = BackgroundAgentLoader(_make_loader_fn()) - async def _go(): - with pytest.raises(RuntimeError, match="before start"): - await loader.await_ready() + with pytest.raises(RuntimeError, match="before start"): + await loader.await_ready() - _run(_go()) - - def test_reraises_real_error_on_subsequent_awaits(self): + async def test_reraises_real_error_on_subsequent_awaits(self): """After a failure, ``await_ready`` must keep raising the real exception — not the "before start()" sentinel — until ``start`` is called again.""" @@ -265,16 +262,13 @@ class TestBackgroundAgentLoaderAwaitReady: loader = BackgroundAgentLoader(_fail) - async def _go(): - loader.start() - with pytest.raises(RuntimeError, match="bad MCP config"): - await loader.await_ready() - with pytest.raises(RuntimeError, match="bad MCP config"): - await loader.await_ready() + loader.start() + with pytest.raises(RuntimeError, match="bad MCP config"): + await loader.await_ready() + with pytest.raises(RuntimeError, match="bad MCP config"): + await loader.await_ready() - _run(_go()) - - def test_needs_restart_flags_failed_load_for_retry(self): + async def test_needs_restart_flags_failed_load_for_retry(self): calls = {"n": 0} def flaky(*, on_mcp_progress=None): @@ -285,17 +279,14 @@ class TestBackgroundAgentLoaderAwaitReady: loader = BackgroundAgentLoader(flaky) - async def _go(): - assert loader.needs_restart # never started - loader.start() - with pytest.raises(RuntimeError): - await loader.await_ready() - assert loader.needs_restart # failed, caller may retry - loader.start() - assert await loader.await_ready() == "SECOND" - assert not loader.needs_restart # success → no retry - - _run(_go()) + assert loader.needs_restart # never started + loader.start() + with pytest.raises(RuntimeError): + await loader.await_ready() + assert loader.needs_restart # failed, caller may retry + loader.start() + assert await loader.await_ready() == "SECOND" + assert not loader.needs_restart # success → no retry class TestBackgroundAgentLoaderAdopt: @@ -305,26 +296,19 @@ class TestBackgroundAgentLoaderAdopt: assert loader.agent == "EXTERNAL" assert not loader.is_pending - def test_adopt_supersedes_in_flight_load(self): + async def test_adopt_supersedes_in_flight_load(self): """A late background completion must not overwrite an adopted agent.""" - import time + slow_loader = _GatedThreadLoader("FROM_BACKGROUND") - def _slow(*, on_mcp_progress=None): - time.sleep(0.08) - return "FROM_BACKGROUND" + loader = BackgroundAgentLoader(slow_loader) - loader = BackgroundAgentLoader(_slow) - - async def _go(): - loader.start() - await asyncio.sleep(0.01) - loader.adopt("FROM_MODEL") - # Give the background thread time to finish and fire its - # done-callback; the generation token should make it a no-op. - await asyncio.sleep(0.1) - assert loader.agent == "FROM_MODEL" - - _run(_go()) + loader.start() + assert await _wait_for_event(slow_loader.started) + loader.adopt("FROM_MODEL") + slow_loader.release.set() + assert await _wait_for_event(slow_loader.finished) + await asyncio.sleep(0) + assert loader.agent == "FROM_MODEL" class TestBackgroundAgentLoaderIsPending: @@ -332,29 +316,22 @@ class TestBackgroundAgentLoaderIsPending: loader = BackgroundAgentLoader(_make_loader_fn()) assert not loader.is_pending - def test_false_after_completion(self): + async def test_false_after_completion(self): loader = BackgroundAgentLoader(_make_loader_fn()) - async def _go(): - loader.start() - await loader.await_ready() + loader.start() + await loader.await_ready() - _run(_go()) assert not loader.is_pending - def test_true_between_start_and_completion(self): - import time + async def test_true_between_start_and_completion(self): + wait_loader = _GatedThreadLoader("ok") - def _wait_loader(*, on_mcp_progress=None): - time.sleep(0.05) - return "ok" + loader = BackgroundAgentLoader(wait_loader) - loader = BackgroundAgentLoader(_wait_loader) - - async def _go(): - loader.start() - assert loader.is_pending - await loader.await_ready() - assert not loader.is_pending - - _run(_go()) + loader.start() + assert await _wait_for_event(wait_loader.started) + assert loader.is_pending + wait_loader.release.set() + await loader.await_ready() + assert not loader.is_pending diff --git a/tests/test_async_notifier.py b/tests/test_async_notifier.py index b7c7769..a3b4d72 100644 --- a/tests/test_async_notifier.py +++ b/tests/test_async_notifier.py @@ -5,6 +5,8 @@ import queue from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock +import pytest + from EvoScientist.cli import async_notifier from EvoScientist.cli.async_notifier import ( dedup_notifications, @@ -28,12 +30,6 @@ def test_notification_dataclass_fields(): def test_notification_queue_is_module_level_fifo(): - # Drain anything left over from other tests - while True: - try: - async_notifier._notification_queue.get_nowait() - except queue.Empty: - break n1 = async_notifier.AsyncTaskNotification("a", "x", "success", "") n2 = async_notifier.AsyncTaskNotification("b", "x", "success", "") async_notifier._notification_queue.put(n1) @@ -51,7 +47,7 @@ def _drain_queue(q): return items -def test_read_async_tasks_from_gateway_reads_state_values(run_async): +async def test_read_async_tasks_from_gateway_reads_state_values(): gateway = FakeGraphGateway( state_values={ "async_tasks": { @@ -60,18 +56,16 @@ def test_read_async_tasks_from_gateway_reads_state_values(run_async): } ) - tasks = run_async( - async_notifier.read_async_tasks_from_gateway( - gateway, - GraphTarget(local_graph=MagicMock()), - "tid", - ) + tasks = await async_notifier.read_async_tasks_from_gateway( + gateway, + GraphTarget(local_graph=MagicMock()), + "tid", ) assert tasks == {"task-1": {"status": "success"}} -def test_watcher_pushes_notification_on_stream_end(run_async): +async def test_watcher_pushes_notification_on_stream_end(): # Stream yields one "values" chunk with the final state, then closes final_state = { "messages": [{"type": "ai", "content": "Quantum superposition is..."}] @@ -87,10 +81,7 @@ def test_watcher_pushes_notification_on_stream_end(run_async): # runs.get is used to fetch terminal status when stream ends client.runs.get = AsyncMock(return_value={"status": "success"}) - _drain_all(async_notifier) - run_async( - async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent") - ) + await async_notifier.watch_run_and_notify(client, "thr-1", "run-1", "writing-agent") notifs = _drain_queue(async_notifier._notification_queue) assert len(notifs) == 1 @@ -99,7 +90,7 @@ def test_watcher_pushes_notification_on_stream_end(run_async): assert notifs[0].status == "success" -def test_watcher_pushes_error_status_on_stream_exception(run_async): +async def test_watcher_pushes_error_status_on_stream_exception(): async def fake_stream(*a, **kw): raise RuntimeError("network broken") yield # unreachable; makes this an async generator @@ -111,14 +102,13 @@ def test_watcher_pushes_error_status_on_stream_exception(run_async): return_value={"status": "error", "error": "network broken"} ) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ")) + await async_notifier.watch_run_and_notify(client, "thr-4", "run-4", "agentZ") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "error" -def test_spawn_watcher_replaces_existing_for_same_thread(run_async): +async def test_spawn_watcher_replaces_existing_for_same_thread(): """A second spawn_watcher with the same thread_id cancels the old watcher and registers the new one — supports update_async_task creating a new run_id on the same thread_id.""" @@ -138,43 +128,34 @@ def test_spawn_watcher_replaces_existing_for_same_thread(run_async): client.runs.join_stream = fake_stream_long client.runs.get = AsyncMock(return_value={"status": "success"}) - async def scenario(): - # Clear all queues and the watcher registries - async_notifier._active_watchers.clear() - async_notifier._watcher_by_thread.clear() - _drain_all(async_notifier) + # First spawn for thread X, run R1 + t1 = async_notifier.spawn_watcher(client, "thr-X", "R1", "agent") + assert t1 is not None + assert async_notifier._watcher_by_thread["thr-X"] is t1 + await asyncio.sleep(0.02) # let it start streaming - # First spawn for thread X, run R1 - t1 = async_notifier.spawn_watcher(client, "thr-X", "R1", "agent") - assert t1 is not None - assert async_notifier._watcher_by_thread["thr-X"] is t1 - await asyncio.sleep(0.02) # let it start streaming + # Second spawn for SAME thread X, NEW run R2 + t2 = async_notifier.spawn_watcher(client, "thr-X", "R2", "agent") + assert t2 is not None + assert t2 is not t1 + assert async_notifier._watcher_by_thread["thr-X"] is t2 - # Second spawn for SAME thread X, NEW run R2 - t2 = async_notifier.spawn_watcher(client, "thr-X", "R2", "agent") - assert t2 is not None - assert t2 is not t1 - assert async_notifier._watcher_by_thread["thr-X"] is t2 + # Old watcher should be cancelled + await asyncio.sleep(0.02) + assert t1.cancelled() or t1.done() - # Old watcher should be cancelled - await asyncio.sleep(0.02) - assert t1.cancelled() or t1.done() + # Cleanup the new task too + t2.cancel() + try: + await t2 + except asyncio.CancelledError: + pass - # Cleanup the new task too - t2.cancel() - try: - await t2 - except asyncio.CancelledError: - pass - - # Cancelled watchers don't push notifications - assert _drain_one_queue_helper(async_notifier._notification_queue) == [] - assert _drain_one_queue_helper(async_notifier._unrouted_queue) == [] - if hasattr(async_notifier, "_notifications_by_thread"): - for q in async_notifier._notifications_by_thread.values(): - assert _drain_one_queue_helper(q) == [] - - run_async(scenario()) + # Cancelled watchers don't push notifications + assert _drain_one_queue_helper(async_notifier._notification_queue) == [] + assert _drain_one_queue_helper(async_notifier._unrouted_queue) == [] + for q in async_notifier._notifications_by_thread.values(): + assert _drain_one_queue_helper(q) == [] # ============================================================================ @@ -324,13 +305,6 @@ def test_format_notification_lines_timeout_uses_warning_icon(): def test_drain_returns_all_pending_and_empties_queue(): """drain_notifications pulls every pending notification and empties queue.""" - # Clear the queue first - while True: - try: - async_notifier._notification_queue.get_nowait() - except queue.Empty: - break - # Add three notifications for tid in ("a", "b", "c"): async_notifier._notification_queue.put( @@ -463,17 +437,12 @@ def test_dedup_preserves_order(): # ============================================================================ -def test_consume_notifications_calls_runner_with_batched_message(run_async): +async def test_consume_notifications_calls_runner_with_batched_message(): """When notifications arrive and agent is idle, consume_notifications fires the supplied async runner once with the formatted batch message and notifs list.""" from EvoScientist.cli import async_notifier as an # Set up two pending notifications, no dedup match - while True: - try: - an._notification_queue.get_nowait() - except queue.Empty: - break an._notification_queue.put(an.AsyncTaskNotification("t1", "wA", "success", "", "")) an._notification_queue.put(an.AsyncTaskNotification("t2", "wB", "success", "", "")) @@ -486,21 +455,15 @@ def test_consume_notifications_calls_runner_with_batched_message(run_async): async def fake_state_reader() -> dict: return {} # no dedup info - run_async(an.consume_notifications(fake_runner, fake_state_reader)) + await an.consume_notifications(fake_runner, fake_state_reader) assert "wA" in captured["text"] assert "wB" in captured["text"] assert len(captured["notifs"]) == 2 -def test_consume_notifications_no_op_when_queue_empty(run_async): +async def test_consume_notifications_no_op_when_queue_empty(): from EvoScientist.cli import async_notifier as an - while True: - try: - an._notification_queue.get_nowait() - except queue.Empty: - break - called = False async def fake_runner(text: str, notifs: list): @@ -510,7 +473,7 @@ def test_consume_notifications_no_op_when_queue_empty(run_async): async def fake_state_reader(): return {} - run_async(an.consume_notifications(fake_runner, fake_state_reader)) + await an.consume_notifications(fake_runner, fake_state_reader) assert called is False @@ -521,7 +484,7 @@ def test_consume_notifications_no_op_when_queue_empty(run_async): # ============================================================================ -def test_notification_consuming_flag_prevents_reentry(run_async): +async def test_notification_consuming_flag_prevents_reentry(): """The _notification_consuming guard prevents two overlapping consumers. Verifies the flag contract used by _consume_notifications_tui: @@ -536,13 +499,6 @@ def test_notification_consuming_flag_prevents_reentry(run_async): """ from EvoScientist.cli import async_notifier as an - # Clear the queue - while True: - try: - an._notification_queue.get_nowait() - except queue.Empty: - break - state = {"inject_count": 0, "consuming": False} async def counting_runner(text: str, notifs: list) -> None: @@ -565,45 +521,42 @@ def test_notification_consuming_flag_prevents_reentry(run_async): n1 = an.AsyncTaskNotification("g1", "writing-agent", "success", "", "") n2 = an.AsyncTaskNotification("g2", "data-agent", "success", "", "") - async def scenario(): - # Scenario 1: normal flow — flag cleared, second consumer runs fine. - await guarded_consume(n1) - assert state["inject_count"] == 1 - assert state["consuming"] is False # finally ran + # Scenario 1: normal flow — flag cleared, second consumer runs fine. + await guarded_consume(n1) + assert state["inject_count"] == 1 + assert state["consuming"] is False # finally ran - state["inject_count"] = 0 - await guarded_consume(n2) - assert state["inject_count"] == 1 - assert state["consuming"] is False + state["inject_count"] = 0 + await guarded_consume(n2) + assert state["inject_count"] == 1 + assert state["consuming"] is False - # Scenario 2: flag pre-set (first consumer in-flight) → second bails. - state["inject_count"] = 0 - state["consuming"] = True # simulate first consumer running - an._notification_queue.put(n1) - await guarded_consume(n1) # should be blocked immediately - assert state["inject_count"] == 0 # runner never called - state["consuming"] = False # cleanup + # Scenario 2: flag pre-set (first consumer in-flight) → second bails. + state["inject_count"] = 0 + state["consuming"] = True # simulate first consumer running + an._notification_queue.put(n1) + await guarded_consume(n1) # should be blocked immediately + assert state["inject_count"] == 0 # runner never called + state["consuming"] = False # cleanup - # Scenario 3: exception in runner → flag still cleared by finally. - async def raising_runner(text: str, notifs: list) -> None: - raise RuntimeError("boom") + # Scenario 3: exception in runner → flag still cleared by finally. + async def raising_runner(text: str, notifs: list) -> None: + raise RuntimeError("boom") - async def guarded_consume_raising(notif): - if state["consuming"]: - return - state["consuming"] = True - try: - an._notification_queue.put(notif) - await an.consume_notifications(raising_runner, fake_state_reader) - except RuntimeError: - pass - finally: - state["consuming"] = False + async def guarded_consume_raising(notif): + if state["consuming"]: + return + state["consuming"] = True + try: + an._notification_queue.put(notif) + await an.consume_notifications(raising_runner, fake_state_reader) + except RuntimeError: + pass + finally: + state["consuming"] = False - await guarded_consume_raising(n2) - assert state["consuming"] is False # cleared despite exception - - run_async(scenario()) + await guarded_consume_raising(n2) + assert state["consuming"] is False # cleared despite exception # ============================================================================ @@ -613,33 +566,42 @@ def test_notification_consuming_flag_prevents_reentry(run_async): def _drain_all(an_mod): """Drain every queue (per-thread + unrouted) so tests start clean.""" - if hasattr(an_mod, "_notification_queue"): + while True: + try: + an_mod._notification_queue.get_nowait() + except queue.Empty: + break + for q in list(an_mod._notifications_by_thread.values()): while True: try: - an_mod._notification_queue.get_nowait() - except queue.Empty: - break - if hasattr(an_mod, "_notifications_by_thread"): - for q in list(an_mod._notifications_by_thread.values()): - while True: - try: - q.get_nowait() - except queue.Empty: - break - if hasattr(an_mod, "_unrouted_queue"): - while True: - try: - an_mod._unrouted_queue.get_nowait() + q.get_nowait() except queue.Empty: break + while True: + try: + an_mod._unrouted_queue.get_nowait() + except queue.Empty: + break -def test_consume_only_drains_matching_thread(run_async): +def _reset_notifier_state(an_mod): + _drain_all(an_mod) + an_mod._active_watchers.clear() + an_mod._watcher_by_thread.clear() + + +@pytest.fixture(autouse=True) +def _clean_async_notifier_state(): + _reset_notifier_state(async_notifier) + yield + _reset_notifier_state(async_notifier) + + +async def test_consume_only_drains_matching_thread(): """Notifications tagged with origin_cli_thread_id only drain when the consumer is invoked with the matching current_thread_id.""" from EvoScientist.cli import async_notifier as an - _drain_all(an) n_a = an.AsyncTaskNotification( "tA", "writing-agent", "success", "", "", origin_cli_thread_id="threadA" ) @@ -657,21 +619,17 @@ def test_consume_only_drains_matching_thread(run_async): async def state_reader() -> dict: return {} - run_async( - an.consume_notifications(runner, state_reader, current_thread_id="threadA") - ) + await an.consume_notifications(runner, state_reader, current_thread_id="threadA") assert captured["runs"] == [["tA"]] # B's notification should still be queued assert an.has_pending_notifications("threadB") - _drain_all(an) -def test_unrouted_notifications_drain_on_any_thread(run_async): +async def test_unrouted_notifications_drain_on_any_thread(): """Notifications without origin_cli_thread_id (legacy / direct put) drain regardless of the current_thread_id arg.""" from EvoScientist.cli import async_notifier as an - _drain_all(an) an._notification_queue.put( an.AsyncTaskNotification("tU", "writing-agent", "success", "", "") ) @@ -684,19 +642,15 @@ def test_unrouted_notifications_drain_on_any_thread(run_async): async def state_reader() -> dict: return {} - run_async( - an.consume_notifications(runner, state_reader, current_thread_id="anything") - ) + await an.consume_notifications(runner, state_reader, current_thread_id="anything") assert [n.task_id for n in captured["notifs"]] == ["tU"] - _drain_all(an) -def test_thread_switch_drains_pending(run_async): +async def test_thread_switch_drains_pending(): """Pending notifications for thread B are not delivered while consumer asks for thread A; once consumer runs with thread B they drain.""" from EvoScientist.cli import async_notifier as an - _drain_all(an) an._enqueue( an.AsyncTaskNotification( "tB", "writing-agent", "success", "", "", origin_cli_thread_id="threadB" @@ -712,25 +666,19 @@ def test_thread_switch_drains_pending(run_async): return {} # First consume in thread A → no drain, B's notif still queued - run_async( - an.consume_notifications(runner, state_reader, current_thread_id="threadA") - ) + await an.consume_notifications(runner, state_reader, current_thread_id="threadA") assert captured["runs"] == [] assert an.has_pending_notifications("threadB") # Now switch to thread B → drains - run_async( - an.consume_notifications(runner, state_reader, current_thread_id="threadB") - ) + await an.consume_notifications(runner, state_reader, current_thread_id="threadB") assert captured["runs"] == [["tB"]] - _drain_all(an) def test_has_pending_notifications_respects_routing(): """has_pending_notifications returns true only for matching or unrouted.""" from EvoScientist.cli import async_notifier as an - _drain_all(an) # Unrouted always counts an._notification_queue.put( an.AsyncTaskNotification("tU", "writing-agent", "success", "", "") @@ -748,7 +696,6 @@ def test_has_pending_notifications_respects_routing(): assert an.has_pending_notifications("threadA") is True assert an.has_pending_notifications("threadB") is False assert an.has_pending_notifications() is False # no unrouted, no current_thread - _drain_all(an) # ============================================================================ @@ -762,7 +709,7 @@ def test_has_pending_notifications_respects_routing(): # ============================================================================ -def test_watcher_reports_error_on_in_band_error_event(run_async): +async def test_watcher_reports_error_on_in_band_error_event(): """SSE error event in the stream → notification.status == 'error'.""" async def fake_stream(*a, **kw): @@ -777,8 +724,7 @@ def test_watcher_reports_error_on_in_band_error_event(run_async): return_value={"status": "success"} ) # would mislead — should NOT be consulted - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE")) + await async_notifier.watch_run_and_notify(client, "thrE", "rE", "agentE") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "error" @@ -786,7 +732,7 @@ def test_watcher_reports_error_on_in_band_error_event(run_async): client.runs.get.assert_not_awaited() -def test_watcher_clean_exit_with_runs_get_success_is_success(run_async): +async def test_watcher_clean_exit_with_runs_get_success_is_success(): """Clean stream exit + runs.get reports success → status=success.""" async def fake_stream(*a, **kw): @@ -798,15 +744,14 @@ def test_watcher_clean_exit_with_runs_get_success_is_success(run_async): client.runs.join_stream = fake_stream client.runs.get = AsyncMock(return_value={"status": "success"}) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS")) + await async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "success" client.runs.get.assert_awaited_once() -def test_watcher_clean_exit_with_runs_get_error_is_race_safe(run_async): +async def test_watcher_clean_exit_with_runs_get_error_is_race_safe(): """Clean stream exit + no in-band error event + runs.get returns 'error' → status=success (race-safe). @@ -827,14 +772,13 @@ def test_watcher_clean_exit_with_runs_get_error_is_race_safe(run_async): client.runs.join_stream = fake_stream client.runs.get = AsyncMock(return_value={"status": "error"}) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS")) + await async_notifier.watch_run_and_notify(client, "thrS", "rS", "agentS") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "success" -def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async): +async def test_watcher_clean_exit_with_runs_get_running_drops_notification(): """Reproduces the production bug: clean SSE close while run is still actually running (HTTP keep-alive timeout under concurrency). @@ -856,11 +800,8 @@ def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async): client.runs.join_stream = fake_stream client.runs.get = AsyncMock(return_value={"status": "running"}) - _drain_all(async_notifier) - run_async( - async_notifier.watch_run_and_notify( - client, "thr-bug", "rB", "data-analysis-agent" - ) + await async_notifier.watch_run_and_notify( + client, "thr-bug", "rB", "data-analysis-agent" ) # No notification should have been enqueued anywhere. @@ -872,7 +813,7 @@ def test_watcher_clean_exit_with_runs_get_running_drops_notification(run_async): assert client.runs.get.await_count >= 1 -def test_watcher_unknown_status_treated_as_non_terminal(run_async): +async def test_watcher_unknown_status_treated_as_non_terminal(): """Future / unrecognized status values should trigger a re-join, not a false-positive notification. @@ -893,8 +834,7 @@ def test_watcher_unknown_status_treated_as_non_terminal(run_async): side_effect=[{"status": "queued"}, {"status": "success"}] ) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrU", "rU", "agentU")) + await async_notifier.watch_run_and_notify(client, "thrU", "rU", "agentU") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "success" @@ -902,7 +842,7 @@ def test_watcher_unknown_status_treated_as_non_terminal(run_async): assert client.runs.get.await_count == 2 -def test_watcher_runs_get_persistent_failure_drops_notification(run_async, monkeypatch): +async def test_watcher_runs_get_persistent_failure_drops_notification(monkeypatch): """If ``runs.get`` keeps raising, the watcher cannot verify terminal state and MUST drop the notification rather than default to ``"success"`` — otherwise a transient server outage reintroduces the @@ -921,22 +861,20 @@ def test_watcher_runs_get_persistent_failure_drops_notification(run_async, monke monkeypatch.setattr(async_notifier.asyncio, "sleep", _no_sleep) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrG", "rG", "agentG")) + await async_notifier.watch_run_and_notify(client, "thrG", "rG", "agentG") # No notification — watcher exhausted the reconnect budget. Check every # queue routing could send to so a future routing change can't make this # test silently false-pass. assert _drain_one_queue_helper(async_notifier._unrouted_queue) == [] assert _drain_one_queue_helper(async_notifier._notification_queue) == [] - if hasattr(async_notifier, "_notifications_by_thread"): - for q in async_notifier._notifications_by_thread.values(): - assert _drain_one_queue_helper(q) == [] + for q in async_notifier._notifications_by_thread.values(): + assert _drain_one_queue_helper(q) == [] # 1 initial + _MAX_RECONNECT_ATTEMPTS retries = 11 calls total. assert client.runs.get.await_count == async_notifier._MAX_RECONNECT_ATTEMPTS + 1 -def test_watcher_runs_get_transient_failure_recovers(run_async, monkeypatch): +async def test_watcher_runs_get_transient_failure_recovers(monkeypatch): """A single ``runs.get`` failure followed by a successful response on retry must produce a correct notification — verifies the bounded retry path actually recovers from transient outages instead of just @@ -957,15 +895,14 @@ def test_watcher_runs_get_transient_failure_recovers(run_async, monkeypatch): monkeypatch.setattr(async_notifier.asyncio, "sleep", _no_sleep) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrT", "rT", "agentT")) + await async_notifier.watch_run_and_notify(client, "thrT", "rT", "agentT") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "success" assert client.runs.get.await_count == 2 -def test_watcher_re_joins_stream_until_terminal_status(run_async): +async def test_watcher_re_joins_stream_until_terminal_status(): """When runs.get returns 'running' on attempt N but a terminal status on attempt N+1, the watcher re-joins, observes the terminal status, and enqueues the notification correctly.""" @@ -980,8 +917,7 @@ def test_watcher_re_joins_stream_until_terminal_status(run_async): side_effect=[{"status": "running"}, {"status": "success"}] ) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrR", "rR", "agentR")) + await async_notifier.watch_run_and_notify(client, "thrR", "rR", "agentR") notif = async_notifier._notification_queue.get_nowait() assert notif.status == "success" @@ -995,15 +931,12 @@ def test_watcher_re_joins_stream_until_terminal_status(run_async): # ============================================================================ -def test_consume_notifications_propagates_inject_exception(run_async): +async def test_consume_notifications_propagates_inject_exception(): """If the run_message callback raises, consume_notifications propagates the exception to the caller — pollers wrap it in try/except so the poller task does not die.""" - import pytest - from EvoScientist.cli import async_notifier as an - _drain_all(an) an._notification_queue.put( an.AsyncTaskNotification("tX", "writing-agent", "success", "", "") ) @@ -1015,11 +948,10 @@ def test_consume_notifications_propagates_inject_exception(run_async): return {} with pytest.raises(RuntimeError, match="kaboom"): - run_async(an.consume_notifications(boom_runner, state_reader)) - _drain_all(an) + await an.consume_notifications(boom_runner, state_reader) -def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(run_async): +async def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(): """When the SSE stream errors AND runs.get returns a non-terminal status (e.g. ``pending`` because the run is still alive), the watcher must NOT enqueue a notification — otherwise the user sees a confusing @@ -1035,15 +967,13 @@ def test_watcher_skips_notification_on_stream_fail_with_nonterminal_status(run_a client.runs.join_stream = fake_stream client.runs.get = AsyncMock(return_value={"status": "pending"}) - _drain_all(async_notifier) - run_async(async_notifier.watch_run_and_notify(client, "thrP", "rP", "agentP")) + await async_notifier.watch_run_and_notify(client, "thrP", "rP", "agentP") # No notification should have been enqueued in any queue. assert _drain_one_queue_helper(async_notifier._unrouted_queue) == [] assert _drain_one_queue_helper(async_notifier._notification_queue) == [] - if hasattr(async_notifier, "_notifications_by_thread"): - for q in async_notifier._notifications_by_thread.values(): - assert _drain_one_queue_helper(q) == [] + for q in async_notifier._notifications_by_thread.values(): + assert _drain_one_queue_helper(q) == [] def _drain_one_queue_helper(q): @@ -1060,8 +990,6 @@ def test_active_watchers_grace_filters_by_thread(): (otherwise consume_notifications grace period would block thread A by up to 3s waiting for thread B's unrelated watchers to finish).""" - async_notifier._active_watchers.clear() - # Sentinel handles — only their identity matters here, not their type handle_a = object() handle_b = object() diff --git a/tests/test_async_watcher_middleware.py b/tests/test_async_watcher_middleware.py index ae81e9d..09295c6 100644 --- a/tests/test_async_watcher_middleware.py +++ b/tests/test_async_watcher_middleware.py @@ -7,7 +7,6 @@ deepagents internals. It hooks into ``awrap_tool_call`` and only fires on from __future__ import annotations -import asyncio from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -105,7 +104,7 @@ def _make_middleware(): return mw, fake_client -def test_middleware_spawns_watcher_on_start_async_task(): +async def test_middleware_spawns_watcher_on_start_async_task(): """A successful start_async_task tool call must spawn one watcher per task.""" from langgraph.types import Command @@ -142,7 +141,7 @@ def test_middleware_spawns_watcher_on_start_async_task(): ) with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn): - result = asyncio.run(mw.awrap_tool_call(request, fake_handler)) + result = await mw.awrap_tool_call(request, fake_handler) assert isinstance(result, Command) assert spawn_calls == [ @@ -150,7 +149,7 @@ def test_middleware_spawns_watcher_on_start_async_task(): ] -def test_middleware_spawns_watcher_on_update_async_task(): +async def test_middleware_spawns_watcher_on_update_async_task(): """A successful update_async_task call must also spawn a (replacement) watcher.""" from langgraph.types import Command @@ -183,7 +182,7 @@ def test_middleware_spawns_watcher_on_update_async_task(): ) with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn): - asyncio.run(mw.awrap_tool_call(request, fake_handler)) + await mw.awrap_tool_call(request, fake_handler) assert len(spawn_calls) == 1 args, kwargs = spawn_calls[0] @@ -195,7 +194,7 @@ def test_middleware_spawns_watcher_on_update_async_task(): assert kwargs["origin_cli_thread_id"] == "cli-thread-A" -def test_middleware_pre_cancels_old_watcher_on_update(): +async def test_middleware_pre_cancels_old_watcher_on_update(): """update_async_task must cancel the existing watcher BEFORE invoking the handler. Otherwise the new run interrupts the old run's stream, which closes @@ -221,14 +220,14 @@ def test_middleware_pre_cancels_old_watcher_on_update(): try: with patch.object(async_notifier, "spawn_watcher"): - asyncio.run(mw.awrap_tool_call(request, fake_handler)) + await mw.awrap_tool_call(request, fake_handler) finally: async_notifier._watcher_by_thread.pop("task-1", None) assert cancel_observed_before_handler["value"] is True -def test_middleware_passes_through_unrelated_tools(): +async def test_middleware_passes_through_unrelated_tools(): """A non-launch tool call must not spawn any watcher and must return result unchanged.""" mw, _ = _make_middleware() @@ -240,13 +239,13 @@ def test_middleware_passes_through_unrelated_tools(): request = _build_request("ls", {"path": "/"}, thread_id="t") with patch.object(async_notifier, "spawn_watcher") as mock_spawn: - result = asyncio.run(mw.awrap_tool_call(request, fake_handler)) + result = await mw.awrap_tool_call(request, fake_handler) assert result is sentinel assert mock_spawn.call_count == 0 -def test_middleware_handles_non_command_results_gracefully(): +async def test_middleware_handles_non_command_results_gracefully(): """If the launch tool returns a string (validation error), no watcher is spawned.""" mw, _ = _make_middleware() @@ -260,13 +259,13 @@ def test_middleware_handles_non_command_results_gracefully(): ) with patch.object(async_notifier, "spawn_watcher") as mock_spawn: - result = asyncio.run(mw.awrap_tool_call(request, fake_handler)) + result = await mw.awrap_tool_call(request, fake_handler) assert result == "Unknown async subagent type `bogus`" assert mock_spawn.call_count == 0 -def test_middleware_origin_thread_id_is_none_when_runtime_config_missing(): +async def test_middleware_origin_thread_id_is_none_when_runtime_config_missing(): """When runtime.config is empty, origin_cli_thread_id must be None (not crash).""" from langgraph.types import Command @@ -299,12 +298,12 @@ def test_middleware_origin_thread_id_is_none_when_runtime_config_missing(): ) with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn): - asyncio.run(mw.awrap_tool_call(request, fake_handler)) + await mw.awrap_tool_call(request, fake_handler) assert captured.get("origin_cli_thread_id") is None -def test_middleware_swallows_spawn_exceptions(): +async def test_middleware_swallows_spawn_exceptions(): """spawn_watcher errors must not propagate up — middleware logs and continues.""" from langgraph.types import Command @@ -336,7 +335,7 @@ def test_middleware_swallows_spawn_exceptions(): with patch.object(async_notifier, "spawn_watcher", side_effect=boom): # Should not raise. - result = asyncio.run(mw.awrap_tool_call(request, fake_handler)) + result = await mw.awrap_tool_call(request, fake_handler) assert isinstance(result, Command) @@ -356,7 +355,9 @@ def test_middleware_swallows_spawn_exceptions(): ), ], ) -def test_middleware_picks_correct_prompt_field_per_tool(tool_name, args, prompt_field): +async def test_middleware_picks_correct_prompt_field_per_tool( + tool_name, args, prompt_field +): """start_async_task uses 'description'; update_async_task uses 'message'.""" from langgraph.types import Command @@ -385,12 +386,12 @@ def test_middleware_picks_correct_prompt_field_per_tool(tool_name, args, prompt_ request = _build_request(tool_name, args, thread_id="t") with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn): - asyncio.run(mw.awrap_tool_call(request, fake_handler)) + await mw.awrap_tool_call(request, fake_handler) assert captured_prompt["value"] == prompt_field -def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained(): +async def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained(): """update_async_task with extra `description` arg must still use `message`. Guards against the previous `args.get('description') or args.get('message')` @@ -432,12 +433,12 @@ def test_middleware_prompt_field_is_tool_name_gated_not_fallback_chained(): ) with patch.object(async_notifier, "spawn_watcher", side_effect=fake_spawn): - asyncio.run(mw.awrap_tool_call(request, fake_handler)) + await mw.awrap_tool_call(request, fake_handler) assert captured_prompt["value"] == "use this" -def test_middleware_pre_cancel_swallows_unexpected_errors(): +async def test_middleware_pre_cancel_swallows_unexpected_errors(): """A faulty old-watcher handle must not block the handler from running.""" from langgraph.types import Command @@ -460,7 +461,7 @@ def test_middleware_pre_cancel_swallows_unexpected_errors(): try: with patch.object(async_notifier, "spawn_watcher"): # Should not raise. - asyncio.run(mw.awrap_tool_call(request, fake_handler)) + await mw.awrap_tool_call(request, fake_handler) finally: async_notifier._watcher_by_thread.pop("t1", None) diff --git a/tests/test_autoskills.py b/tests/test_autoskills.py index ed1cd3e..20d225e 100644 --- a/tests/test_autoskills.py +++ b/tests/test_autoskills.py @@ -1,6 +1,5 @@ from __future__ import annotations -import asyncio import json from types import SimpleNamespace @@ -917,16 +916,16 @@ class _AsyncFakeCrons: return [{"cron_id": "cron-async"}] -def test_alist_autoskill_schedules_uses_async_client_and_explicit_limit(monkeypatch): +async def test_alist_autoskill_schedules_uses_async_client_and_explicit_limit( + monkeypatch, +): crons = _AsyncFakeCrons() client = SimpleNamespace(crons=crons) monkeypatch.setattr("langgraph_sdk.get_client", lambda **_kwargs: client) - rows = asyncio.run( - alist_autoskill_schedules( - EvoScientistConfig(), - limit=3, - ) + rows = await alist_autoskill_schedules( + EvoScientistConfig(), + limit=3, ) assert rows == [{"cron_id": "cron-async"}] diff --git a/tests/test_background.py b/tests/test_background.py index d7716f8..1bc99cc 100644 --- a/tests/test_background.py +++ b/tests/test_background.py @@ -91,7 +91,7 @@ def test_stop_already_finished_is_graceful(tmp_path): assert "already finished" in bg.stop(pid) -def test_exited_elapsed_is_frozen(tmp_path): +def test_exited_elapsed_is_frozen(tmp_path, monkeypatch): """Elapsed for an exited process freezes at its runtime, it must not keep growing.""" pid = bg.launch(_true_cmd(), str(tmp_path)) assert _wait_until(lambda: bg._PROCESSES[pid].finished_ts is not None) @@ -99,7 +99,7 @@ def test_exited_elapsed_is_frozen(tmp_path): proc = bg._PROCESSES[pid] assert proc.finished_ts is not None first = bg._elapsed(proc) - time.sleep(1.1) # intentional: prove elapsed stays frozen, not ticking up + monkeypatch.setattr(bg.time, "time", lambda: proc.finished_ts + 100.0) assert bg._elapsed(proc) == first diff --git a/tests/test_bus_integration.py b/tests/test_bus_integration.py index fe70165..5f15c9a 100644 --- a/tests/test_bus_integration.py +++ b/tests/test_bus_integration.py @@ -12,7 +12,6 @@ import pytest from EvoScientist.channels.bus.events import InboundMessage from EvoScientist.channels.bus.message_bus import MessageBus from EvoScientist.channels.channel_manager import ChannelManager -from tests.conftest import run_async as _run from tests.fakes import QueueFakeChannel as FakeChannel @@ -58,7 +57,7 @@ def clean_channel_state(): class TestBusInboundConsumer: """Test the _bus_inbound_consumer queue bridge.""" - def test_processes_inbound_and_publishes_outbound(self): + async def test_processes_inbound_and_publishes_outbound(self): """InboundMessage -> queue -> response -> OutboundMessage flow.""" from EvoScientist.cli.channel import ( _bus_inbound_consumer, @@ -68,54 +67,51 @@ class TestBusInboundConsumer: _drain_queue(_message_queue) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="hello agent", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="hello agent", ) + ) - # Wait for consumer to enqueue the message - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + # Wait for consumer to enqueue the message + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - msg = _message_queue.get_nowait() - assert msg.content == "hello agent" - assert msg.sender == "user1" - assert msg.channel_type == "fake" + msg = _message_queue.get_nowait() + assert msg.content == "hello agent" + assert msg.sender == "user1" + assert msg.channel_type == "fake" - # Simulate main-thread response - _set_channel_response(msg.msg_id, "Reply to: hello agent") + # Simulate main-thread response + _set_channel_response(msg.msg_id, "Reply to: hello agent") - outbound = await asyncio.wait_for( - bus.consume_outbound(), - timeout=2.0, - ) - assert outbound.channel == "fake" - assert outbound.chat_id == "chat1" - assert "Reply to: hello agent" in outbound.content + outbound = await asyncio.wait_for( + bus.consume_outbound(), + timeout=2.0, + ) + assert outbound.channel == "fake" + assert outbound.chat_id == "chat1" + assert "Reply to: hello agent" in outbound.content - consumer.cancel() - try: - await consumer - except asyncio.CancelledError: - pass + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass - _run(_test()) - - def test_no_response_fallback(self): + async def test_no_response_fallback(self): """Empty response is replaced with 'No response' fallback.""" from EvoScientist.cli.channel import ( _bus_inbound_consumer, @@ -125,47 +121,44 @@ class TestBusInboundConsumer: _drain_queue(_message_queue) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="test", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="test", ) + ) - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - msg = _message_queue.get_nowait() - # Set empty response — falsy, so consumer falls back to "No response" - _set_channel_response(msg.msg_id, "") + msg = _message_queue.get_nowait() + # Set empty response — falsy, so consumer falls back to "No response" + _set_channel_response(msg.msg_id, "") - outbound = await asyncio.wait_for( - bus.consume_outbound(), - timeout=2.0, - ) - assert outbound.content == "No response" + outbound = await asyncio.wait_for( + bus.consume_outbound(), + timeout=2.0, + ) + assert outbound.content == "No response" - consumer.cancel() - try: - await consumer - except asyncio.CancelledError: - pass + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass - _run(_test()) - - def test_late_response_after_timeout_still_publishes(self, monkeypatch): + async def test_late_response_after_timeout_still_publishes(self, monkeypatch): """A response that arrives after the bridge timeout is still forwarded.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import ( @@ -179,56 +172,53 @@ class TestBusInboundConsumer: _drain_queue(_message_queue) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="slow request", - message_id="msg-123", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="slow request", + message_id="msg-123", ) + ) - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - msg = _message_queue.get_nowait() + msg = _message_queue.get_nowait() - notice = await asyncio.wait_for( - bus.consume_outbound(), - timeout=1.0, - ) - assert "Still working on it" in notice.content - assert notice.reply_to == "msg-123" + notice = await asyncio.wait_for( + bus.consume_outbound(), + timeout=1.0, + ) + assert "Still working on it" in notice.content + assert notice.reply_to == "msg-123" - _set_channel_response(msg.msg_id, "final answer") + _set_channel_response(msg.msg_id, "final answer") - outbound = await asyncio.wait_for( - bus.consume_outbound(), - timeout=1.0, - ) - assert outbound.content == "final answer" - assert outbound.reply_to == "msg-123" + outbound = await asyncio.wait_for( + bus.consume_outbound(), + timeout=1.0, + ) + assert outbound.content == "final answer" + assert outbound.reply_to == "msg-123" - consumer.cancel() - try: - await consumer - except asyncio.CancelledError: - pass + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass - _run(_test()) - - def test_late_timeout_keeps_active_request_cancellable(self, monkeypatch): + async def test_late_timeout_keeps_active_request_cancellable(self, monkeypatch): """Late timeout must not discard an active request's cancel scope.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import ( @@ -243,191 +233,179 @@ class TestBusInboundConsumer: monkeypatch.setattr(channel_mod, "_RESPONSE_TIMEOUT", 0.05) monkeypatch.setattr(channel_mod, "_LATE_RESPONSE_TIMEOUT", 0.05) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - task = asyncio.create_task( - _handle_bus_message( - bus, - manager, - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="still running", - message_id="msg-active", - ), - ) - ) - - queued = None - for _ in range(20): - with _message_queue.mutex: - queued = _message_queue.queue[0] if _message_queue.queue else None - if queued is not None: - break - await asyncio.sleep(0.05) - - assert queued is not None - assert _claim_channel_request(queued) is True - - notice = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0) - assert "Still working on it" in notice.content - - await task - - assert _channel_request_state(queued.msg_id) == "active" - cancel_scope = _channel_message_cancel_scope(queued) - assert not display_mod.is_stream_cancel_requested(cancel_scope) - - await _handle_bus_message( + task = asyncio.create_task( + _handle_bus_message( bus, manager, InboundMessage( channel="fake", sender_id="user1", chat_id="chat1", - content="/stop", - message_id="msg-stop-active", + content="still running", + message_id="msg-active", ), ) + ) - ack = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0) - assert ack.content == "Stopped." - assert ack.reply_to == "msg-stop-active" - assert display_mod.is_stream_cancel_requested(cancel_scope) + queued = None + for _ in range(20): + with _message_queue.mutex: + queued = _message_queue.queue[0] if _message_queue.queue else None + if queued is not None: + break + await asyncio.sleep(0.05) - _run(_test()) + assert queued is not None + assert _claim_channel_request(queued) is True - def test_cancelled_wait_cleans_pending_response(self): + notice = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0) + assert "Still working on it" in notice.content + + await task + + assert _channel_request_state(queued.msg_id) == "active" + cancel_scope = _channel_message_cancel_scope(queued) + assert not display_mod.is_stream_cancel_requested(cancel_scope) + + await _handle_bus_message( + bus, + manager, + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="/stop", + message_id="msg-stop-active", + ), + ) + + ack = await asyncio.wait_for(bus.consume_outbound(), timeout=1.0) + assert ack.content == "Stopped." + assert ack.reply_to == "msg-stop-active" + assert display_mod.is_stream_cancel_requested(cancel_scope) + + async def test_cancelled_wait_cleans_pending_response(self): """Cancelling a pending bus message should not leak its response slot.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import _handle_bus_message, _message_queue - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - task = asyncio.create_task( - _handle_bus_message( - bus, - manager, - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="cancel me", - ), - ) + task = asyncio.create_task( + _handle_bus_message( + bus, + manager, + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="cancel me", + ), ) + ) - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - queued = _message_queue.get_nowait() - with channel_mod._response_lock: - assert queued.msg_id in channel_mod._pending_responses + queued = _message_queue.get_nowait() + with channel_mod._response_lock: + assert queued.msg_id in channel_mod._pending_responses - task.cancel() - with pytest.raises(asyncio.CancelledError): - await task + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task - with channel_mod._response_lock: - assert queued.msg_id not in channel_mod._pending_responses + with channel_mod._response_lock: + assert queued.msg_id not in channel_mod._pending_responses - _run(_test()) - - def test_consumer_shutdown_cleans_pending_response(self): + async def test_consumer_shutdown_cleans_pending_response(self): """Stopping the consumer should cancel late waits and clear state.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="slow shutdown", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="slow shutdown", ) + ) - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - queued = _message_queue.get_nowait() - with channel_mod._response_lock: - assert queued.msg_id in channel_mod._pending_responses + queued = _message_queue.get_nowait() + with channel_mod._response_lock: + assert queued.msg_id in channel_mod._pending_responses - consumer.cancel() - await consumer + consumer.cancel() + await consumer - with channel_mod._response_lock: - assert queued.msg_id not in channel_mod._pending_responses + with channel_mod._response_lock: + assert queued.msg_id not in channel_mod._pending_responses - _run(_test()) - - def test_stop_during_hitl_wait_releases_wait_and_acks(self): + async def test_stop_during_hitl_wait_releases_wait_and_acks(self): """`/stop` should wake pending HITL wait and publish immediate ack.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import _bus_inbound_consumer, _message_queue - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - hitl_event = channel_mod._register_hitl_wait("fake", "chat1") - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + hitl_event = channel_mod._register_hitl_wait("fake", "chat1") + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="/stop", - message_id="m-stop-1", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="/stop", + message_id="m-stop-1", ) + ) - for _ in range(20): - if hitl_event.is_set(): - break - await asyncio.sleep(0.05) - assert hitl_event.is_set() - assert channel_mod._pop_hitl_reply("fake", "chat1") == "/stop" + for _ in range(20): + if hitl_event.is_set(): + break + await asyncio.sleep(0.05) + assert hitl_event.is_set() + assert channel_mod._pop_hitl_reply("fake", "chat1") == "/stop" - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) - assert outbound.content == "Stopped." - assert outbound.reply_to == "m-stop-1" - assert _message_queue.empty() + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) + assert outbound.content == "Stopped." + assert outbound.reply_to == "m-stop-1" + assert _message_queue.empty() - consumer.cancel() - try: - await consumer - except asyncio.CancelledError: - pass + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass - _run(_test()) - - def test_stop_cancels_queued_request_before_main_thread_processes_it(self): + async def test_stop_cancels_queued_request_before_main_thread_processes_it(self): """`/stop` should cancel a queued request instead of only acking.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import ( @@ -436,73 +414,70 @@ class TestBusInboundConsumer: _message_queue, ) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - task = asyncio.create_task( - _handle_bus_message( - bus, - manager, - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="please work", - message_id="m-work-1", - ), - ) - ) - - queued = None - for _ in range(20): - with _message_queue.mutex: - queued = _message_queue.queue[0] if _message_queue.queue else None - if queued is not None: - break - await asyncio.sleep(0.05) - - assert queued is not None - with channel_mod._response_lock: - assert queued.msg_id in channel_mod._pending_responses - - await _handle_bus_message( + task = asyncio.create_task( + _handle_bus_message( bus, manager, InboundMessage( channel="fake", sender_id="user1", chat_id="chat1", - content="/stop", - message_id="m-stop-2", + content="please work", + message_id="m-work-1", ), ) + ) - with pytest.raises(asyncio.CancelledError): - await task + queued = None + for _ in range(20): + with _message_queue.mutex: + queued = _message_queue.queue[0] if _message_queue.queue else None + if queued is not None: + break + await asyncio.sleep(0.05) - skipped = _message_queue.get_nowait() - assert skipped.msg_id == queued.msg_id - assert _claim_or_complete_channel_request(skipped) is False + assert queued is not None + with channel_mod._response_lock: + assert queued.msg_id in channel_mod._pending_responses - with channel_mod._response_lock: - assert queued.msg_id not in channel_mod._pending_responses - with channel_mod._channel_request_lock: - assert queued.msg_id not in channel_mod._channel_requests - assert queued.msg_id not in channel_mod._cancelled_channel_messages - assert "fake:chat1" not in channel_mod._session_requests + await _handle_bus_message( + bus, + manager, + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="/stop", + message_id="m-stop-2", + ), + ) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) - assert outbound.content == "Stopped." - assert outbound.reply_to == "m-stop-2" - with pytest.raises(asyncio.TimeoutError): - await asyncio.wait_for(bus.consume_outbound(), timeout=0.2) + with pytest.raises(asyncio.CancelledError): + await task - _run(_test()) + skipped = _message_queue.get_nowait() + assert skipped.msg_id == queued.msg_id + assert _claim_or_complete_channel_request(skipped) is False - def test_stop_leaves_resolved_response_available_for_delivery(self): + with channel_mod._response_lock: + assert queued.msg_id not in channel_mod._pending_responses + with channel_mod._channel_request_lock: + assert queued.msg_id not in channel_mod._channel_requests + assert queued.msg_id not in channel_mod._cancelled_channel_messages + assert "fake:chat1" not in channel_mod._session_requests + + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) + assert outbound.content == "Stopped." + assert outbound.reply_to == "m-stop-2" + with pytest.raises(asyncio.TimeoutError): + await asyncio.wait_for(bus.consume_outbound(), timeout=0.2) + + async def test_stop_leaves_resolved_response_available_for_delivery(self): """`/stop` must not steal a response whose waiter already resolved.""" from EvoScientist.cli import channel as channel_mod from EvoScientist.cli.channel import ( @@ -515,42 +490,39 @@ class TestBusInboundConsumer: _set_channel_response, ) - async def _test(): - msg = ChannelMessage( - msg_id="msg-resolved", - content="already answered", - sender="user1", - channel_type="fake", - metadata={}, - channel_ref=None, - bus_ref=None, - chat_id="chat1", - message_id="m-resolved", - ) + msg = ChannelMessage( + msg_id="msg-resolved", + content="already answered", + sender="user1", + channel_type="fake", + metadata={}, + channel_ref=None, + bus_ref=None, + chat_id="chat1", + message_id="m-resolved", + ) - waiter = _enqueue_channel_message(msg) - assert _claim_channel_request(msg) is True + waiter = _enqueue_channel_message(msg) + assert _claim_channel_request(msg) is True - _set_channel_response(msg.msg_id, "final answer") - assert await asyncio.wait_for(asyncio.shield(waiter), timeout=1.0) == ( - "final answer" - ) + _set_channel_response(msg.msg_id, "final answer") + assert await asyncio.wait_for(asyncio.shield(waiter), timeout=1.0) == ( + "final answer" + ) - cancelled_count, active_count = _cancel_channel_session("fake", "chat1") - assert cancelled_count == 0 - assert active_count == 0 + cancelled_count, active_count = _cancel_channel_session("fake", "chat1") + assert cancelled_count == 0 + assert active_count == 0 - with channel_mod._response_lock: - assert msg.msg_id in channel_mod._pending_responses - with channel_mod._channel_request_lock: - assert msg.msg_id not in channel_mod._cancelled_channel_messages + with channel_mod._response_lock: + assert msg.msg_id in channel_mod._pending_responses + with channel_mod._channel_request_lock: + assert msg.msg_id not in channel_mod._cancelled_channel_messages - assert _pop_channel_response(msg.msg_id) == "final answer" - _complete_channel_request(msg.msg_id) + assert _pop_channel_response(msg.msg_id) == "final answer" + _complete_channel_request(msg.msg_id) - _run(_test()) - - def test_message_counting(self): + async def test_message_counting(self): """Messages are counted via record_message.""" from EvoScientist.cli.channel import ( _bus_inbound_consumer, @@ -560,45 +532,42 @@ class TestBusInboundConsumer: _drain_queue(_message_queue) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="u1", - chat_id="c1", - content="test", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="u1", + chat_id="c1", + content="test", ) + ) - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - msg = _message_queue.get_nowait() - _set_channel_response(msg.msg_id, "ok") + msg = _message_queue.get_nowait() + _set_channel_response(msg.msg_id, "ok") - await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) + await asyncio.wait_for(bus.consume_outbound(), timeout=2.0) - assert manager._message_counts["fake"]["received"] == 1 - assert manager._message_counts["fake"]["sent"] == 1 + assert manager._message_counts["fake"]["received"] == 1 + assert manager._message_counts["fake"]["sent"] == 1 - consumer.cancel() - try: - await consumer - except asyncio.CancelledError: - pass + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass - _run(_test()) - - def test_channel_message_carries_metadata(self): + async def test_channel_message_carries_metadata(self): """ChannelMessage carries metadata, chat_id, and message_id.""" from EvoScientist.cli.channel import ( _bus_inbound_consumer, @@ -608,49 +577,46 @@ class TestBusInboundConsumer: _drain_queue(_message_queue) - async def _test(): - bus = MessageBus() - manager = ChannelManager(bus) - ch = FakeChannel() - manager.register(ch) + bus = MessageBus() + manager = ChannelManager(bus) + ch = FakeChannel() + manager.register(ch) - consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) + consumer = asyncio.create_task(_bus_inbound_consumer(bus, manager)) - await bus.publish_inbound( - InboundMessage( - channel="fake", - sender_id="user1", - chat_id="chat1", - content="with metadata", - metadata={"key": "value"}, - message_id="msg-123", - ) + await bus.publish_inbound( + InboundMessage( + channel="fake", + sender_id="user1", + chat_id="chat1", + content="with metadata", + metadata={"key": "value"}, + message_id="msg-123", ) + ) - for _ in range(20): - if not _message_queue.empty(): - break - await asyncio.sleep(0.05) + for _ in range(20): + if not _message_queue.empty(): + break + await asyncio.sleep(0.05) - msg = _message_queue.get_nowait() - assert msg.content == "with metadata" - assert msg.metadata == {"key": "value"} - assert msg.chat_id == "chat1" - assert msg.message_id == "msg-123" - assert msg.channel_ref is ch + msg = _message_queue.get_nowait() + assert msg.content == "with metadata" + assert msg.metadata == {"key": "value"} + assert msg.chat_id == "chat1" + assert msg.message_id == "msg-123" + assert msg.channel_ref is ch - _set_channel_response(msg.msg_id, "done") + _set_channel_response(msg.msg_id, "done") - outbound = await asyncio.wait_for( - bus.consume_outbound(), - timeout=2.0, - ) - assert outbound.reply_to == "msg-123" + outbound = await asyncio.wait_for( + bus.consume_outbound(), + timeout=2.0, + ) + assert outbound.reply_to == "msg-123" - consumer.cancel() - try: - await consumer - except asyncio.CancelledError: - pass - - _run(_test()) + consumer.cancel() + try: + await consumer + except asyncio.CancelledError: + pass diff --git a/tests/test_channel_command.py b/tests/test_channel_command.py index a466785..91a31c8 100644 --- a/tests/test_channel_command.py +++ b/tests/test_channel_command.py @@ -3,8 +3,6 @@ from types import SimpleNamespace from unittest.mock import MagicMock, patch -from tests.conftest import run_async as _run - def _ctx(): from EvoScientist.commands.base import ChannelRuntime, CommandContext @@ -55,7 +53,7 @@ class TestNeedsAgent: class TestStartPath: """Start flow must propagate agent/thread_id globals.""" - def test_start_binds_channel_runtime(self): + async def test_start_binds_channel_runtime(self): from EvoScientist.commands.implementation.channel import ChannelCommand ctx, _ui = _ctx() @@ -77,11 +75,11 @@ class TestStartPath: return_value=config, ), ): - _run(ChannelCommand().execute(ctx, ["telegram"])) + await ChannelCommand().execute(ctx, ["telegram"]) assert ctx.channel_runtime.agent is ctx.agent assert ctx.channel_runtime.thread_id == "tid-42" - def test_start_propagates_send_thinking(self): + async def test_start_propagates_send_thinking(self): """send_thinking flag must reach _start_channels_bus_mode.""" from EvoScientist.commands.implementation.channel import ChannelCommand @@ -111,14 +109,14 @@ class TestStartPath: return_value=config, ), ): - _run(ChannelCommand().execute(ctx, ["telegram"])) + await ChannelCommand().execute(ctx, ["telegram"]) assert captured["agent"] is ctx.agent assert captured["thread_id"] == "tid-42" assert captured["send_thinking"] is False class TestAddToRunningPath: - def test_add_to_running_binds_channel_runtime(self): + async def test_add_to_running_binds_channel_runtime(self): from EvoScientist.commands.implementation.channel import ChannelCommand ctx, _ui = _ctx() @@ -140,11 +138,11 @@ class TestAddToRunningPath: return_value=config, ), ): - _run(ChannelCommand().execute(ctx, ["discord"])) + await ChannelCommand().execute(ctx, ["discord"]) assert ctx.channel_runtime.agent is ctx.agent assert ctx.channel_runtime.thread_id == "tid-42" - def test_add_to_running_propagates_send_thinking(self): + async def test_add_to_running_propagates_send_thinking(self): """Adding to a running bus must honor config.channel_send_thinking.""" from EvoScientist.commands.implementation.channel import ChannelCommand @@ -173,13 +171,13 @@ class TestAddToRunningPath: return_value=config, ), ): - _run(ChannelCommand().execute(ctx, ["discord"])) + await ChannelCommand().execute(ctx, ["discord"]) assert captured["channel_type"] == "discord" assert captured["send_thinking"] is True class TestStatusPath: - def test_status_without_running_channels(self): + async def test_status_without_running_channels(self): from EvoScientist.commands.implementation.channel import ChannelCommand ctx, ui = _ctx() @@ -198,6 +196,6 @@ class TestStatusPath: return_value=config, ), ): - _run(ChannelCommand().execute(ctx, ["status"])) + await ChannelCommand().execute(ctx, ["status"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("No messaging channels" in m for m in msgs) diff --git a/tests/test_channel_command_ui.py b/tests/test_channel_command_ui.py index 2d65a17..a5c8298 100644 --- a/tests/test_channel_command_ui.py +++ b/tests/test_channel_command_ui.py @@ -8,7 +8,6 @@ import pytest from EvoScientist.commands.channel_ui import ChannelCommandUI from EvoScientist.gateway import ThreadStore -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -57,7 +56,7 @@ def _sent_text(bus_ref) -> str: ) -def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate(): +async def test_handle_session_resume_sends_history_back_to_channel_without_local_duplicate(): callback = AsyncMock() bus_ref = SimpleNamespace(publish_outbound=AsyncMock()) @@ -72,7 +71,7 @@ def test_handle_session_resume_sends_history_back_to_channel_without_local_dupli thread_store=thread_store, ) - _run(_run_resume(ui, "thread-42", "/workspace")) + await _run_resume(ui, "thread-42", "/workspace") callback.assert_awaited_once_with("thread-42", "/workspace") assert thread_store.calls == [("get_thread_messages", "thread-42")] @@ -84,7 +83,7 @@ def test_handle_session_resume_sends_history_back_to_channel_without_local_dupli assert "EvoScientist: Here is the saved answer." in text -def test_handle_session_resume_propagates_callback_abort_without_history(): +async def test_handle_session_resume_propagates_callback_abort_without_history(): callback = AsyncMock(side_effect=RuntimeError("workspace conflict")) bus_ref = SimpleNamespace(publish_outbound=AsyncMock()) thread_store = FakeThreadStore() @@ -95,7 +94,7 @@ def test_handle_session_resume_propagates_callback_abort_without_history(): ) with pytest.raises(RuntimeError, match="workspace conflict"): - _run(_run_resume(ui, "thread-42", "/workspace")) + await _run_resume(ui, "thread-42", "/workspace") callback.assert_awaited_once_with("thread-42", "/workspace") assert thread_store.calls == [] @@ -103,7 +102,7 @@ def test_handle_session_resume_propagates_callback_abort_without_history(): assert captured == [] -def test_handle_session_resume_reports_history_load_error(): +async def test_handle_session_resume_reports_history_load_error(): callback = AsyncMock() bus_ref = SimpleNamespace(publish_outbound=AsyncMock()) ui, captured = _make_ui( @@ -114,7 +113,7 @@ def test_handle_session_resume_reports_history_load_error(): ), ) - _run(_run_resume(ui, "thread-42", "/workspace")) + await _run_resume(ui, "thread-42", "/workspace") callback.assert_awaited_once_with("thread-42", "/workspace") assert captured == [] @@ -123,7 +122,7 @@ def test_handle_session_resume_reports_history_load_error(): assert "history unavailable: db locked" in text -def test_handle_session_resume_distinguishes_non_displayable_messages(): +async def test_handle_session_resume_distinguishes_non_displayable_messages(): bus_ref = SimpleNamespace(publish_outbound=AsyncMock()) ui, captured = _make_ui( bus_ref=bus_ref, @@ -132,7 +131,7 @@ def test_handle_session_resume_distinguishes_non_displayable_messages(): ), ) - _run(_run_resume(ui, "thread-42", "/workspace")) + await _run_resume(ui, "thread-42", "/workspace") assert captured == [ "Resumed session: thread-42\nNo displayable messages in this session." diff --git a/tests/test_channel_comprehensive.py b/tests/test_channel_comprehensive.py index c22d941..274c29a 100644 --- a/tests/test_channel_comprehensive.py +++ b/tests/test_channel_comprehensive.py @@ -15,7 +15,6 @@ Test groups: from __future__ import annotations import asyncio -import time from datetime import datetime from unittest.mock import AsyncMock, MagicMock @@ -44,10 +43,37 @@ from EvoScientist.channels.retry import RetryConfig, RetryInfo, retry_async # ═══════════════════════════════════════════════════════════════════ # Helpers # ═══════════════════════════════════════════════════════════════════ -from tests.conftest import run_async as _run from tests.fakes import FakeChannelConfig as _FakeConfig from tests.fakes import FakeGraphGateway, StubChannel + +class ManualClock: + def __init__(self) -> None: + self._now = 0.0 + + def __call__(self) -> float: + return self._now + + def advance(self, seconds: float) -> None: + self._now += seconds + + +async def _flush_debounce(ch: StubChannel, sender: str) -> None: + task = ch._debounce_tasks.get(sender) + if task is not None: + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + await ch._process_buffered_messages(sender) + + +async def _wait_for_async(predicate) -> None: + while not predicate(): + await asyncio.sleep(0) + + # ═══════════════════════════════════════════════════════════════════ # 1. DedupCache # ═══════════════════════════════════════════════════════════════════ @@ -69,9 +95,10 @@ class TestDedupCache: assert dc.is_duplicate("") is False def test_ttl_expiry(self): - dc = DedupCache(ttl_seconds=0.05) + clock = ManualClock() + dc = DedupCache(ttl_seconds=0.05, clock=clock) dc.is_duplicate("msg_001") - time.sleep(0.1) + clock.advance(0.051) # After TTL, the entry should be pruned assert dc.is_duplicate("msg_001") is False @@ -109,7 +136,7 @@ class TestDedupCache: class TestRetryAsync: - def test_success_on_first_attempt(self): + async def test_success_on_first_attempt(self): call_count = 0 async def _fn(): @@ -117,11 +144,11 @@ class TestRetryAsync: call_count += 1 return "ok" - result = _run(retry_async(_fn)) + result = await retry_async(_fn) assert result == "ok" assert call_count == 1 - def test_retries_on_failure_then_succeeds(self): + async def test_retries_on_failure_then_succeeds(self): attempts = [] async def _fn(): @@ -130,28 +157,24 @@ class TestRetryAsync: raise RuntimeError("transient") return "recovered" - result = _run( - retry_async( - _fn, - config=RetryConfig(attempts=5, min_delay_s=0.01, max_delay_s=0.05), - ) + result = await retry_async( + _fn, + config=RetryConfig(attempts=5, min_delay_s=0.01, max_delay_s=0.05), ) assert result == "recovered" assert len(attempts) == 3 - def test_exhausts_retries_raises(self): + async def test_exhausts_retries_raises(self): async def _fn(): raise ValueError("permanent") with pytest.raises(ValueError, match="permanent"): - _run( - retry_async( - _fn, - config=RetryConfig(attempts=2, min_delay_s=0.01), - ) + await retry_async( + _fn, + config=RetryConfig(attempts=2, min_delay_s=0.01), ) - def test_should_retry_false_aborts(self): + async def test_should_retry_false_aborts(self): """[B-01] should_retry returning False should abort immediately.""" call_count = 0 @@ -161,16 +184,14 @@ class TestRetryAsync: raise PermissionError("forbidden") with pytest.raises(PermissionError): - _run( - retry_async( - _fn, - config=RetryConfig(attempts=5, min_delay_s=0.01), - should_retry=lambda exc, _: False, - ) + await retry_async( + _fn, + config=RetryConfig(attempts=5, min_delay_s=0.01), + should_retry=lambda exc, _: False, ) assert call_count == 1 # No retry happened - def test_server_retry_after_respected(self): + async def test_server_retry_after_respected(self): """retry_after_s callback provides server-supplied delay.""" delays = [] @@ -182,20 +203,16 @@ class TestRetryAsync: def _on_retry(info: RetryInfo): delays.append(info.delay_s) - _run( - retry_async( - _fn, - config=RetryConfig( - attempts=3, min_delay_s=0.01, max_delay_s=10, jitter=0 - ), - retry_after_s=lambda _: 0.5, - on_retry=_on_retry, - ) + await retry_async( + _fn, + config=RetryConfig(attempts=3, min_delay_s=0.01, max_delay_s=10, jitter=0), + retry_after_s=lambda _: 0.5, + on_retry=_on_retry, ) assert len(delays) == 1 assert delays[0] >= 0.5 - def test_jitter_applied(self): + async def test_jitter_applied(self): """With jitter > 0, delays should vary.""" delays = [] @@ -204,14 +221,12 @@ class TestRetryAsync: raise RuntimeError("fail") return "ok" - _run( - retry_async( - _fn, - config=RetryConfig( - attempts=10, min_delay_s=0.01, max_delay_s=1.0, jitter=0.5 - ), - on_retry=lambda info: delays.append(info.delay_s), - ) + await retry_async( + _fn, + config=RetryConfig( + attempts=10, min_delay_s=0.01, max_delay_s=1.0, jitter=0.5 + ), + on_retry=lambda info: delays.append(info.delay_s), ) # With 50% jitter, not all delays should be identical if len(delays) > 1: @@ -426,103 +441,88 @@ class TestMarkdownUtils: class TestChannelSend: - def test_send_single_chunk(self): - async def _test(): - ch = StubChannel() - msg = OutboundMessage( - channel="stub", - chat_id="c1", - content="hello", - metadata={"chat_id": "c1"}, - ) - ok = await ch.send(msg) - assert ok is True - assert len(ch._sent_chunks) == 1 - assert ch._sent_chunks[0][0] == "c1" - assert ch._sent_chunks[0][2] == "hello" # raw + async def test_send_single_chunk(self): + ch = StubChannel() + msg = OutboundMessage( + channel="stub", + chat_id="c1", + content="hello", + metadata={"chat_id": "c1"}, + ) + ok = await ch.send(msg) + assert ok is True + assert len(ch._sent_chunks) == 1 + assert ch._sent_chunks[0][0] == "c1" + assert ch._sent_chunks[0][2] == "hello" # raw - _run(_test()) + async def test_send_multi_chunk(self): + cfg = _FakeConfig(text_chunk_limit=10) + ch = StubChannel(cfg) + msg = OutboundMessage( + channel="stub", + chat_id="c1", + content="hello world this is a long message", + metadata={"chat_id": "c1"}, + ) + ok = await ch.send(msg) + assert ok is True + assert len(ch._sent_chunks) > 1 - def test_send_multi_chunk(self): - async def _test(): - cfg = _FakeConfig(text_chunk_limit=10) - ch = StubChannel(cfg) - msg = OutboundMessage( - channel="stub", - chat_id="c1", - content="hello world this is a long message", - metadata={"chat_id": "c1"}, - ) - ok = await ch.send(msg) - assert ok is True - assert len(ch._sent_chunks) > 1 + async def test_send_returns_false_when_not_ready(self): + ch = StubChannel() + ch._is_ready = lambda: False + msg = OutboundMessage(channel="stub", chat_id="c1", content="hi") + ok = await ch.send(msg) + assert ok is False - _run(_test()) - - def test_send_returns_false_when_not_ready(self): - async def _test(): - ch = StubChannel() - ch._is_ready = lambda: False - msg = OutboundMessage(channel="stub", chat_id="c1", content="hi") - ok = await ch.send(msg) - assert ok is False - - _run(_test()) - - def test_send_per_chat_lock_serializes(self): + async def test_send_per_chat_lock_serializes(self): """[B-03] Per-chat locks prevent message reordering.""" - async def _test(): - ch = StubChannel() - order = [] + ch = StubChannel() + order = [] - original_send_chunk = ch._send_chunk + original_send_chunk = ch._send_chunk - async def slow_send(chat_id, fmt, raw, reply_to, meta): - order.append(raw) - await asyncio.sleep(0.05) - await original_send_chunk(chat_id, fmt, raw, reply_to, meta) + async def slow_send(chat_id, fmt, raw, reply_to, meta): + order.append(raw) + await asyncio.sleep(0.05) + await original_send_chunk(chat_id, fmt, raw, reply_to, meta) - ch._send_chunk = slow_send + ch._send_chunk = slow_send - msg1 = OutboundMessage( - channel="stub", - chat_id="c1", - content="first", - metadata={"chat_id": "c1"}, - ) - msg2 = OutboundMessage( - channel="stub", - chat_id="c1", - content="second", - metadata={"chat_id": "c1"}, - ) + msg1 = OutboundMessage( + channel="stub", + chat_id="c1", + content="first", + metadata={"chat_id": "c1"}, + ) + msg2 = OutboundMessage( + channel="stub", + chat_id="c1", + content="second", + metadata={"chat_id": "c1"}, + ) - await asyncio.gather(ch.send(msg1), ch.send(msg2)) - # Both complete; order may vary but no interleaving within a single send - assert len(order) == 2 + await asyncio.gather(ch.send(msg1), ch.send(msg2)) + # Both complete; order may vary but no interleaving within a single send + assert len(order) == 2 - _run(_test()) - - def test_reply_to_only_on_first_chunk(self): + async def test_reply_to_only_on_first_chunk(self): """reply_to should only be passed to the first chunk.""" - async def _test(): - cfg = _FakeConfig(text_chunk_limit=10) - ch = StubChannel(cfg) - msg = OutboundMessage( - channel="stub", - chat_id="c1", - content="a very long message that will be split into multiple parts", - reply_to="msg_42", - metadata={"chat_id": "c1"}, - ) - await ch.send(msg) - reply_tos = [c[3] for c in ch._sent_chunks] - assert reply_tos[0] == "msg_42" - assert all(r is None for r in reply_tos[1:]) - - _run(_test()) + cfg = _FakeConfig(text_chunk_limit=10) + ch = StubChannel(cfg) + msg = OutboundMessage( + channel="stub", + chat_id="c1", + content="a very long message that will be split into multiple parts", + reply_to="msg_42", + metadata={"chat_id": "c1"}, + ) + await ch.send(msg) + reply_tos = [c[3] for c in ch._sent_chunks] + assert reply_tos[0] == "msg_42" + assert all(r is None for r in reply_tos[1:]) class TestChannelAllowList: @@ -608,25 +608,19 @@ class TestChannelBuildInbound: assert msg.content == "hello" assert msg.media == ["/path/img.jpg"] - def test_drops_disallowed_sender(self): - async def _test(): - cfg = _FakeConfig(allowed_senders=["alice"]) - ch = StubChannel(cfg) - raw = RawIncoming(sender_id="eve", chat_id="c1", text="hack") - await ch._enqueue_raw(raw) - assert ch._queue.qsize() == 0 + async def test_drops_disallowed_sender(self): + cfg = _FakeConfig(allowed_senders=["alice"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="eve", chat_id="c1", text="hack") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 0 - _run(_test()) - - def test_drops_disallowed_channel(self): - async def _test(): - cfg = _FakeConfig(allowed_channels=["c1"]) - ch = StubChannel(cfg) - raw = RawIncoming(sender_id="u1", chat_id="c2", text="hello") - await ch._enqueue_raw(raw) - assert ch._queue.qsize() == 0 - - _run(_test()) + async def test_drops_disallowed_channel(self): + cfg = _FakeConfig(allowed_channels=["c1"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="u1", chat_id="c2", text="hello") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 0 def test_drops_empty_content_no_media(self): ch = StubChannel() @@ -669,267 +663,221 @@ class TestChannelBuildInbound: class TestInboundPipeline: """Tests for the new middleware-based inbound pipeline in _enqueue_raw().""" - def test_pipeline_dedup(self): + async def test_pipeline_dedup(self): """Duplicate messages are dropped by the pipeline.""" - async def _test(): - ch = StubChannel() - raw = RawIncoming( - sender_id="u1", chat_id="c1", text="hello", message_id="m1" - ) - await ch._enqueue_raw(raw) - await ch._enqueue_raw(raw) - assert ch._queue.qsize() == 1 + ch = StubChannel() + raw = RawIncoming(sender_id="u1", chat_id="c1", text="hello", message_id="m1") + await ch._enqueue_raw(raw) + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 1 - _run(_test()) - - def test_pipeline_allowlist_blocks(self): + async def test_pipeline_allowlist_blocks(self): """Non-allowed senders are blocked by the pipeline.""" - async def _test(): - cfg = _FakeConfig(allowed_senders=["alice"]) - ch = StubChannel(cfg) - raw = RawIncoming(sender_id="eve", chat_id="c1", text="hack") - await ch._enqueue_raw(raw) - assert ch._queue.qsize() == 0 + cfg = _FakeConfig(allowed_senders=["alice"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="eve", chat_id="c1", text="hack") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 0 - _run(_test()) - - def test_pipeline_allowlist_passes(self): + async def test_pipeline_allowlist_passes(self): """Allowed senders pass through the pipeline.""" - async def _test(): - cfg = _FakeConfig(allowed_senders=["alice"]) - ch = StubChannel(cfg) - raw = RawIncoming(sender_id="alice", chat_id="c1", text="hello") - await ch._enqueue_raw(raw) - assert ch._queue.qsize() == 1 + cfg = _FakeConfig(allowed_senders=["alice"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="alice", chat_id="c1", text="hello") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 1 - _run(_test()) - - def test_pipeline_channel_allowlist_blocks(self): + async def test_pipeline_channel_allowlist_blocks(self): """Non-allowed channels are blocked by the pipeline.""" - async def _test(): - cfg = _FakeConfig(allowed_channels=["c1"]) - ch = StubChannel(cfg) - raw = RawIncoming(sender_id="u1", chat_id="c2", text="hello") - await ch._enqueue_raw(raw) - assert ch._queue.qsize() == 0 + cfg = _FakeConfig(allowed_channels=["c1"]) + ch = StubChannel(cfg) + raw = RawIncoming(sender_id="u1", chat_id="c2", text="hello") + await ch._enqueue_raw(raw) + assert ch._queue.qsize() == 0 - _run(_test()) - - def test_pipeline_inbound_has_is_group(self): + async def test_pipeline_inbound_has_is_group(self): """InboundMessage carries is_group and was_mentioned from RawIncoming.""" - async def _test(): - ch = StubChannel() - raw = RawIncoming( - sender_id="u1", - chat_id="c1", - text="hello", - is_group=True, - was_mentioned=True, - ) - await ch._enqueue_raw(raw) - msg = await ch._queue.get() - assert msg.is_group is True - assert msg.was_mentioned is True - - _run(_test()) + ch = StubChannel() + raw = RawIncoming( + sender_id="u1", + chat_id="c1", + text="hello", + is_group=True, + was_mentioned=True, + ) + await ch._enqueue_raw(raw) + msg = await ch._queue.get() + assert msg.is_group is True + assert msg.was_mentioned is True class TestChannelDebounce: - def test_single_message_processed(self): + async def test_single_message_processed(self): """A single message should be published after debounce delay.""" - async def _test(): - bus = MessageBus() - ch = StubChannel() - ch.set_bus(bus) - ch.initial_debounce = 0.05 - ch.max_debounce = 0.1 + bus = MessageBus() + ch = StubChannel() + ch.set_bus(bus) + ch.initial_debounce = 0.05 + ch.max_debounce = 0.1 + msg = InboundMessage( + channel="stub", + sender_id="u1", + chat_id="c1", + content="hello", + message_id="m1", + metadata={"chat_id": "c1"}, + ) + await ch.queue_message(msg) + await _flush_debounce(ch, "u1") + + # Check bus received the message + assert bus.inbound.qsize() == 1 + received = await bus.consume_inbound() + assert received.content == "hello" + + async def test_rapid_messages_merged(self): + """[B-05] Multiple rapid messages should be merged.""" + + bus = MessageBus() + ch = StubChannel() + ch.set_bus(bus) + ch.initial_debounce = 0.1 + ch.max_debounce = 0.3 + + for i in range(3): msg = InboundMessage( channel="stub", sender_id="u1", chat_id="c1", - content="hello", - message_id="m1", + content=f"part{i}", + message_id=f"m{i}", metadata={"chat_id": "c1"}, ) await ch.queue_message(msg) - await asyncio.sleep(0.2) - # Check bus received the message - assert bus.inbound.qsize() == 1 - received = await bus.consume_inbound() - assert received.content == "hello" + await _flush_debounce(ch, "u1") + assert bus.inbound.qsize() == 1 + received = await bus.consume_inbound() + assert "part0" in received.content + assert "part1" in received.content + assert "part2" in received.content - _run(_test()) - - def test_rapid_messages_merged(self): - """[B-05] Multiple rapid messages should be merged.""" - - async def _test(): - bus = MessageBus() - ch = StubChannel() - ch.set_bus(bus) - ch.initial_debounce = 0.1 - ch.max_debounce = 0.3 - - for i in range(3): - msg = InboundMessage( - channel="stub", - sender_id="u1", - chat_id="c1", - content=f"part{i}", - message_id=f"m{i}", - metadata={"chat_id": "c1"}, - ) - await ch.queue_message(msg) - await asyncio.sleep(0.01) - - await asyncio.sleep(0.5) - assert bus.inbound.qsize() == 1 - received = await bus.consume_inbound() - assert "part0" in received.content - assert "part1" in received.content - assert "part2" in received.content - - _run(_test()) - - def test_dedup_skips_duplicate(self): + async def test_dedup_skips_duplicate(self): """Dedup is now handled in _enqueue_raw pipeline, not queue_message.""" - async def _test(): - ch = StubChannel() + ch = StubChannel() - raw = RawIncoming( - sender_id="u1", - chat_id="c1", - text="hello", - message_id="m1", - ) - await ch._enqueue_raw(raw) - await ch._enqueue_raw(raw) # duplicate + raw = RawIncoming( + sender_id="u1", + chat_id="c1", + text="hello", + message_id="m1", + ) + await ch._enqueue_raw(raw) + await ch._enqueue_raw(raw) # duplicate - # Only one should be enqueued (dedup catches second) - assert ch._queue.qsize() == 1 + # Only one should be enqueued (dedup catches second) + assert ch._queue.qsize() == 1 - _run(_test()) - - def test_debounce_metadata_from_first_message(self): + async def test_debounce_metadata_from_first_message(self): """[B-05] Metadata from the first message in a debounce window is kept.""" - async def _test(): - bus = MessageBus() - ch = StubChannel() - ch.set_bus(bus) - ch.initial_debounce = 0.1 + bus = MessageBus() + ch = StubChannel() + ch.set_bus(bus) + ch.initial_debounce = 0.1 - msg1 = InboundMessage( - channel="stub", - sender_id="u1", - chat_id="c1", - content="first", - message_id="m1", - metadata={"chat_id": "c1", "key": "val1"}, - ) - msg2 = InboundMessage( - channel="stub", - sender_id="u1", - chat_id="c1", - content="second", - message_id="m2", - metadata={"chat_id": "c2", "key": "val2"}, - ) - await ch.queue_message(msg1) - await asyncio.sleep(0.01) - await ch.queue_message(msg2) - await asyncio.sleep(0.3) + msg1 = InboundMessage( + channel="stub", + sender_id="u1", + chat_id="c1", + content="first", + message_id="m1", + metadata={"chat_id": "c1", "key": "val1"}, + ) + msg2 = InboundMessage( + channel="stub", + sender_id="u1", + chat_id="c1", + content="second", + message_id="m2", + metadata={"chat_id": "c2", "key": "val2"}, + ) + await ch.queue_message(msg1) + await ch.queue_message(msg2) + await _flush_debounce(ch, "u1") - received = await bus.consume_inbound() - # BUG: metadata is from msg1 only; msg2's metadata is lost - assert received.metadata["key"] == "val1" - - _run(_test()) + received = await bus.consume_inbound() + # BUG: metadata is from msg1 only; msg2's metadata is lost + assert received.metadata["key"] == "val1" class TestChannelTyping: - def test_start_and_stop_typing(self): - async def _test(): - ch = StubChannel() - await ch.start_typing("c1") - assert "c1" in ch._typing_tasks - await asyncio.sleep(0.1) - await ch.stop_typing("c1") - assert "c1" not in ch._typing_tasks + async def test_start_and_stop_typing(self): + ch = StubChannel() + await ch.start_typing("c1") + assert "c1" in ch._typing_tasks + await asyncio.sleep(0) + await ch.stop_typing("c1") + assert "c1" not in ch._typing_tasks - _run(_test()) + async def test_double_start_cancels_previous(self): + ch = StubChannel() + await ch.start_typing("c1") + task1 = ch._typing_tasks["c1"] + await ch.start_typing("c1") + task2 = ch._typing_tasks["c1"] + assert task1 is not task2 + # Allow the event loop to process the cancellation + await asyncio.sleep(0) + assert task1.cancelled() or task1.done() + await ch.stop_typing("c1") - def test_double_start_cancels_previous(self): - async def _test(): - ch = StubChannel() - await ch.start_typing("c1") - task1 = ch._typing_tasks["c1"] - await ch.start_typing("c1") - task2 = ch._typing_tasks["c1"] - assert task1 is not task2 - # Allow the event loop to process the cancellation - await asyncio.sleep(0) - assert task1.cancelled() or task1.done() - await ch.stop_typing("c1") - - _run(_test()) - - def test_stop_typing_idempotent(self): - async def _test(): - ch = StubChannel() - # Should not raise even if never started - await ch.stop_typing("nonexistent") - - _run(_test()) + async def test_stop_typing_idempotent(self): + ch = StubChannel() + # Should not raise even if never started + await ch.stop_typing("nonexistent") class TestChannelReconnect: - def test_run_reconnects_on_error(self): + async def test_run_reconnects_on_error(self): """Channel.run() should reconnect with backoff on transient errors.""" - async def _test(): - ch = StubChannel() - start_count = 0 - original_start = ch.start + ch = StubChannel() + start_count = 0 + original_start = ch.start - async def flaky_start(): - nonlocal start_count - start_count += 1 - if start_count <= 2: - raise ConnectionError("transient") - await original_start() - # Stop after successful start to end the test - ch._running = False + async def flaky_start(): + nonlocal start_count + start_count += 1 + if start_count <= 2: + raise ConnectionError("transient") + await original_start() + # Stop after successful start to end the test + ch._running = False - ch.start = flaky_start - await ch.run() - assert start_count == 3 + ch.start = flaky_start + await ch.run() + assert start_count == 3 - _run(_test()) - - def test_run_stops_on_channel_error(self): + async def test_run_stops_on_channel_error(self): """ChannelError should stop the channel permanently.""" - async def _test(): - ch = StubChannel() + ch = StubChannel() - async def fatal_start(): - raise ChannelError("fatal") + async def fatal_start(): + raise ChannelError("fatal") - ch.start = fatal_start - await ch.run() - assert ch._running is False - - _run(_test()) + ch.start = fatal_start + await ch.run() + assert ch._running is False class TestExtractRetryAfter: @@ -969,14 +917,11 @@ class TestChannelAttachments: assert result is not None assert "too large" in result - def test_send_media_returns_false_when_not_ready(self): - async def _test(): - ch = StubChannel() - ch._is_ready = lambda: False - ok = await ch.send_media("r1", "/path/file.txt") - assert ok is False - - _run(_test()) + async def test_send_media_returns_false_when_not_ready(self): + ch = StubChannel() + ch._is_ready = lambda: False + ok = await ch.send_media("r1", "/path/file.txt") + assert ok is False # ═══════════════════════════════════════════════════════════════════ @@ -1023,129 +968,135 @@ class TestChannelManagerRegister: class TestChannelManagerDispatch: - def test_dispatch_routes_to_channel(self): - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - ch = StubChannel() - # Override send to track calls - sent = [] - ch.send = AsyncMock(side_effect=lambda m: sent.append(m) or True) - mgr.register(ch) + async def test_dispatch_routes_to_channel(self): + bus = MessageBus() + mgr = ChannelManager(bus) + ch = StubChannel() + # Override send to track calls + sent = [] + sent_event = asyncio.Event() - task = asyncio.create_task(mgr._dispatch_outbound()) - await bus.publish_outbound( - OutboundMessage( - channel="stub", - chat_id="c1", - content="hello", - ) + async def send(msg): + sent.append(msg) + sent_event.set() + return True + + ch.send = send + mgr.register(ch) + + task = asyncio.create_task(mgr._dispatch_outbound()) + await bus.publish_outbound( + OutboundMessage( + channel="stub", + chat_id="c1", + content="hello", ) - await asyncio.sleep(0.1) - task.cancel() - try: - await task - except asyncio.CancelledError: - pass + ) + await asyncio.wait_for(sent_event.wait(), timeout=1.0) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass - assert len(sent) == 1 - assert sent[0].content == "hello" + assert len(sent) == 1 + assert sent[0].content == "hello" - _run(_test()) - - def test_dispatch_unknown_channel_logged(self): + async def test_dispatch_unknown_channel_logged(self): """Messages to unknown channels should be logged, not crash.""" - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) + bus = MessageBus() + mgr = ChannelManager(bus) - task = asyncio.create_task(mgr._dispatch_outbound()) - await bus.publish_outbound( - OutboundMessage( - channel="nonexistent", - chat_id="c1", - content="hello", - ) + task = asyncio.create_task(mgr._dispatch_outbound()) + await bus.publish_outbound( + OutboundMessage( + channel="nonexistent", + chat_id="c1", + content="hello", ) - await asyncio.sleep(0.1) - task.cancel() - try: - await task - except asyncio.CancelledError: - pass - # Should not raise + ) + await asyncio.wait_for( + _wait_for_async(lambda: bus.outbound_size == 0), + timeout=1.0, + ) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass + # Should not raise - _run(_test()) - - def test_dispatch_send_return_false_counts_failure(self): + async def test_dispatch_send_return_false_counts_failure(self): """send() returning False should mark the delivery as failed.""" - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - ch = StubChannel() + bus = MessageBus() + mgr = ChannelManager(bus) + ch = StubChannel() + failed_event = asyncio.Event() - async def failing_send(msg): - return False # Indicates failure + async def failing_send(msg): + failed_event.set() + return False # Indicates failure - ch.send = failing_send - mgr.register(ch) + ch.send = failing_send + mgr.register(ch) - task = asyncio.create_task(mgr._dispatch_outbound()) - await bus.publish_outbound( - OutboundMessage( - channel="stub", - chat_id="c1", - content="hello", - ) + task = asyncio.create_task(mgr._dispatch_outbound()) + await bus.publish_outbound( + OutboundMessage( + channel="stub", + chat_id="c1", + content="hello", ) - await asyncio.sleep(0.1) - task.cancel() - try: - await task - except asyncio.CancelledError: - pass + ) + await asyncio.wait_for(failed_event.wait(), timeout=1.0) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass - health = mgr._health["stub"] - assert health.total_successes == 0 - assert health.total_failures == 1 - assert health.consecutive_failures == 1 + health = mgr._health["stub"] + assert health.total_successes == 0 + assert health.total_failures == 1 + assert health.consecutive_failures == 1 - _run(_test()) - - def test_dispatch_send_media_return_false_counts_failure(self): + async def test_dispatch_send_media_return_false_counts_failure(self): """send_media() returning False should mark the delivery as failed.""" - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - ch = StubChannel() - ch.send_media = AsyncMock(return_value=False) - mgr.register(ch) + bus = MessageBus() + mgr = ChannelManager(bus) + ch = StubChannel() + failed_event = asyncio.Event() - task = asyncio.create_task(mgr._dispatch_outbound()) - await bus.publish_outbound( - OutboundMessage( - channel="stub", - chat_id="c1", - content="", - media=["/tmp/file.png"], - ) + async def failing_send_media(**kwargs): + failed_event.set() + return False + + ch.send_media = failing_send_media + mgr.register(ch) + + task = asyncio.create_task(mgr._dispatch_outbound()) + await bus.publish_outbound( + OutboundMessage( + channel="stub", + chat_id="c1", + content="", + media=["/tmp/file.png"], ) - await asyncio.sleep(0.1) - task.cancel() - try: - await task - except asyncio.CancelledError: - pass + ) + await asyncio.wait_for(failed_event.wait(), timeout=1.0) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass - health = mgr._health["stub"] - assert health.total_successes == 0 - assert health.total_failures == 1 - assert health.consecutive_failures == 1 - - _run(_test()) + health = mgr._health["stub"] + assert health.total_successes == 0 + assert health.total_failures == 1 + assert health.consecutive_failures == 1 class TestChannelManagerHealth: @@ -1175,83 +1126,71 @@ class TestChannelManagerDynamicOps: """[B-15] add_channel uses channel_type as key for start_times but register() uses channel.name — potential mismatch.""" - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - # We can't easily test add_channel without registry, - # but we can verify the key mismatch concern - ch = StubChannel() - ch.name = "custom_name" - mgr.register(ch) - assert "custom_name" in mgr._channels - # If add_channel used "other_type" but channel.name is "custom_name", - # start_times would be keyed differently + bus = MessageBus() + mgr = ChannelManager(bus) + # We can't easily test add_channel without registry, + # but we can verify the key mismatch concern + ch = StubChannel() + ch.name = "custom_name" + mgr.register(ch) + assert "custom_name" in mgr._channels + # If add_channel used "other_type" but channel.name is "custom_name", + # start_times would be keyed differently - _run(_test()) - - def test_remove_channel(self): + async def test_remove_channel(self): """[B-14] remove_channel removes from dict but doesn't cancel task.""" - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - ch = StubChannel() - mgr.register(ch) - assert "stub" in mgr._channels + bus = MessageBus() + mgr = ChannelManager(bus) + ch = StubChannel() + mgr.register(ch) + assert "stub" in mgr._channels - await mgr.remove_channel("stub") - assert "stub" not in mgr._channels + await mgr.remove_channel("stub") + assert "stub" not in mgr._channels - _run(_test()) - - def test_remove_nonexistent_channel(self): - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - await mgr.remove_channel("ghost") # should not raise - - _run(_test()) + async def test_remove_nonexistent_channel(self): + bus = MessageBus() + mgr = ChannelManager(bus) + await mgr.remove_channel("ghost") # should not raise class TestChannelManagerDrain: - def test_stop_all_drains_outbound(self): - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus, drain_timeout=1.0) - ch = StubChannel() - sent = [] - ch.send = AsyncMock(side_effect=lambda m: sent.append(m) or True) - mgr.register(ch) + async def test_stop_all_drains_outbound(self): + bus = MessageBus() + mgr = ChannelManager(bus, drain_timeout=1.0) + ch = StubChannel() + sent = [] + ch.send = AsyncMock(side_effect=lambda m: sent.append(m) or True) + mgr.register(ch) - # Pre-load an outbound message - await bus.publish_outbound( - OutboundMessage( - channel="stub", - chat_id="c1", - content="drain me", - ) + # Pre-load an outbound message + await bus.publish_outbound( + OutboundMessage( + channel="stub", + chat_id="c1", + content="drain me", ) + ) - await mgr.stop_all() - # The drain loop should have sent it - assert len(sent) == 1 - assert sent[0].content == "drain me" + await mgr.stop_all() + # The drain loop should have sent it + assert len(sent) == 1 + assert sent[0].content == "drain me" - _run(_test()) - - def test_stop_all_drains_media_and_counts_only_success(self, caplog): - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus, drain_timeout=1.0) - ch = StubChannel() - sent = [] - media_sent = [] - ch.send = AsyncMock(side_effect=lambda m: sent.append(m) or False) - ch.send_media = AsyncMock( - side_effect=lambda **kw: media_sent.append(kw) or True - ) - mgr.register(ch) + async def test_stop_all_drains_media_and_counts_only_success(self, caplog): + bus = MessageBus() + mgr = ChannelManager(bus, drain_timeout=1.0) + ch = StubChannel() + sent = [] + media_sent = [] + ch.send = AsyncMock(side_effect=lambda m: sent.append(m) or False) + ch.send_media = AsyncMock( + side_effect=lambda **kw: media_sent.append(kw) or True + ) + mgr.register(ch) + with caplog.at_level("INFO"): await bus.publish_outbound( OutboundMessage( channel="stub", @@ -1266,8 +1205,6 @@ class TestChannelManagerDrain: assert len(sent) == 1 assert len(media_sent) == 1 - with caplog.at_level("INFO"): - _run(_test()) assert "Outbound drain:" not in caplog.text @@ -1378,25 +1315,25 @@ class TestInboundConsumer: msg = BusInbound(channel="tg", sender_id="u1", chat_id="c1", content="hi") assert msg.session_key == "tg:c1" - def test_get_thread_id_creates_unique(self): + async def test_get_thread_id_creates_unique(self): consumer = self._make_consumer( graph_gateway=FakeGraphGateway( generated_thread_ids=["thread-a", "thread-b"] ) ) - tid1 = _run(consumer._get_thread_id("user_a")) - tid2 = _run(consumer._get_thread_id("user_b")) + tid1 = await consumer._get_thread_id("user_a") + tid2 = await consumer._get_thread_id("user_b") assert tid1 != tid2 - def test_get_thread_id_returns_same_for_same_sender(self): + async def test_get_thread_id_returns_same_for_same_sender(self): consumer = self._make_consumer( graph_gateway=FakeGraphGateway(generated_thread_ids=["thread-a"]) ) - tid1 = _run(consumer._get_thread_id("user_a")) - tid2 = _run(consumer._get_thread_id("user_a")) + tid1 = await consumer._get_thread_id("user_a") + tid2 = await consumer._get_thread_id("user_a") assert tid1 == tid2 - def test_shared_thread_id_bug(self): + async def test_shared_thread_id_bug(self): """[B-20] If thread_id is non-empty, senders get unique thread IDs with shared prefix.""" bus = MessageBus() mgr = ChannelManager(bus) @@ -1408,14 +1345,14 @@ class TestInboundConsumer: thread_id="shared_thread", # Non-empty! graph_gateway=FakeGraphGateway(), ) - tid1 = _run(consumer._get_thread_id("alice")) - tid2 = _run(consumer._get_thread_id("bob")) + tid1 = await consumer._get_thread_id("alice") + tid2 = await consumer._get_thread_id("bob") # Fixed: Each sender gets a unique thread_id using thread_id as prefix assert tid1 != tid2 assert tid1 == "shared_thread:alice" assert tid2 == "shared_thread:bob" - def test_session_eviction_is_lru(self): + async def test_session_eviction_is_lru(self): """Sessions use LRU eviction: recently accessed senders are kept.""" consumer = self._make_consumer() consumer._sessions.clear() @@ -1425,7 +1362,7 @@ class TestInboundConsumer: consumer._sessions[f"user_{i}"] = f"thread_{i}" # Access "user_0" via _get_thread_id (triggers LRU move_to_end) - _run(consumer._get_thread_id("user_0")) + await consumer._get_thread_id("user_0") # "user_0" should now be at the end (most recently used) oldest = next(iter(consumer._sessions)) @@ -1439,17 +1376,17 @@ class TestInboundConsumer: assert m["total_failures"] == 0 assert m["total_timeouts"] == 0 - def test_stop_graceful(self): - async def _test(): - consumer = self._make_consumer() - # Start and immediately stop - task = asyncio.create_task(consumer.run()) - await asyncio.sleep(0.1) - await consumer.stop() - await task - assert consumer._stopping is True - - _run(_test()) + async def test_stop_graceful(self): + consumer = self._make_consumer() + # Start and immediately stop + task = asyncio.create_task(consumer.run()) + await asyncio.sleep(0) + await consumer.stop() + await consumer.bus.publish_inbound( + BusInbound(channel="stub", sender_id="u1", chat_id="c1", content="wake") + ) + await task + assert consumer._stopping is True class TestInboundConsumerErrorHandling: @@ -1457,26 +1394,23 @@ class TestInboundConsumerErrorHandling: """[B-22] Exception messages are sent directly to users.""" # This test documents that internal error details are exposed - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - ch = StubChannel() - mgr.register(ch) + bus = MessageBus() + mgr = ChannelManager(bus) + ch = StubChannel() + mgr.register(ch) - _consumer = InboundConsumer( - bus=bus, - manager=mgr, - agent=MagicMock(), - thread_id="", - graph_gateway=FakeGraphGateway(), - ) + _consumer = InboundConsumer( + bus=bus, + manager=mgr, + agent=MagicMock(), + thread_id="", + graph_gateway=FakeGraphGateway(), + ) - # The error message format includes the raw exception - # This should be sanitized in production - error_msg = f"Error: {RuntimeError('secret internal path /etc/passwd')}" - assert "/etc/passwd" in error_msg # Documents the leak - - _run(_test()) + # The error message format includes the raw exception + # This should be sanitized in production + error_msg = f"Error: {RuntimeError('secret internal path /etc/passwd')}" + assert "/etc/passwd" in error_msg # Documents the leak # ═══════════════════════════════════════════════════════════════════ @@ -1492,47 +1426,36 @@ class TestMessageBus: ``TestChannelManagerDispatch`` for that coverage. """ - def test_publish_consume_inbound(self): - async def _test(): - bus = MessageBus() - msg = BusInbound( - channel="tg", sender_id="u1", chat_id="c1", content="hello" + async def test_publish_consume_inbound(self): + bus = MessageBus() + msg = BusInbound(channel="tg", sender_id="u1", chat_id="c1", content="hello") + await bus.publish_inbound(msg) + assert bus.inbound_size == 1 + received = await bus.consume_inbound() + assert received.content == "hello" + assert bus.inbound_size == 0 + + async def test_publish_consume_outbound(self): + bus = MessageBus() + msg = BusOutbound(channel="tg", chat_id="c1", content="reply") + await bus.publish_outbound(msg) + assert bus.outbound_size == 1 + received = await bus.consume_outbound() + assert received.content == "reply" + + async def test_queue_sizes(self): + bus = MessageBus() + assert bus.inbound_size == 0 + assert bus.outbound_size == 0 + await bus.publish_inbound( + BusInbound( + channel="x", + sender_id="u", + chat_id="c", + content="a", ) - await bus.publish_inbound(msg) - assert bus.inbound_size == 1 - received = await bus.consume_inbound() - assert received.content == "hello" - assert bus.inbound_size == 0 - - _run(_test()) - - def test_publish_consume_outbound(self): - async def _test(): - bus = MessageBus() - msg = BusOutbound(channel="tg", chat_id="c1", content="reply") - await bus.publish_outbound(msg) - assert bus.outbound_size == 1 - received = await bus.consume_outbound() - assert received.content == "reply" - - _run(_test()) - - def test_queue_sizes(self): - async def _test(): - bus = MessageBus() - assert bus.inbound_size == 0 - assert bus.outbound_size == 0 - await bus.publish_inbound( - BusInbound( - channel="x", - sender_id="u", - chat_id="c", - content="a", - ) - ) - assert bus.inbound_size == 1 - - _run(_test()) + ) + assert bus.inbound_size == 1 # ═══════════════════════════════════════════════════════════════════ @@ -1569,120 +1492,113 @@ class TestEvents: class TestIntegration: - def test_full_inbound_pipeline(self): + async def test_full_inbound_pipeline(self): """Raw message → build_inbound → queue_message → bus.""" - async def _test(): - bus = MessageBus() - ch = StubChannel() - ch.set_bus(bus) - ch.initial_debounce = 0.05 + bus = MessageBus() + ch = StubChannel() + ch.set_bus(bus) + ch.initial_debounce = 0.05 - raw = RawIncoming( - sender_id="user1", - chat_id="chat1", - text="integration test", - message_id="int_001", - ) - await ch._enqueue_raw(raw) + raw = RawIncoming( + sender_id="user1", + chat_id="chat1", + text="integration test", + message_id="int_001", + ) + await ch._enqueue_raw(raw) - # _enqueue_raw puts on internal queue, not bus - assert ch._queue.qsize() == 1 - inbound = await ch._queue.get() - assert inbound.content == "integration test" + # _enqueue_raw puts on internal queue, not bus + assert ch._queue.qsize() == 1 + inbound = await ch._queue.get() + assert inbound.content == "integration test" - # Now simulate the bus path via queue_message - await ch.queue_message(inbound) - await asyncio.sleep(0.2) - assert bus.inbound_size == 1 + # Now simulate the bus path via queue_message + await ch.queue_message(inbound) + await _flush_debounce(ch, "user1") + assert bus.inbound_size == 1 - _run(_test()) - - def test_outbound_dispatch_with_media(self): + async def test_outbound_dispatch_with_media(self): """Dispatch routes media alongside text content.""" - async def _test(): - bus = MessageBus() - mgr = ChannelManager(bus) - ch = StubChannel() - media_sent = [] - ch.send_media = AsyncMock( - side_effect=lambda **kw: media_sent.append(kw) or True, + bus = MessageBus() + mgr = ChannelManager(bus) + ch = StubChannel() + media_sent = [] + media_event = asyncio.Event() + + async def send_media(**kw): + media_sent.append(kw) + media_event.set() + return True + + ch.send_media = send_media + ch.send = AsyncMock(return_value=True) + mgr.register(ch) + + task = asyncio.create_task(mgr._dispatch_outbound()) + await bus.publish_outbound( + OutboundMessage( + channel="stub", + chat_id="c1", + content="see attached", + media=["/path/doc.pdf"], ) - ch.send = AsyncMock(return_value=True) - mgr.register(ch) + ) + await asyncio.wait_for(media_event.wait(), timeout=1.0) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass - task = asyncio.create_task(mgr._dispatch_outbound()) - await bus.publish_outbound( - OutboundMessage( - channel="stub", - chat_id="c1", - content="see attached", - media=["/path/doc.pdf"], - ) - ) - await asyncio.sleep(0.1) - task.cancel() - try: - await task - except asyncio.CancelledError: - pass + assert len(media_sent) == 1 - assert len(media_sent) == 1 - - _run(_test()) - - def test_debounce_lost_on_stop(self): + async def test_debounce_lost_on_stop(self): """Buffered messages should be flushed when stop() is called.""" - async def _test(): - bus = MessageBus() - ch = StubChannel() - ch.set_bus(bus) - ch.initial_debounce = 5.0 # Long debounce + bus = MessageBus() + ch = StubChannel() + ch.set_bus(bus) + ch.initial_debounce = 5.0 # Long debounce - msg = InboundMessage( - channel="stub", - sender_id="u1", - chat_id="c1", - content="will be lost", - message_id="m1", - metadata={"chat_id": "c1"}, - ) - await ch.queue_message(msg) - # Message is buffered but debounce hasn't fired yet + msg = InboundMessage( + channel="stub", + sender_id="u1", + chat_id="c1", + content="will be lost", + message_id="m1", + metadata={"chat_id": "c1"}, + ) + await ch.queue_message(msg) + # Message is buffered but debounce hasn't fired yet - assert len(ch._message_buffers) == 1 + assert len(ch._message_buffers) == 1 - # Stop the channel — debounce tasks are cancelled - ch._running = True - await ch.stop() + # Stop the channel — debounce tasks are cancelled + ch._running = True + await ch.stop() - assert bus.inbound_size == 1 - flushed = await bus.consume_inbound() - assert flushed.content == "will be lost" + assert bus.inbound_size == 1 + flushed = await bus.consume_inbound() + assert flushed.content == "will be lost" - _run(_test()) - - def test_send_locks_bounded_growth(self): + async def test_send_locks_bounded_growth(self): """_send_locks stays bounded via LRU eviction of unlocked entries.""" - async def _test(): - ch = StubChannel() - ch._send_locks_max = 10 # Small limit for testing - for i in range(20): - msg = OutboundMessage( - channel="stub", - chat_id=f"chat_{i}", - content="hi", - metadata={"chat_id": f"chat_{i}"}, - ) - await ch.send(msg) + ch = StubChannel() + ch._send_locks_max = 10 # Small limit for testing + for i in range(20): + msg = OutboundMessage( + channel="stub", + chat_id=f"chat_{i}", + content="hi", + metadata={"chat_id": f"chat_{i}"}, + ) + await ch.send(msg) - # Should be bounded at max + 1 (the newly inserted entry) - assert len(ch._send_locks) <= ch._send_locks_max + 1 - - _run(_test()) + # Should be bounded at max + 1 (the newly inserted entry) + assert len(ch._send_locks) <= ch._send_locks_max + 1 # ═══════════════════════════════════════════════════════════════════ @@ -1707,16 +1623,13 @@ class TestEdgeCases: for _ in range(100): assert dc.is_duplicate("x") is True - def test_channel_send_empty_content(self): - async def _test(): - ch = StubChannel() - msg = OutboundMessage(channel="stub", chat_id="c1", content="") - ok = await ch.send(msg) - # Empty content goes through chunk_text which returns [] - assert ok is True - assert len(ch._sent_chunks) == 0 - - _run(_test()) + async def test_channel_send_empty_content(self): + ch = StubChannel() + msg = OutboundMessage(channel="stub", chat_id="c1", content="") + ok = await ch.send(msg) + # Empty content goes through chunk_text which returns [] + assert ok is True + assert len(ch._sent_chunks) == 0 def test_raw_incoming_defaults(self): raw = RawIncoming(sender_id="u1", chat_id="c1") diff --git a/tests/test_channel_debug.py b/tests/test_channel_debug.py index 3abd1f1..f81af68 100644 --- a/tests/test_channel_debug.py +++ b/tests/test_channel_debug.py @@ -11,8 +11,6 @@ from EvoScientist.channels.debug import ( emit_debug_event_if, ) -from .conftest import run_async - def test_debug_trace_enabled_from_bool(): assert debug_trace_enabled(True) is True @@ -75,10 +73,10 @@ def _make_channel_context(*, debug_trace=True, name="test_channel"): return {"channel": channel} -def test_middleware_dedup_emits_structured_event(caplog): +async def test_middleware_dedup_emits_structured_event(caplog): from EvoScientist.channels.middleware import DedupMiddleware - async def _run(): + with caplog.at_level(logging.DEBUG): mw = DedupMiddleware() ctx = _make_channel_context() raw = _make_raw(message_id="dup1") @@ -91,49 +89,40 @@ def test_middleware_dedup_emits_structured_event(caplog): caplog.clear() result = await mw.process_inbound(raw, ctx) assert result is None - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "middleware_dedup_drop" in caplog.text assert "message_id=dup1" in caplog.text -def test_middleware_allowlist_emits_structured_event(caplog): +async def test_middleware_allowlist_emits_structured_event(caplog): from EvoScientist.channels.middleware import AllowListMiddleware - async def _run(): + with caplog.at_level(logging.DEBUG): mw = AllowListMiddleware(allowed_senders={"allowed_user"}) ctx = _make_channel_context() raw = _make_raw(sender_id="blocked_user") result = await mw.process_inbound(raw, ctx) assert result is None - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "middleware_allowlist_drop" in caplog.text assert "reason=sender_not_allowed" in caplog.text -def test_middleware_mention_gating_emits_structured_event(caplog): +async def test_middleware_mention_gating_emits_structured_event(caplog): from EvoScientist.channels.middleware import MentionGatingMiddleware - async def _run(): + with caplog.at_level(logging.DEBUG): mw = MentionGatingMiddleware(require_mention="group") ctx = _make_channel_context() raw = _make_raw(is_group=True, was_mentioned=False) result = await mw.process_inbound(raw, ctx) assert result is None - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "middleware_mention_drop" in caplog.text assert "policy=group" in caplog.text -def test_typing_manager_emits_trace_events(caplog): +async def test_typing_manager_emits_trace_events(caplog): from EvoScientist.channels.middleware import TypingManager - async def _run(): + with caplog.at_level(logging.DEBUG): send_action = AsyncMock(side_effect=RuntimeError("typing api down")) mgr = TypingManager( send_action, @@ -144,17 +133,14 @@ def test_typing_manager_emits_trace_events(caplog): await mgr.start("chat1") await asyncio.sleep(0.01) await mgr.stop("chat1") - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "typing_error" in caplog.text assert "chat_id=chat1" in caplog.text -def test_ack_reaction_emits_error_traces(caplog): +async def test_ack_reaction_emits_error_traces(caplog): from EvoScientist.channels.middleware import AckReactionMiddleware - async def _run(): + with caplog.at_level(logging.DEBUG): send_fn = AsyncMock() remove_fn = AsyncMock(side_effect=RuntimeError("remove failed")) ack = AckReactionMiddleware( @@ -178,16 +164,13 @@ def test_ack_reaction_emits_error_traces(caplog): send_fn.reset_mock() send_fn.side_effect = RuntimeError("api down") await ack.send_ack("chat2", "msg2") - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "ack_send_error" in caplog.text assert "ack_remove_error" in caplog.text assert "api down" in caplog.text assert "remove failed" in caplog.text -def test_inbound_raw_event_emitted(caplog): +async def test_inbound_raw_event_emitted(caplog): """Integration-style: _enqueue_raw emits inbound_raw at the top.""" from EvoScientist.channels.base import Channel, RawIncoming @@ -218,19 +201,16 @@ def test_inbound_raw_event_emitted(caplog): config.ack_scope = "off" config.dedup_ttl = 3600 - async def _run(): + with caplog.at_level(logging.DEBUG): with patch.object(Channel, "__abstractmethods__", set()): ch = _TestChannel(config) raw = RawIncoming(sender_id="u1", chat_id="c1", text="hi", message_id="m1") await ch._enqueue_raw(raw) - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "inbound_raw" in caplog.text assert "sender_id=u1" in caplog.text -def test_format_fallback_emits_event(caplog): +async def test_format_fallback_emits_event(caplog): """_send_with_format_fallback emits outbound_format_fallback on fallback.""" from EvoScientist.channels.base import Channel @@ -268,13 +248,10 @@ def test_format_fallback_emits_event(caplog): if call_count == 1: raise ValueError("parse error in formatted text") - async def _run(): + with caplog.at_level(logging.DEBUG): with patch.object(Channel, "__abstractmethods__", set()): ch = _TestChannel(config) await ch._send_with_format_fallback(_failing_send, "hi", "hi") - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "outbound_format_fallback" in caplog.text assert call_count == 2 @@ -304,7 +281,7 @@ def test_trace_mixin_trace_event(caplog): assert "key=val" in caplog.text -def test_standalone_dispatcher_treats_false_send_as_error(caplog): +async def test_standalone_dispatcher_treats_false_send_as_error(caplog): from EvoScientist.channels.bus import MessageBus from EvoScientist.channels.bus.events import OutboundMessage from EvoScientist.channels.standalone import standalone_outbound_dispatcher @@ -315,7 +292,7 @@ def test_standalone_dispatcher_treats_false_send_as_error(caplog): channel.send = AsyncMock(return_value=False) bus = MessageBus() - async def _run(): + with caplog.at_level(logging.DEBUG): task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel)) await bus.publish_outbound( OutboundMessage(channel="test", chat_id="c1", content="hi") @@ -326,14 +303,11 @@ def test_standalone_dispatcher_treats_false_send_as_error(caplog): await task except asyncio.CancelledError: pass - - with caplog.at_level(logging.DEBUG): - run_async(_run()) assert "standalone_dispatch_error" in caplog.text assert "send() returned False" in caplog.text -def test_standalone_dispatcher_sends_media(): +async def test_standalone_dispatcher_sends_media(): from EvoScientist.channels.bus import MessageBus from EvoScientist.channels.bus.events import OutboundMessage from EvoScientist.channels.standalone import standalone_outbound_dispatcher @@ -345,21 +319,16 @@ def test_standalone_dispatcher_sends_media(): channel.send_media = AsyncMock(return_value=True) bus = MessageBus() - async def _run(): - task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel)) - await bus.publish_outbound( - OutboundMessage( - channel="test", chat_id="c1", content="", media=["/tmp/a.png"] - ) - ) - await asyncio.sleep(0.05) - task.cancel() - try: - await task - except asyncio.CancelledError: - pass - - run_async(_run()) + task = asyncio.create_task(standalone_outbound_dispatcher(bus, channel)) + await bus.publish_outbound( + OutboundMessage(channel="test", chat_id="c1", content="", media=["/tmp/a.png"]) + ) + await asyncio.sleep(0.05) + task.cancel() + try: + await task + except asyncio.CancelledError: + pass channel.send_media.assert_awaited_once_with( recipient="c1", file_path="/tmp/a.png", diff --git a/tests/test_cli_channel_slash.py b/tests/test_cli_channel_slash.py index aa41728..c3d8916 100644 --- a/tests/test_cli_channel_slash.py +++ b/tests/test_cli_channel_slash.py @@ -14,7 +14,6 @@ from EvoScientist.cli.channel import ( from EvoScientist.cli.channel import ( dispatch_channel_slash_command as _dispatch_channel_slash_command, ) -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -43,25 +42,23 @@ def _make_msg( ) -def test_non_slash_returns_false(): +async def test_non_slash_returns_false(): """Plain text messages must fall through to the agent.""" msg = _make_msg(content="hello agent") append = MagicMock() - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, ) assert handled is False append.assert_not_called() -def test_unresolved_slash_returns_false(): +async def test_unresolved_slash_returns_false(): """Unknown slash commands must fall through (matches TUI behavior).""" msg = _make_msg(content="/unknown-cmd") append = MagicMock() @@ -69,20 +66,18 @@ def test_unresolved_slash_returns_false(): "EvoScientist.commands.manager.manager.resolve", return_value=None, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, ) assert handled is False -def test_successful_slash_execution_sets_response_and_breadcrumb(): +async def test_successful_slash_execution_sets_response_and_breadcrumb(): """Known slash command: cmd_manager.execute ran, helper returns True, sends a confirmation to the channel user, and appends a local log line.""" msg = _make_msg() @@ -100,15 +95,13 @@ def test_successful_slash_execution_sets_response_and_breadcrumb(): ) as mock_execute, patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent="fake-agent", - thread_id="t1", - workspace_dir="/tmp", - checkpointer=None, - append_system=append, - ) + handled = await dispatch_channel_slash_command( + msg, + agent="fake-agent", + thread_id="t1", + workspace_dir="/tmp", + checkpointer=None, + append_system=append, ) assert handled is True mock_execute.assert_awaited_once() @@ -119,7 +112,7 @@ def test_successful_slash_execution_sets_response_and_breadcrumb(): assert any("Executed command from" in t for t in breadcrumbs) -def test_slash_dispatch_passes_graph_gateway_to_command_context(): +async def test_slash_dispatch_passes_graph_gateway_to_command_context(): msg = _make_msg() fake_cmd = MagicMock() fake_cmd.needs_agent.return_value = False @@ -142,23 +135,21 @@ def test_slash_dispatch_passes_graph_gateway_to_command_context(): ), patch("EvoScientist.cli.channel._set_channel_response"), ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent="fake-agent", - thread_id="t1", - workspace_dir="/tmp", - checkpointer=None, - append_system=append, - graph_gateway=graph_gateway, - ) + handled = await dispatch_channel_slash_command( + msg, + agent="fake-agent", + thread_id="t1", + workspace_dir="/tmp", + checkpointer=None, + append_system=append, + graph_gateway=graph_gateway, ) assert handled is True assert captured["graph_gateway"] is graph_gateway -def test_needs_agent_awaits_loader_and_passes_result(): +async def test_needs_agent_awaits_loader_and_passes_result(): """Commands with needs_agent=True must await the loader and the resulting agent must flow through the CommandContext.""" msg = _make_msg() @@ -182,16 +173,14 @@ def test_needs_agent_awaits_loader_and_passes_result(): ) as mock_execute, patch("EvoScientist.cli.channel._set_channel_response"), ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - await_agent_ready=_await_ready, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, + await_agent_ready=_await_ready, ) assert handled is True await_called.assert_called_once() @@ -201,7 +190,7 @@ def test_needs_agent_awaits_loader_and_passes_result(): assert ctx_arg.agent == "ready-agent" -def test_await_agent_ready_failure_sets_error_response(): +async def test_await_agent_ready_failure_sets_error_response(): msg = _make_msg() fake_cmd = MagicMock() fake_cmd.needs_agent.return_value = True @@ -217,16 +206,14 @@ def test_await_agent_ready_failure_sets_error_response(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - await_agent_ready=_await_ready, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, + await_agent_ready=_await_ready, ) assert handled is True mock_set_resp.assert_called_once() @@ -235,7 +222,7 @@ def test_await_agent_ready_failure_sets_error_response(): assert "agent blew up" in resp_text -def test_cmd_manager_raises_returns_true_with_error(): +async def test_cmd_manager_raises_returns_true_with_error(): """If cmd_manager.execute raises past its own try/except, the helper must absorb it, return True, and report via _set_channel_response.""" msg = _make_msg() @@ -253,15 +240,13 @@ def test_cmd_manager_raises_returns_true_with_error(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, ) assert handled is True mock_set_resp.assert_called_once() @@ -270,7 +255,7 @@ def test_cmd_manager_raises_returns_true_with_error(): assert "boom" in resp_text -def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd(): +async def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd(): """After a successful slash execute, the on_cmd_completed hook must be awaited with (ctx, original_agent, cmd) so Rich CLI can adopt an ``/model`` agent swap and refresh status for state-mutating commands.""" @@ -302,16 +287,14 @@ def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd(): ), patch("EvoScientist.cli.channel._set_channel_response"), ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent="original-agent", - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - on_cmd_completed=_on_completed, - ) + handled = await dispatch_channel_slash_command( + msg, + agent="original-agent", + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, + on_cmd_completed=_on_completed, ) assert handled is True assert captured["ctx_agent"] == "swapped-agent" @@ -319,7 +302,7 @@ def test_on_cmd_completed_awaited_with_ctx_original_agent_and_cmd(): assert captured["cmd_name"] == "/model" -def test_on_cmd_completed_receives_cmd_for_new_and_compact(): +async def test_on_cmd_completed_receives_cmd_for_new_and_compact(): """``/new`` / ``/compact`` invoked via channel must flow the cmd into the hook so the callback can still refresh status when the agent didn't swap — mirrors REPL ``interactive.py:1027-1030``.""" @@ -343,21 +326,19 @@ def test_on_cmd_completed_receives_cmd_for_new_and_compact(): ), patch("EvoScientist.cli.channel._set_channel_response"), ): - _run( - dispatch_channel_slash_command( - _make_msg(content=cmd_name), - agent="same-agent", - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=_on_completed, - ) + await dispatch_channel_slash_command( + _make_msg(content=cmd_name), + agent="same-agent", + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=_on_completed, ) assert captured["cmd_name"] == cmd_name, cmd_name -def test_on_cmd_completed_skipped_on_fall_through_and_error(): +async def test_on_cmd_completed_skipped_on_fall_through_and_error(): """The hook must NOT fire for unresolved slash, non-slash text, or when cmd_manager.execute raised.""" fake_cmd = MagicMock() @@ -369,16 +350,14 @@ def test_on_cmd_completed_skipped_on_fall_through_and_error(): # Non-slash with patch("EvoScientist.cli.channel._set_channel_response"): - _run( - dispatch_channel_slash_command( - _make_msg(content="hi"), - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=_noop, - ) + await dispatch_channel_slash_command( + _make_msg(content="hi"), + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=_noop, ) # Unresolved slash with ( @@ -388,16 +367,14 @@ def test_on_cmd_completed_skipped_on_fall_through_and_error(): ), patch("EvoScientist.cli.channel._set_channel_response"), ): - _run( - dispatch_channel_slash_command( - _make_msg(content="/nope"), - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=_noop, - ) + await dispatch_channel_slash_command( + _make_msg(content="/nope"), + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=_noop, ) # Execute raises with ( @@ -411,22 +388,20 @@ def test_on_cmd_completed_skipped_on_fall_through_and_error(): ), patch("EvoScientist.cli.channel._set_channel_response"), ): - _run( - dispatch_channel_slash_command( - _make_msg(), - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=_noop, - ) + await dispatch_channel_slash_command( + _make_msg(), + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=_noop, ) completed.assert_not_called() -def test_command_error_skips_completion_hook_and_reports_error(): +async def test_command_error_skips_completion_hook_and_reports_error(): """A command caught as failed by CommandManager must not look successful.""" msg = _make_msg(content="/resume abc") fake_cmd = MagicMock() @@ -450,16 +425,14 @@ def test_command_error_skips_completion_hook_and_reports_error(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="old-thread", - workspace_dir="/old-workspace", - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=completed, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="old-thread", + workspace_dir="/old-workspace", + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=completed, ) assert handled is True @@ -467,7 +440,7 @@ def test_command_error_skips_completion_hook_and_reports_error(): mock_set_resp.assert_called_once_with("msg-1", "Command error: workspace conflict") -def test_empty_command_error_still_reports_error(): +async def test_empty_command_error_still_reports_error(): """An empty string error is still a command failure sentinel.""" msg = _make_msg(content="/resume abc") fake_cmd = MagicMock() @@ -489,16 +462,14 @@ def test_empty_command_error_still_reports_error(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="old-thread", - workspace_dir="/old-workspace", - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=completed, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="old-thread", + workspace_dir="/old-workspace", + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=completed, ) assert handled is True @@ -506,7 +477,7 @@ def test_empty_command_error_still_reports_error(): mock_set_resp.assert_called_once_with("msg-1", "Command error: (no details)") -def test_on_cmd_completed_exception_is_absorbed(): +async def test_on_cmd_completed_exception_is_absorbed(): """A raising hook must NOT prevent the channel response from being set.""" msg = _make_msg() fake_cmd = MagicMock() @@ -526,23 +497,21 @@ def test_on_cmd_completed_exception_is_absorbed(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent="orig", - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=MagicMock(), - on_cmd_completed=_boom, - ) + handled = await dispatch_channel_slash_command( + msg, + agent="orig", + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=MagicMock(), + on_cmd_completed=_boom, ) assert handled is True mock_set_resp.assert_called_once() assert "Command executed" in mock_set_resp.call_args[0][1] -def test_top_level_exception_is_absorbed(): +async def test_top_level_exception_is_absorbed(): """Last-ditch safety net: if anything inside the dispatch pipeline raises unexpectedly (lazy import failure, ChannelCommandUI ctor, terminal I/O from append_system, ...), the helper must NOT @@ -557,15 +526,13 @@ def test_top_level_exception_is_absorbed(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=MagicMock(), - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=MagicMock(), ) assert handled is True mock_set_resp.assert_called_once() @@ -574,7 +541,7 @@ def test_top_level_exception_is_absorbed(): assert "exploded during resolve" in resp_text -def test_cmd_execute_returning_false_falls_through(): +async def test_cmd_execute_returning_false_falls_through(): """When cmd_manager.execute returns False (empty/unparseable input), the helper must return False so the caller falls through to the agent.""" msg = _make_msg(content="/") @@ -592,15 +559,13 @@ def test_cmd_execute_returning_false_falls_through(): ), patch("EvoScientist.cli.channel._set_channel_response") as mock_set_resp, ): - handled = _run( - dispatch_channel_slash_command( - msg, - agent=None, - thread_id="t1", - workspace_dir=None, - checkpointer=None, - append_system=append, - ) + handled = await dispatch_channel_slash_command( + msg, + agent=None, + thread_id="t1", + workspace_dir=None, + checkpointer=None, + append_system=append, ) assert handled is False mock_set_resp.assert_not_called() diff --git a/tests/test_cli_tui_dispatch.py b/tests/test_cli_tui_dispatch.py index 631a1d5..508ea9a 100644 --- a/tests/test_cli_tui_dispatch.py +++ b/tests/test_cli_tui_dispatch.py @@ -1,6 +1,5 @@ """Tests for CLI interactive UI backend dispatch.""" -import asyncio from types import SimpleNamespace import pytest @@ -101,7 +100,7 @@ def test_background_agent_server_starts_even_when_async_subagents_disabled( assert calls == [(config, "/tmp/workspace")] -def test_resume_workspace_sync_runs_even_when_async_subagents_disabled( +async def test_resume_workspace_sync_runs_even_when_async_subagents_disabled( monkeypatch, ): import EvoScientist.cli.commands as cmds @@ -117,11 +116,9 @@ def test_resume_workspace_sync_runs_even_when_async_subagents_disabled( ) config = SimpleNamespace(enable_async_subagents=False) - asyncio.run( - cmds._sync_background_agent_server_workspace( - config, - workspace_dir="/tmp/resumed-workspace", - ) + await cmds._sync_background_agent_server_workspace( + config, + workspace_dir="/tmp/resumed-workspace", ) assert calls == [(config, "/tmp/resumed-workspace")] diff --git a/tests/test_compact_command.py b/tests/test_compact_command.py index 74444be..7d8b45a 100644 --- a/tests/test_compact_command.py +++ b/tests/test_compact_command.py @@ -4,13 +4,12 @@ from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch from EvoScientist.gateway import GraphTarget -from tests.conftest import run_async as _run from tests.fakes import FakeCommandUI, FakeGraphGateway _TARGET = GraphTarget() -def _compact( +async def _compact( graph_gateway: FakeGraphGateway, *, thread_id: str = "tid-1", @@ -18,30 +17,28 @@ def _compact( ): from EvoScientist.cli.commands import compact_conversation - return _run( - compact_conversation( - graph_gateway=graph_gateway, - thread_id=thread_id, - target=_TARGET, - input_tokens_hint=input_tokens_hint, - ) + return await compact_conversation( + graph_gateway=graph_gateway, + thread_id=thread_id, + target=_TARGET, + input_tokens_hint=input_tokens_hint, ) class TestCompactGuards: """Guard conditions that return early without touching the middleware.""" - def test_empty_messages(self): + async def test_empty_messages(self): graph_gateway = FakeGraphGateway(state_values={"messages": []}) - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "noop" assert "no messages" in result.message - def test_state_read_failure(self): + async def test_state_read_failure(self): graph_gateway = FakeGraphGateway(state_error=RuntimeError("DB gone")) - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "error" assert "Failed to read state" in result.message @@ -49,7 +46,7 @@ class TestCompactGuards: class TestCompactCutoffZero: """When cutoff == 0, conversation is within retention budget.""" - def test_nothing_to_compact_short_conversation(self): + async def test_nothing_to_compact_short_conversation(self): msgs = [MagicMock() for _ in range(3)] graph_gateway = FakeGraphGateway(state_values={"messages": msgs}) @@ -79,7 +76,7 @@ class TestCompactCutoffZero: return_value=500, ), ): - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "noop" assert "within the retention budget" in result.message @@ -89,7 +86,7 @@ class TestCompactCutoffZero: class TestCompactNegligibleSavings: """When cutoff > 0 but savings are too small to be worth it.""" - def test_skip_when_few_messages_and_low_tokens(self): + async def test_skip_when_few_messages_and_low_tokens(self): msgs = [MagicMock() for _ in range(15)] graph_gateway = FakeGraphGateway( state_values={"messages": msgs, "_summarization_event": None} @@ -126,14 +123,14 @@ class TestCompactNegligibleSavings: side_effect=lambda x: next(token_values), ), ): - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "noop" assert "not worth" in result.message # No LLM call should have been made mock_middleware_inst._acreate_summary.assert_not_called() - def test_still_compacts_when_few_messages_but_high_tokens(self): + async def test_still_compacts_when_few_messages_but_high_tokens(self): """2 messages but they account for >2% of tokens — should compact.""" from langchain_core.messages import HumanMessage @@ -178,7 +175,7 @@ class TestCompactNegligibleSavings: side_effect=lambda x: next(token_values), ), ): - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "ok" assert len(graph_gateway.updated_states) == 1 @@ -187,7 +184,7 @@ class TestCompactNegligibleSavings: class TestCompactSuccess: """Normal compaction flow.""" - def test_manual_threshold_blocks_low_context_compaction(self): + async def test_manual_threshold_blocks_low_context_compaction(self): msgs = [MagicMock() for _ in range(20)] graph_gateway = FakeGraphGateway( state_values={"messages": msgs, "_summarization_event": None} @@ -217,7 +214,7 @@ class TestCompactSuccess: return_value=30_000, ), ): - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "noop" assert "40%" in result.message @@ -225,7 +222,7 @@ class TestCompactSuccess: mock_middleware_inst._determine_cutoff_index.assert_not_called() mock_middleware_inst._acreate_summary.assert_not_called() - def test_successful_compaction(self): + async def test_successful_compaction(self): from langchain_core.messages import HumanMessage msgs = [MagicMock() for _ in range(20)] @@ -273,7 +270,7 @@ class TestCompactSuccess: side_effect=lambda x: next(token_values), ), ): - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "ok" assert result.messages_compacted == 15 @@ -291,7 +288,7 @@ class TestCompactSuccess: assert "_summarization_event" in event_data assert event_data["_summarization_event"]["cutoff_index"] == 15 - def test_offload_failure_non_fatal(self): + async def test_offload_failure_non_fatal(self): """Offload failure should not prevent compaction.""" from langchain_core.messages import HumanMessage @@ -335,7 +332,7 @@ class TestCompactSuccess: return_value=1000, ), ): - result = _compact(graph_gateway) + result = await _compact(graph_gateway) assert result.status == "ok" assert len(graph_gateway.updated_states) == 1 @@ -377,7 +374,7 @@ class TestRenderCompactResult: class TestCompactCommandUI: """TUI-specific compact progress indicator behavior.""" - def test_command_uses_tui_indicator_when_available(self): + async def test_command_uses_tui_indicator_when_available(self): from EvoScientist.cli.commands import CompactResult from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.session import CompactCommand @@ -413,7 +410,7 @@ class TestCompactCommandUI: return_value="summary-panel", ), ): - _run(CompactCommand().execute(ctx, [])) + await CompactCommand().execute(ctx, []) assert ui.started == 1 assert ui.stopped == 1 diff --git a/tests/test_config.py b/tests/test_config.py index e68b1e5..8799da5 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -242,6 +242,23 @@ class TestLoadSaveReset: assert data["provider"] == "openai" assert data["model"] == "gpt-4o" + def test_save_restricts_config_permissions(self, temp_config_dir, clean_env): + """Config file permissions should not depend on the process umask.""" + original_umask = os.umask(0) + try: + save_config(EvoScientistConfig(anthropic_api_key="test-key")) + finally: + os.umask(original_umask) + + config_path = get_config_path() + if os.name == "nt": + assert config_path.exists() + # Windows reports pseudo-permission bits, so we don't test them here. + return + + assert config_path.parent.stat().st_mode & 0o777 == 0o700 + assert config_path.stat().st_mode & 0o777 == 0o600 + def test_load_reads_saved_config(self, temp_config_dir, clean_env): """Test that load reads previously saved config.""" original = EvoScientistConfig( diff --git a/tests/test_configurable_model_middleware.py b/tests/test_configurable_model_middleware.py index 73930c4..6314186 100644 --- a/tests/test_configurable_model_middleware.py +++ b/tests/test_configurable_model_middleware.py @@ -15,7 +15,6 @@ from EvoScientist.middleware.configurable_model import ( ConfigurableModelMiddleware, _read_model_override, ) -from tests.conftest import run_async as _run @contextmanager @@ -134,7 +133,7 @@ class TestPassThrough: handler.assert_called_once_with(req) req.override.assert_not_called() - def test_async_no_override_passes_request_unchanged(self): + async def test_async_no_override_passes_request_unchanged(self): mw = ConfigurableModelMiddleware() req = _make_request() @@ -143,7 +142,7 @@ class TestPassThrough: return "ok" with _patched_config({}): - result = _run(mw.awrap_model_call(req, handler)) + result = await mw.awrap_model_call(req, handler) assert result == "ok" req.override.assert_not_called() @@ -185,7 +184,7 @@ class TestModelOverride: assert called_with is not req assert called_with.model is new_model - def test_async_override_path_parity(self): + async def test_async_override_path_parity(self): mw = ConfigurableModelMiddleware() req = _make_request() new_model = MagicMock() @@ -202,7 +201,7 @@ class TestModelOverride: "EvoScientist.llm.get_chat_model", return_value=new_model ) as mock_get, ): - result = _run(mw.awrap_model_call(req, handler)) + result = await mw.awrap_model_call(req, handler) assert result == "ok" mock_get.assert_called_once_with(model="claude-opus-4-8", provider="anthropic") @@ -316,7 +315,7 @@ class TestResolveFailure: handler.assert_called_once_with(req) req.override.assert_not_called() - def test_async_falls_back_when_resolve_raises(self): + async def test_async_falls_back_when_resolve_raises(self): mw = ConfigurableModelMiddleware() req = _make_request() @@ -333,7 +332,7 @@ class TestResolveFailure: side_effect=ValueError("unknown model"), ), ): - result = _run(mw.awrap_model_call(req, handler)) + result = await mw.awrap_model_call(req, handler) assert result == "ok" assert called == [req] diff --git a/tests/test_context_overflow_middleware.py b/tests/test_context_overflow_middleware.py index 51976da..09574bc 100644 --- a/tests/test_context_overflow_middleware.py +++ b/tests/test_context_overflow_middleware.py @@ -66,7 +66,6 @@ def test_wrap_model_call_raises_context_overflow(): assert handler.call_count == 1 -@pytest.mark.anyio async def test_awrap_model_call_raises_context_overflow(): # Setup mocks msgs = [HumanMessage(content=f"msg {i}") for i in range(10)] @@ -91,7 +90,6 @@ async def test_awrap_model_call_raises_context_overflow(): assert handler.call_count == 1 -@pytest.mark.anyio async def test_awrap_model_call_passes_through_other_errors(): request = ModelRequest( messages=[], diff --git a/tests/test_current_command.py b/tests/test_current_command.py index 352c073..4296faf 100644 --- a/tests/test_current_command.py +++ b/tests/test_current_command.py @@ -2,11 +2,9 @@ from unittest.mock import MagicMock -from tests.conftest import run_async as _run - class TestCurrentCommand: - def test_prints_thread_workspace_and_memory(self): + async def test_prints_thread_workspace_and_memory(self): from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.general import CurrentCommand @@ -17,14 +15,14 @@ class TestCurrentCommand: ui=ui, workspace_dir="/tmp/ws", ) - _run(CurrentCommand().execute(ctx, [])) + await CurrentCommand().execute(ctx, []) # Three append_system calls: Thread, Workspace, Memory dir. calls = [c.args[0] for c in ui.append_system.call_args_list] assert any("Thread: abc123" in s for s in calls) assert any("Workspace:" in s for s in calls) assert any("Memory dir:" in s for s in calls) - def test_skips_workspace_when_none(self): + async def test_skips_workspace_when_none(self): from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.general import CurrentCommand @@ -35,7 +33,7 @@ class TestCurrentCommand: ui=ui, workspace_dir=None, ) - _run(CurrentCommand().execute(ctx, [])) + await CurrentCommand().execute(ctx, []) calls = [c.args[0] for c in ui.append_system.call_args_list] assert any("Thread: abc123" in s for s in calls) assert not any("Workspace:" in s for s in calls) diff --git a/tests/test_delete_command.py b/tests/test_delete_command.py index 082a8ea..bf41611 100644 --- a/tests/test_delete_command.py +++ b/tests/test_delete_command.py @@ -2,7 +2,6 @@ from unittest.mock import AsyncMock, MagicMock -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -21,17 +20,17 @@ def _ctx(thread_id="current", thread_store=None): class TestDeleteCommand: - def test_refuses_to_delete_current(self): + async def test_refuses_to_delete_current(self): from EvoScientist.commands.implementation.session import DeleteCommand thread_store = FakeThreadStore(resolved_thread_id="current", deleted=True) ctx, ui = _ctx(thread_id="current", thread_store=thread_store) - _run(DeleteCommand().execute(ctx, ["current"])) + await DeleteCommand().execute(ctx, ["current"]) assert ("delete_thread", "current") not in thread_store.calls msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Cannot delete the current session" in m for m in msgs) - def test_happy_path_success(self): + async def test_happy_path_success(self): from EvoScientist.commands.implementation.session import DeleteCommand ctx, ui = _ctx( @@ -41,45 +40,45 @@ class TestDeleteCommand: deleted=True, ), ) - _run(DeleteCommand().execute(ctx, ["other-thread"])) + await DeleteCommand().execute(ctx, ["other-thread"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Deleted session other-thread" in m for m in msgs) - def test_not_found(self): + async def test_not_found(self): from EvoScientist.commands.implementation.session import DeleteCommand ctx, ui = _ctx() - _run(DeleteCommand().execute(ctx, ["missing"])) + await DeleteCommand().execute(ctx, ["missing"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("not found" in m for m in msgs) - def test_ambiguous_prefix(self): + async def test_ambiguous_prefix(self): from EvoScientist.commands.implementation.session import DeleteCommand ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"])) - _run(DeleteCommand().execute(ctx, ["abc"])) + await DeleteCommand().execute(ctx, ["abc"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Ambiguous" in m for m in msgs) - def test_prefix_resolves_to_unique_match(self): + async def test_prefix_resolves_to_unique_match(self): from EvoScientist.commands.implementation.session import DeleteCommand ctx, ui = _ctx( thread_store=FakeThreadStore(resolved_thread_id="abc-one", deleted=True) ) - _run(DeleteCommand().execute(ctx, ["abc"])) + await DeleteCommand().execute(ctx, ["abc"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Deleted session abc-one" in m for m in msgs) - def test_no_arg_empty_sessions_prints_notice(self): + async def test_no_arg_empty_sessions_prints_notice(self): from EvoScientist.commands.implementation.session import DeleteCommand ctx, ui = _ctx() - _run(DeleteCommand().execute(ctx, [])) + await DeleteCommand().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("No sessions to delete" in m for m in msgs) - def test_no_arg_calls_picker_returns_none(self): + async def test_no_arg_calls_picker_returns_none(self): """When no arg and picker returns None, nothing is deleted.""" from EvoScientist.commands.implementation.session import DeleteCommand @@ -96,5 +95,5 @@ class TestDeleteCommand: ] store = FakeThreadStore(threads=threads) ctx.graph_gateway = FakeGraphGateway(thread_store=store) - _run(DeleteCommand().execute(ctx, [])) + await DeleteCommand().execute(ctx, []) ui.wait_for_thread_pick.assert_awaited_once() diff --git a/tests/test_dingtalk_channel.py b/tests/test_dingtalk_channel.py index f2fd924..be865a0 100644 --- a/tests/test_dingtalk_channel.py +++ b/tests/test_dingtalk_channel.py @@ -7,7 +7,6 @@ import pytest from EvoScientist.channels.base import ChannelError, OutboundMessage from EvoScientist.channels.dingtalk.channel import DingTalkChannel, DingTalkConfig -from tests.conftest import run_async as _run class TestDingTalkConfig: @@ -41,30 +40,30 @@ class TestDingTalkChannel: assert channel._running is False assert channel.name == "dingtalk" - def test_start_raises_without_credentials(self): + async def test_start_raises_without_credentials(self): config = DingTalkConfig(client_id="", client_secret="") channel = DingTalkChannel(config) with pytest.raises(ChannelError, match="client_id and client_secret"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_client_id(self): + async def test_start_raises_without_client_id(self): config = DingTalkConfig(client_id="", client_secret="secret") channel = DingTalkChannel(config) with pytest.raises(ChannelError, match="client_id and client_secret"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_client_secret(self): + async def test_start_raises_without_client_secret(self): config = DingTalkConfig(client_id="id", client_secret="") channel = DingTalkChannel(config) with pytest.raises(ChannelError, match="client_id and client_secret"): - _run(channel.start()) + await channel.start() - def test_stop_when_not_running(self): + async def test_stop_when_not_running(self): config = DingTalkConfig(client_id="test-id", client_secret="test-secret") channel = DingTalkChannel(config) - _run(channel.stop()) + await channel.stop() - def test_send_returns_false_without_client(self): + async def test_send_returns_false_without_client(self): config = DingTalkConfig(client_id="test-id", client_secret="test-secret") channel = DingTalkChannel(config) msg = OutboundMessage( @@ -73,7 +72,7 @@ class TestDingTalkChannel: content="hello", metadata={"chat_id": "user123"}, ) - result = _run(channel.send(msg)) + result = await channel.send(msg) assert result is False def test_capabilities(self): @@ -130,20 +129,20 @@ class TestDingTalkWsMessageParsing: channel._token_expires = 9999999999 return channel - def test_system_ping_ack(self): + async def test_system_ping_ack(self): channel = self._make_channel() data = { "type": "SYSTEM", "headers": {"topic": "ping", "messageId": "ping-1"}, "data": "pong-data", } - _run(channel._on_ws_message(data)) + await channel._on_ws_message(data) channel._ws_session.send_str.assert_called_once() sent = json.loads(channel._ws_session.send_str.call_args[0][0]) assert sent["code"] == 200 assert sent["data"] == "pong-data" - def test_callback_text_message(self): + async def test_callback_text_message(self): channel = self._make_channel() channel._enqueue_raw = AsyncMock() @@ -158,14 +157,14 @@ class TestDingTalkWsMessageParsing: "headers": {"messageId": "msg-1", "contentType": "application/json"}, "data": json.dumps(payload), } - _run(channel._on_ws_message(data)) + await channel._on_ws_message(data) channel._enqueue_raw.assert_called_once() raw = channel._enqueue_raw.call_args[0][0] assert raw.text == "hello bot" assert raw.sender_id == "staff123" assert raw.is_group is False - def test_callback_group_message_mention(self): + async def test_callback_group_message_mention(self): channel = self._make_channel() channel._enqueue_raw = AsyncMock() @@ -181,12 +180,12 @@ class TestDingTalkWsMessageParsing: "headers": {"messageId": "msg-2"}, "data": json.dumps(payload), } - _run(channel._on_ws_message(data)) + await channel._on_ws_message(data) raw = channel._enqueue_raw.call_args[0][0] assert raw.is_group is True assert raw.was_mentioned is True - def test_callback_group_no_mention(self): + async def test_callback_group_no_mention(self): channel = self._make_channel() channel._enqueue_raw = AsyncMock() @@ -201,12 +200,12 @@ class TestDingTalkWsMessageParsing: "headers": {"messageId": "msg-3"}, "data": json.dumps(payload), } - _run(channel._on_ws_message(data)) + await channel._on_ws_message(data) raw = channel._enqueue_raw.call_args[0][0] assert raw.is_group is True assert raw.was_mentioned is False - def test_ignores_non_callback(self): + async def test_ignores_non_callback(self): channel = self._make_channel() channel._enqueue_raw = AsyncMock() @@ -215,10 +214,10 @@ class TestDingTalkWsMessageParsing: "headers": {"messageId": "msg-x"}, "data": "{}", } - _run(channel._on_ws_message(data)) + await channel._on_ws_message(data) channel._enqueue_raw.assert_not_called() - def test_ignores_empty_content(self): + async def test_ignores_empty_content(self): channel = self._make_channel() channel._enqueue_raw = AsyncMock() @@ -232,20 +231,20 @@ class TestDingTalkWsMessageParsing: "headers": {"messageId": "msg-e"}, "data": json.dumps(payload), } - _run(channel._on_ws_message(data)) + await channel._on_ws_message(data) channel._enqueue_raw.assert_not_called() - def test_non_dict_data_ignored(self): + async def test_non_dict_data_ignored(self): channel = self._make_channel() channel._enqueue_raw = AsyncMock() - _run(channel._on_ws_message("not a dict")) + await channel._on_ws_message("not a dict") channel._enqueue_raw.assert_not_called() class TestDingTalkSendChunk: """Test _send_chunk with mocked HTTP client.""" - def test_send_chunk_calls_api(self): + async def test_send_chunk_calls_api(self): config = DingTalkConfig(client_id="test-app", client_secret="test-secret") channel = DingTalkChannel(config) channel._access_token = "fake-token" @@ -256,7 +255,7 @@ class TestDingTalkSendChunk: channel._http_client = MagicMock() channel._http_client.post = AsyncMock(return_value=mock_response) - _run(channel._send_chunk("user1", "formatted", "raw text", None, {})) + await channel._send_chunk("user1", "formatted", "raw text", None, {}) channel._http_client.post.assert_called_once() call_args = channel._http_client.post.call_args body = call_args.kwargs.get("json") or call_args[1].get("json") @@ -273,21 +272,21 @@ class TestDingTalkChannelRegistration: class TestDingTalkProbe: - def test_missing_credentials(self): + async def test_missing_credentials(self): from EvoScientist.channels.dingtalk.probe import validate_dingtalk - ok, msg = _run(validate_dingtalk("", "")) + ok, msg = await validate_dingtalk("", "") assert ok is False assert "required" in msg - def test_missing_client_id(self): + async def test_missing_client_id(self): from EvoScientist.channels.dingtalk.probe import validate_dingtalk - ok, _msg = _run(validate_dingtalk("", "secret")) + ok, _msg = await validate_dingtalk("", "secret") assert ok is False - def test_missing_client_secret(self): + async def test_missing_client_secret(self): from EvoScientist.channels.dingtalk.probe import validate_dingtalk - ok, _msg = _run(validate_dingtalk("id", "")) + ok, _msg = await validate_dingtalk("id", "") assert ok is False diff --git a/tests/test_discord_channel.py b/tests/test_discord_channel.py index 5cd1342..18eec2e 100644 --- a/tests/test_discord_channel.py +++ b/tests/test_discord_channel.py @@ -4,7 +4,6 @@ import pytest from EvoScientist.channels.base import ChannelError from EvoScientist.channels.discord.channel import DiscordChannel, DiscordConfig -from tests.conftest import run_async as _run class TestDiscordChannel: @@ -14,18 +13,18 @@ class TestDiscordChannel: assert channel.config is config assert channel._running is False - def test_start_raises_without_token_or_library(self): + async def test_start_raises_without_token_or_library(self): config = DiscordConfig(bot_token="") channel = DiscordChannel(config) with pytest.raises(ChannelError): - _run(channel.start()) + await channel.start() - def test_stop_when_not_running(self): + async def test_stop_when_not_running(self): config = DiscordConfig(bot_token="test") channel = DiscordChannel(config) - _run(channel.stop()) + await channel.stop() - def test_send_returns_false_without_client(self): + async def test_send_returns_false_without_client(self): from EvoScientist.channels.base import OutboundMessage config = DiscordConfig(bot_token="test") @@ -36,5 +35,5 @@ class TestDiscordChannel: content="hello", metadata={"chat_id": "123"}, ) - result = _run(channel.send(msg)) + result = await channel.send(msg) assert result is False diff --git a/tests/test_evoskills_command.py b/tests/test_evoskills_command.py index 45fcd6e..cf5f1c7 100644 --- a/tests/test_evoskills_command.py +++ b/tests/test_evoskills_command.py @@ -2,8 +2,6 @@ from unittest.mock import AsyncMock, MagicMock, patch -from tests.conftest import run_async as _run - def _ctx(supports_interactive=True): from EvoScientist.commands.base import CommandContext @@ -31,7 +29,7 @@ _INDEX = [ class TestInstallSkills: - def test_picker_cancel_no_install(self): + async def test_picker_cancel_no_install(self): from EvoScientist.commands.implementation.skills import InstallSkills ctx, ui = _ctx() @@ -45,12 +43,12 @@ class TestInstallSkills: "EvoScientist.tools.skills_manager.install_skill", ) as install_mock, ): - _run(InstallSkills().execute(ctx, [])) + await InstallSkills().execute(ctx, []) install_mock.assert_not_called() msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Browse cancelled" in m for m in msgs) - def test_picker_returns_selections_installs_each(self): + async def test_picker_returns_selections_installs_each(self): from EvoScientist.commands.implementation.skills import InstallSkills ctx, ui = _ctx() @@ -68,10 +66,10 @@ class TestInstallSkills: return_value={"success": True, "name": "x"}, ) as install_mock, ): - _run(InstallSkills().execute(ctx, [])) + await InstallSkills().execute(ctx, []) assert install_mock.call_count == 2 - def test_channel_auto_install_on_tag(self): + async def test_channel_auto_install_on_tag(self): """Non-interactive UI + tag arg → auto-installs matching skills.""" from EvoScientist.commands.implementation.skills import InstallSkills @@ -86,12 +84,12 @@ class TestInstallSkills: return_value={"success": True, "name": "x"}, ) as install_mock, ): - _run(InstallSkills().execute(ctx, ["core"])) + await InstallSkills().execute(ctx, ["core"]) # "core" matches research-ideation only → 1 install, no picker call assert install_mock.call_count == 1 ui.wait_for_skill_browse.assert_not_called() - def test_fetch_failure_prints_error(self): + async def test_fetch_failure_prints_error(self): from EvoScientist.commands.implementation.skills import InstallSkills ctx, ui = _ctx() @@ -99,6 +97,6 @@ class TestInstallSkills: "EvoScientist.tools.skills_manager.fetch_remote_skill_index", side_effect=RuntimeError("network fail"), ): - _run(InstallSkills().execute(ctx, [])) + await InstallSkills().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Failed to fetch" in m for m in msgs) diff --git a/tests/test_exit_command.py b/tests/test_exit_command.py index ee8c613..9df172a 100644 --- a/tests/test_exit_command.py +++ b/tests/test_exit_command.py @@ -2,11 +2,9 @@ from unittest.mock import MagicMock -from tests.conftest import run_async as _run - class TestExitCommand: - def test_execute_calls_force_quit(self): + async def test_execute_calls_force_quit(self): from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.session import ExitCommand @@ -17,7 +15,7 @@ class TestExitCommand: ui=ui, ) cmd = ExitCommand() - _run(cmd.execute(ctx, [])) + await cmd.execute(ctx, []) ui.force_quit.assert_called_once() def test_aliases_registered(self): diff --git a/tests/test_feishu_channel.py b/tests/test_feishu_channel.py index 623bc06..f5f409c 100644 --- a/tests/test_feishu_channel.py +++ b/tests/test_feishu_channel.py @@ -14,7 +14,6 @@ from EvoScientist.channels.feishu.channel import ( _parse_inline_elements, _parse_inline_text, ) -from tests.conftest import run_async as _run class TestFeishuConfig: @@ -57,24 +56,24 @@ class TestFeishuChannel: assert channel._running is False assert channel.name == "feishu" - def test_start_raises_without_app_id(self): + async def test_start_raises_without_app_id(self): config = FeishuConfig(app_id="", app_secret="test-secret") channel = FeishuChannel(config) with pytest.raises(ChannelError, match="app_id"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_app_secret(self): + async def test_start_raises_without_app_secret(self): config = FeishuConfig(app_id="test-id", app_secret="") channel = FeishuChannel(config) with pytest.raises(ChannelError, match="app_secret"): - _run(channel.start()) + await channel.start() - def test_stop_when_not_running(self): + async def test_stop_when_not_running(self): config = FeishuConfig(app_id="test-id", app_secret="test-secret") channel = FeishuChannel(config) - _run(channel.stop()) + await channel.stop() - def test_send_returns_false_without_client(self): + async def test_send_returns_false_without_client(self): config = FeishuConfig(app_id="test-id", app_secret="test-secret") channel = FeishuChannel(config) msg = OutboundMessage( @@ -83,7 +82,7 @@ class TestFeishuChannel: content="hello", metadata={"chat_id": "oc_test"}, ) - result = _run(channel.send(msg)) + result = await channel.send(msg) assert result is False def test_capabilities(self): @@ -201,7 +200,7 @@ class TestFeishuWebhookEvent: channel._enqueue_raw = AsyncMock() return channel - def test_text_message_v2(self): + async def test_text_message_v2(self): channel = self._make_channel() event = { "sender": { @@ -217,7 +216,7 @@ class TestFeishuWebhookEvent: "create_time": "1700000000000", }, } - _run(channel._on_message(event)) + await channel._on_message(event) channel._enqueue_raw.assert_called_once() raw = channel._enqueue_raw.call_args[0][0] assert raw.text == "hello feishu" @@ -225,7 +224,7 @@ class TestFeishuWebhookEvent: assert raw.chat_id == "oc_chat1" assert raw.is_group is False - def test_group_message_with_mention(self): + async def test_group_message_with_mention(self): channel = self._make_channel() event = { "sender": { @@ -242,13 +241,13 @@ class TestFeishuWebhookEvent: "mentions": [{"key": "@_user_1", "id": {}}], }, } - _run(channel._on_message(event)) + await channel._on_message(event) raw = channel._enqueue_raw.call_args[0][0] assert raw.is_group is True assert raw.was_mentioned is True assert channel._mention_names == ["@_user_1"] - def test_group_message_no_mention(self): + async def test_group_message_no_mention(self): channel = self._make_channel() event = { "sender": { @@ -264,12 +263,12 @@ class TestFeishuWebhookEvent: "create_time": "1700000000000", }, } - _run(channel._on_message(event)) + await channel._on_message(event) raw = channel._enqueue_raw.call_args[0][0] assert raw.is_group is True assert raw.was_mentioned is False - def test_skips_bot_messages(self): + async def test_skips_bot_messages(self): channel = self._make_channel() event = { "sender": { @@ -283,10 +282,10 @@ class TestFeishuWebhookEvent: "content": json.dumps({"text": "bot reply"}), }, } - _run(channel._on_message(event)) + await channel._on_message(event) channel._enqueue_raw.assert_not_called() - def test_post_message(self): + async def test_post_message(self): channel = self._make_channel() post_content = { "zh_cn": { @@ -308,12 +307,12 @@ class TestFeishuWebhookEvent: "create_time": "1700000000000", }, } - _run(channel._on_message(event)) + await channel._on_message(event) raw = channel._enqueue_raw.call_args[0][0] assert "Test" in raw.text assert "Post body" in raw.text - def test_unsupported_msg_type_annotation(self): + async def test_unsupported_msg_type_annotation(self): channel = self._make_channel() event = { "sender": { @@ -329,7 +328,7 @@ class TestFeishuWebhookEvent: "create_time": "1700000000000", }, } - _run(channel._on_message(event)) + await channel._on_message(event) raw = channel._enqueue_raw.call_args[0][0] assert "share_chat" in raw.text @@ -337,7 +336,7 @@ class TestFeishuWebhookEvent: class TestFeishuSendChunk: """Test _send_chunk with mocked HTTP client.""" - def test_send_chunk_post_format(self): + async def test_send_chunk_post_format(self): config = FeishuConfig(app_id="test-app", app_secret="test-secret") channel = FeishuChannel(config) channel._access_token = "fake-token" @@ -348,14 +347,14 @@ class TestFeishuSendChunk: channel._http_client = MagicMock() channel._http_client.post = AsyncMock(return_value=mock_response) - _run(channel._send_chunk("oc_chat1", "formatted", "raw **text**", None, {})) + await channel._send_chunk("oc_chat1", "formatted", "raw **text**", None, {}) channel._http_client.post.assert_called() # Should try post format first call_args = channel._http_client.post.call_args body = call_args.kwargs.get("json") or call_args[1].get("json") assert body["receive_id"] == "oc_chat1" - def test_send_chunk_with_reply(self): + async def test_send_chunk_with_reply(self): config = FeishuConfig(app_id="test-app", app_secret="test-secret") channel = FeishuChannel(config) channel._access_token = "fake-token" @@ -366,7 +365,7 @@ class TestFeishuSendChunk: channel._http_client = MagicMock() channel._http_client.post = AsyncMock(return_value=mock_response) - _run(channel._send_chunk("oc_chat1", "reply", "reply text", "om_reply_id", {})) + await channel._send_chunk("oc_chat1", "reply", "reply text", "om_reply_id", {}) # Should call the reply API endpoint first_call_url = channel._http_client.post.call_args_list[0][0][0] assert "reply" in first_call_url @@ -484,17 +483,17 @@ class TestFeishuChannelRegistration: class TestFeishuProbe: - def test_missing_app_id(self): + async def test_missing_app_id(self): from EvoScientist.channels.feishu.probe import validate_feishu_credentials - ok, msg = _run(validate_feishu_credentials("", "secret")) + ok, msg = await validate_feishu_credentials("", "secret") assert ok is False assert "app_id" in msg - def test_missing_app_secret(self): + async def test_missing_app_secret(self): from EvoScientist.channels.feishu.probe import validate_feishu_credentials - ok, msg = _run(validate_feishu_credentials("id", "")) + ok, msg = await validate_feishu_credentials("id", "") assert ok is False assert "app_secret" in msg @@ -510,7 +509,7 @@ class TestFeishuWebSocketMode: ) assert config.subscription_mode == "websocket" - def test_start_websocket_raises_without_lark_oapi(self): + async def test_start_websocket_raises_without_lark_oapi(self): config = FeishuConfig( app_id="test-id", app_secret="test-secret", @@ -520,9 +519,9 @@ class TestFeishuWebSocketMode: # Temporarily hide lark_oapi if it's installed with patch.dict(sys.modules, {"lark_oapi": None}): with pytest.raises(ChannelError, match="lark-oapi"): - _run(channel.start()) + await channel.start() - def test_start_webhook_mode_still_works(self): + async def test_start_webhook_mode_still_works(self): """Ensure subscription_mode='webhook' still validates as before.""" config = FeishuConfig( app_id="", @@ -531,9 +530,9 @@ class TestFeishuWebSocketMode: ) channel = FeishuChannel(config) with pytest.raises(ChannelError, match="app_id"): - _run(channel.start()) + await channel.start() - def test_invalid_subscription_mode_raises(self): + async def test_invalid_subscription_mode_raises(self): config = FeishuConfig( app_id="test-id", app_secret="test-secret", @@ -541,9 +540,9 @@ class TestFeishuWebSocketMode: ) channel = FeishuChannel(config) with pytest.raises(ChannelError, match="Invalid feishu_subscription_mode"): - _run(channel.start()) + await channel.start() - def test_on_lark_sdk_message_bridges_to_on_message(self): + async def test_on_lark_sdk_message_bridges_to_on_message(self): """Test that _on_lark_sdk_message enqueues event dict via queue.""" import queue as queue_mod @@ -594,14 +593,14 @@ class TestFeishuWebSocketMode: ) # Verify the consumer processes it correctly - _run(channel._on_message(event_dict)) + await channel._on_message(event_dict) channel._enqueue_raw.assert_called_once() raw = channel._enqueue_raw.call_args[0][0] assert raw.text == "hello from websocket" assert raw.sender_id == "ou_test_ws" assert raw.is_group is False - def test_cleanup_websocket_mode(self): + async def test_cleanup_websocket_mode(self): config = FeishuConfig( app_id="test-id", app_secret="test-secret", @@ -617,7 +616,7 @@ class TestFeishuWebSocketMode: channel._ws_consumer_task = None channel._access_token = "fake-token" - _run(channel._cleanup()) + await channel._cleanup() mock_client.aclose.assert_called_once() assert channel._http_client is None diff --git a/tests/test_gateway_background_runs.py b/tests/test_gateway_background_runs.py index 27ff1a0..ae7ea78 100644 --- a/tests/test_gateway_background_runs.py +++ b/tests/test_gateway_background_runs.py @@ -1,6 +1,5 @@ from __future__ import annotations -import asyncio from types import SimpleNamespace from unittest.mock import MagicMock @@ -179,7 +178,7 @@ def test_launch_background_run_deletes_thread_when_run_creation_fails(monkeypatc fake_client.threads.delete.assert_called_once_with("thread-1") -def test_async_launch_background_run_deletes_thread_when_run_creation_fails( +async def test_async_launch_background_run_deletes_thread_when_run_creation_fails( monkeypatch, ): monkeypatch.setattr( @@ -204,11 +203,8 @@ def test_async_launch_background_run_deletes_thread_when_run_creation_fails( lambda **_kwargs: SimpleNamespace(threads=_Threads(), runs=_Runs()), ) - async def run() -> None: - with pytest.raises(RuntimeError, match="run creation failed"): - await background_runs.alaunch_background_run(_request()) - - asyncio.run(run()) + with pytest.raises(RuntimeError, match="run creation failed"): + await background_runs.alaunch_background_run(_request()) assert deleted == ["thread-1"] @@ -278,7 +274,7 @@ def test_sync_status_watcher_preserves_thread_on_poll_failure( assert deleted == [] -def test_async_status_watcher_aborts_and_deletes_thread_on_error_status(): +async def test_async_status_watcher_aborts_and_deletes_thread_on_error_status(): finished: list[background_runs.BackgroundRun] = [] aborted: list[background_runs.BackgroundRun] = [] deleted: list[str] = [] @@ -291,29 +287,26 @@ def test_async_status_watcher_aborts_and_deletes_thread_on_error_status(): async def delete(self, thread_id: str): deleted.append(thread_id) - async def run() -> None: - await background_runs.awatch_background_run( - SimpleNamespace(runs=_Runs(), threads=_Threads()), - thread_id="thread-1", - run_id="run-1", - name="test worker", - hooks=background_runs.BackgroundRunHooks( - on_finished=finished.append, - on_aborted=aborted.append, - ), - watcher_config=background_runs.BackgroundRunWatcherConfig( - poll_interval_seconds=0, - ), - ) - - asyncio.run(run()) + await background_runs.awatch_background_run( + SimpleNamespace(runs=_Runs(), threads=_Threads()), + thread_id="thread-1", + run_id="run-1", + name="test worker", + hooks=background_runs.BackgroundRunHooks( + on_finished=finished.append, + on_aborted=aborted.append, + ), + watcher_config=background_runs.BackgroundRunWatcherConfig( + poll_interval_seconds=0, + ), + ) assert finished == [] assert [run.run_id for run in aborted] == ["run-1"] assert deleted == ["thread-1"] -def test_async_status_watcher_preserves_run_url(): +async def test_async_status_watcher_preserves_run_url(): finished: list[background_runs.BackgroundRun] = [] class _Runs: @@ -324,21 +317,18 @@ def test_async_status_watcher_preserves_run_url(): async def delete(self, _thread_id: str): return None - async def run() -> None: - await background_runs.awatch_background_run( - SimpleNamespace(runs=_Runs(), threads=_Threads()), - url="http://worker.example", - thread_id="thread-1", - run_id="run-1", - name="test worker", - hooks=background_runs.BackgroundRunHooks( - on_finished=finished.append, - ), - watcher_config=background_runs.BackgroundRunWatcherConfig( - poll_interval_seconds=0, - ), - ) - - asyncio.run(run()) + await background_runs.awatch_background_run( + SimpleNamespace(runs=_Runs(), threads=_Threads()), + url="http://worker.example", + thread_id="thread-1", + run_id="run-1", + name="test worker", + hooks=background_runs.BackgroundRunHooks( + on_finished=finished.append, + ), + watcher_config=background_runs.BackgroundRunWatcherConfig( + poll_interval_seconds=0, + ), + ) assert [run.url for run in finished] == ["http://worker.example"] diff --git a/tests/test_graph_gateway.py b/tests/test_graph_gateway.py index a773a48..ebd2859 100644 --- a/tests/test_graph_gateway.py +++ b/tests/test_graph_gateway.py @@ -20,7 +20,6 @@ from EvoScientist.gateway import ( ) from EvoScientist.gateway.server import _THREAD_SEARCH_LIMIT from EvoScientist.stream import display as display_mod -from tests.conftest import run_async from tests.fakes import ( FakeGraphGateway, FakeLangGraphClient, @@ -30,7 +29,7 @@ from tests.fakes import ( ) -def test_local_gateway_streams_from_injected_streamer(): +async def test_local_gateway_streams_from_injected_streamer(): seen: dict[str, Any] = {} async def _streamer(agent, message, thread_id, **kwargs): @@ -60,7 +59,7 @@ def test_local_gateway_streams_from_injected_streamer(): return [event async for event in gateway.stream_events(request)] with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer): - events = run_async(_collect()) + events = await _collect() assert events == [ {"type": "text", "content": "hi"}, @@ -75,7 +74,7 @@ def test_local_gateway_streams_from_injected_streamer(): } -def test_local_graph_gateway_delegates_thread_operations(): +async def test_local_graph_gateway_delegates_thread_operations(): thread_store = FakeThreadStore( generated_thread_id="new12345", threads=[{"thread_id": "abc12345"}], @@ -102,7 +101,7 @@ def test_local_graph_gateway_delegates_thread_operations(): "deleted": await gateway.delete_thread("abc12345"), } - result = run_async(_run()) + result = await _run() assert result["created"] == "new12345" assert result["threads"] == [{"thread_id": "abc12345"}] @@ -132,16 +131,14 @@ def test_local_graph_gateway_delegates_thread_operations(): ] -def test_local_graph_gateway_reads_state_values(): +async def test_local_graph_gateway_reads_state_values(): agent = MagicMock() agent.aget_state = AsyncMock( return_value=SimpleNamespace(values={"async_tasks": {"task-1": {}}}) ) gateway = LocalGraphGateway() - values = run_async( - gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345") - ) + values = await gateway.get_state_values(GraphTarget(local_graph=agent), "abc12345") assert values == {"async_tasks": {"task-1": {}}} agent.aget_state.assert_awaited_once_with( @@ -149,17 +146,15 @@ def test_local_graph_gateway_reads_state_values(): ) -def test_local_graph_gateway_updates_state_values(): +async def test_local_graph_gateway_updates_state_values(): agent = MagicMock() agent.aupdate_state = AsyncMock() gateway = LocalGraphGateway() - run_async( - gateway.update_state_values( - GraphTarget(local_graph=agent), - "abc12345", - {"_summarization_event": {"cutoff_index": 2}}, - ) + await gateway.update_state_values( + GraphTarget(local_graph=agent), + "abc12345", + {"_summarization_event": {"cutoff_index": 2}}, ) agent.aupdate_state.assert_awaited_once_with( @@ -169,7 +164,7 @@ def test_local_graph_gateway_updates_state_values(): ) -def test_local_stream_events_delegates_aclose_to_inner(): +async def test_local_stream_events_delegates_aclose_to_inner(): cleanup_ran = False async def _streamer(_agent, _message, _thread_id, **_kwargs): @@ -194,7 +189,7 @@ def test_local_stream_events_delegates_aclose_to_inner(): assert cleanup_ran is True with patch("EvoScientist.stream.events.stream_agent_events", new=_streamer): - run_async(_run()) + await _run() def test_run_streaming_can_consume_injected_gateway(): @@ -228,7 +223,7 @@ def test_run_streaming_can_consume_injected_gateway(): ] -def test_resume_command_consumes_context_gateway(): +async def test_resume_command_consumes_context_gateway(): from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.session import ResumeCommand @@ -246,7 +241,7 @@ def test_resume_command_consumes_context_gateway(): graph_gateway=FakeGraphGateway(thread_store=thread_store), ) - run_async(ResumeCommand().execute(ctx, ["abc"])) + await ResumeCommand().execute(ctx, ["abc"]) assert ctx.thread_id == "abc12345" assert ctx.workspace_dir == "/restored" @@ -287,7 +282,7 @@ def test_cmd_run_passes_local_graph_gateway(monkeypatch): assert seen["gateway"].thread_store is thread_store -def test_langgraph_server_thread_store_delegates_to_sdk_threads(): +async def test_langgraph_server_thread_store_delegates_to_sdk_threads(): threads = FakeLangGraphThreadsClient( threads=[ { @@ -335,7 +330,7 @@ def test_langgraph_server_thread_store_delegates_to_sdk_threads(): "deleted": await store.delete_thread("abc12345"), } - result = run_async(_run()) + result = await _run() assert result["created"] == "server-thread" assert len(threads.created) == 1 @@ -379,7 +374,7 @@ def test_langgraph_server_thread_store_delegates_to_sdk_threads(): assert threads.deleted == ["abc12345"] -def test_langgraph_server_thread_store_limit_zero_pages_all_threads(): +async def test_langgraph_server_thread_store_limit_zero_pages_all_threads(): rows = [ { "thread_id": f"thread-{index}", @@ -392,7 +387,7 @@ def test_langgraph_server_thread_store_limit_zero_pages_all_threads(): client=FakeLangGraphClient(threads), ) - result = run_async(store.list_threads(limit=0)) + result = await store.list_threads(limit=0) assert [row["thread_id"] for row in result] == [ f"thread-{index}" for index in range(_THREAD_SEARCH_LIMIT + 1) @@ -403,7 +398,7 @@ def test_langgraph_server_thread_store_limit_zero_pages_all_threads(): ] -def test_langgraph_server_thread_store_positive_limit_uses_single_search(): +async def test_langgraph_server_thread_store_positive_limit_uses_single_search(): threads = FakeLangGraphThreadsClient( threads=[ { @@ -417,7 +412,7 @@ def test_langgraph_server_thread_store_positive_limit_uses_single_search(): client=FakeLangGraphClient(threads), ) - result = run_async(store.list_threads(limit=2)) + result = await store.list_threads(limit=2) assert [row["thread_id"] for row in result] == ["thread-0", "thread-1"] assert [(search["limit"], search["offset"]) for search in threads.searches] == [ @@ -425,7 +420,7 @@ def test_langgraph_server_thread_store_positive_limit_uses_single_search(): ] -def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup(): +async def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup(): threads = FakeLangGraphThreadsClient( threads=[ { @@ -438,14 +433,14 @@ def test_langgraph_server_thread_store_prefix_resolution_skips_exact_lookup(): client=FakeLangGraphClient(threads), ) - result = run_async(store.resolve_thread_id_prefix("abc")) + result = await store.resolve_thread_id_prefix("abc") assert result == ("abc12345", []) assert threads.gets == [] assert len(threads.searches) == 1 -def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads(): +async def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads(): rows = [ { "thread_id": f"thread-{index}", @@ -464,7 +459,7 @@ def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads(): client=FakeLangGraphClient(threads), ) - result = run_async(store.resolve_thread_id_prefix("older-thread")) + result = await store.resolve_thread_id_prefix("older-thread") assert result == ("older-thread-match", []) assert [(search["limit"], search["offset"]) for search in threads.searches] == [ @@ -473,7 +468,7 @@ def test_langgraph_server_thread_store_prefix_resolution_pages_all_threads(): ] -def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup(): +async def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup(): thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f" threads = FakeLangGraphThreadsClient( threads=[ @@ -487,14 +482,14 @@ def test_langgraph_server_thread_store_uuid_resolution_uses_exact_lookup(): client=FakeLangGraphClient(threads), ) - result = run_async(store.resolve_thread_id_prefix(thread_id)) + result = await store.resolve_thread_id_prefix(thread_id) assert result == (thread_id, []) assert threads.gets == [thread_id] assert threads.searches == [] -def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id(): +async def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id(): thread_id = "019ed9e4-4253-7f62-b50f-f0470a4b3c9f" threads = FakeLangGraphThreadsClient( threads=[ @@ -508,7 +503,7 @@ def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id(): client=FakeLangGraphClient(threads), ) - result = run_async(store.resolve_thread_id_prefix(thread_id)) + result = await store.resolve_thread_id_prefix(thread_id) assert result == (None, []) assert threads.gets == [thread_id] @@ -517,7 +512,7 @@ def test_langgraph_server_thread_store_uuid_resolution_filters_graph_id(): ] -def test_langgraph_server_thread_store_clones_thread_with_metadata(): +async def test_langgraph_server_thread_store_clones_thread_with_metadata(): clone_metadata = { "clone_purpose": "memory_extraction", "source_thread_id": "source-thread", @@ -534,8 +529,8 @@ def test_langgraph_server_thread_store_clones_thread_with_metadata(): client=FakeLangGraphClient(threads), ) - cloned_thread_id = run_async( - store.clone_thread("source-thread", metadata=clone_metadata) + cloned_thread_id = await store.clone_thread( + "source-thread", metadata=clone_metadata ) assert cloned_thread_id == "source-thread-copy" @@ -552,7 +547,7 @@ def test_langgraph_server_thread_store_clones_thread_with_metadata(): } -def test_langgraph_server_thread_store_rejects_copy_without_thread_id(): +async def test_langgraph_server_thread_store_rejects_copy_without_thread_id(): threads = FakeLangGraphThreadsClient( threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}], copy_response=None, @@ -565,10 +560,10 @@ def test_langgraph_server_thread_store_rejects_copy_without_thread_id(): await store.clone_thread("source-thread") with pytest.raises(RuntimeError, match="did not return a cloned thread id"): - run_async(_run()) + await _run() -def test_langgraph_server_gateway_clones_thread(): +async def test_langgraph_server_gateway_clones_thread(): threads = FakeLangGraphThreadsClient( threads=[{"thread_id": "source-thread", "metadata": {"graph_id": "agent"}}] ) @@ -578,12 +573,10 @@ def test_langgraph_server_gateway_clones_thread(): ) ) - cloned_thread_id = run_async( - gateway.clone_thread( - "source-thread", - metadata={"clone_purpose": "manual"}, - target=GraphTarget(graph_id="agent"), - ) + cloned_thread_id = await gateway.clone_thread( + "source-thread", + metadata={"clone_purpose": "manual"}, + target=GraphTarget(graph_id="agent"), ) assert cloned_thread_id == "source-thread-copy" @@ -592,12 +585,12 @@ def test_langgraph_server_gateway_clones_thread(): ] -def test_local_graph_gateway_clone_thread_is_explicitly_unsupported(): +async def test_local_graph_gateway_clone_thread_is_explicitly_unsupported(): async def _run(): await LocalGraphGateway().clone_thread("source-thread") with pytest.raises(NotImplementedError, match="does not support thread cloning"): - run_async(_run()) + await _run() def test_runtime_gateways_can_use_langgraph_server_backend(): @@ -616,7 +609,7 @@ def test_runtime_gateways_can_use_langgraph_server_backend(): assert gateway.thread_store is runtime_gateways.thread_store -def test_langgraph_server_gateway_reads_state_values(): +async def test_langgraph_server_gateway_reads_state_values(): threads = FakeLangGraphThreadsClient( threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], states={"abc12345": {"values": {"async_tasks": {"task-1": {}}}}}, @@ -627,12 +620,12 @@ def test_langgraph_server_gateway_reads_state_values(): ) ) - values = run_async(gateway.get_state_values(GraphTarget(), "abc12345")) + values = await gateway.get_state_values(GraphTarget(), "abc12345") assert values == {"async_tasks": {"task-1": {}}} -def test_langgraph_server_gateway_messages_apply_summarization_event(): +async def test_langgraph_server_gateway_messages_apply_summarization_event(): threads = FakeLangGraphThreadsClient( threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], states={ @@ -658,7 +651,7 @@ def test_langgraph_server_gateway_messages_apply_summarization_event(): ) ) - messages = run_async(gateway.get_thread_messages("abc12345")) + messages = await gateway.get_thread_messages("abc12345") assert len(messages) == 2 assert isinstance(messages[0], AIMessage) @@ -667,7 +660,7 @@ def test_langgraph_server_gateway_messages_apply_summarization_event(): assert messages[1].content == "third" -def test_langgraph_server_gateway_updates_state_values(): +async def test_langgraph_server_gateway_updates_state_values(): threads = FakeLangGraphThreadsClient( threads=[{"thread_id": "abc12345", "metadata": {"graph_id": "EvoScientist"}}], ) @@ -677,12 +670,10 @@ def test_langgraph_server_gateway_updates_state_values(): ) ) - run_async( - gateway.update_state_values( - GraphTarget(), - "abc12345", - {"_summarization_event": {"cutoff_index": 2}}, - ) + await gateway.update_state_values( + GraphTarget(), + "abc12345", + {"_summarization_event": {"cutoff_index": 2}}, ) assert threads.state_updates == [ @@ -690,7 +681,7 @@ def test_langgraph_server_gateway_updates_state_values(): ] -def test_langgraph_server_gateway_streams_root_protocol_events(): +async def test_langgraph_server_gateway_streams_root_protocol_events(): stream = FakeLangGraphThreadStream( "abc12345", events=[ @@ -737,7 +728,7 @@ def test_langgraph_server_gateway_streams_root_protocol_events(): ) ] - events = run_async(_collect()) + events = await _collect() assert len(threads.created) == 1 assert threads.created[0]["thread_id"] == "abc12345" @@ -804,7 +795,7 @@ def _root_message_finish() -> dict[str, object]: } -def _collect_server_gateway_stream( +async def _collect_server_gateway_stream( events: list[dict[str, object]], *, state_messages: list[dict[str, object]] | None = None, @@ -832,11 +823,11 @@ def _collect_server_gateway_stream( ) ] - return run_async(_collect()) + return await _collect() -def test_langgraph_server_gateway_streams_value_message_snapshots(): - events = _collect_server_gateway_stream( +async def test_langgraph_server_gateway_streams_value_message_snapshots(): + events = await _collect_server_gateway_stream( [ _value_snapshot([_OLD_AI, _HUMAN]), _value_snapshot([_OLD_AI, _HUMAN, _NEW_AI]), @@ -850,8 +841,8 @@ def test_langgraph_server_gateway_streams_value_message_snapshots(): ] -def test_langgraph_server_gateway_values_do_not_duplicate_message_stream(): - events = _collect_server_gateway_stream( +async def test_langgraph_server_gateway_values_do_not_duplicate_message_stream(): + events = await _collect_server_gateway_stream( [ _root_text_delta("new"), _root_message_finish(), @@ -866,8 +857,8 @@ def test_langgraph_server_gateway_values_do_not_duplicate_message_stream(): ] -def test_langgraph_server_gateway_ignores_non_root_value_messages(): - events = _collect_server_gateway_stream( +async def test_langgraph_server_gateway_ignores_non_root_value_messages(): + events = await _collect_server_gateway_stream( [ _value_snapshot( [{"type": "ai", "content": "subagent text", "id": "subagent-ai"}], @@ -880,7 +871,7 @@ def test_langgraph_server_gateway_ignores_non_root_value_messages(): assert events[-1] == {"type": "done", "content": "", "response": ""} -def test_langgraph_server_gateway_emits_state_interrupt_before_done(): +async def test_langgraph_server_gateway_emits_state_interrupt_before_done(): stream = FakeLangGraphThreadStream( "abc12345", events=[], @@ -930,7 +921,7 @@ def test_langgraph_server_gateway_emits_state_interrupt_before_done(): ) ] - events = run_async(_collect()) + events = await _collect() assert events == [ { @@ -954,7 +945,7 @@ def test_langgraph_server_gateway_emits_state_interrupt_before_done(): ] -def test_langgraph_server_gateway_streams_subagent_protocol_events(): +async def test_langgraph_server_gateway_streams_subagent_protocol_events(): stream = FakeLangGraphThreadStream( "abc12345", events=[ @@ -1003,7 +994,7 @@ def test_langgraph_server_gateway_streams_subagent_protocol_events(): ) ] - events = run_async(_collect()) + events = await _collect() assert events == [ { @@ -1028,7 +1019,7 @@ def test_langgraph_server_gateway_streams_subagent_protocol_events(): ] -def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream(): +async def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream(): from langgraph.types import Command stream = FakeLangGraphThreadStream( @@ -1058,7 +1049,7 @@ def test_langgraph_server_gateway_resumes_interrupt_with_thread_stream(): ) ] - events = run_async(_collect()) + events = await _collect() assert stream.run.starts == [] assert stream.run.responses == [ diff --git a/tests/test_hitl.py b/tests/test_hitl.py index d235815..61d1f80 100644 --- a/tests/test_hitl.py +++ b/tests/test_hitl.py @@ -378,7 +378,7 @@ class TestHitlConfig: class TestInterruptEventParsing: - def test_interrupt_from_updates_mode(self): + async def test_interrupt_from_updates_mode(self): """__interrupt__ in updates mode yields interrupt event.""" interrupt_data = { "__interrupt__": [ @@ -405,7 +405,7 @@ class TestInterruptEventParsing: protocol_event("updates", interrupt_data), ] ) - events = collect_events(agent, message="test", thread_id="thread-1") + events = await collect_events(agent, message="test", thread_id="thread-1") types = [e["type"] for e in events] assert "interrupt" in types @@ -415,14 +415,14 @@ class TestInterruptEventParsing: assert interrupt_ev["action_requests"][0]["name"] == "execute" assert interrupt_ev["interrupt_id"] == "main" - def test_updates_without_interrupt_skipped(self): + async def test_updates_without_interrupt_skipped(self): """Regular updates mode data is skipped as before.""" agent = FakeV3Agent( [ protocol_event("updates", {"some_node": {"key": "value"}}), ] ) - events = collect_events(agent, message="test", thread_id="thread-1") + events = await collect_events(agent, message="test", thread_id="thread-1") types = [e["type"] for e in events] assert "interrupt" not in types diff --git a/tests/test_install_skill_command.py b/tests/test_install_skill_command.py index 9214637..2a0ad68 100644 --- a/tests/test_install_skill_command.py +++ b/tests/test_install_skill_command.py @@ -2,8 +2,6 @@ from unittest.mock import MagicMock, patch -from tests.conftest import run_async as _run - def _ctx(): from EvoScientist.commands.base import CommandContext @@ -14,15 +12,15 @@ def _ctx(): class TestInstallSkill: - def test_usage_message_when_no_args(self): + async def test_usage_message_when_no_args(self): from EvoScientist.commands.implementation.skills import InstallSkill ctx, ui = _ctx() - _run(InstallSkill().execute(ctx, [])) + await InstallSkill().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Usage:" in m for m in msgs) - def test_happy_path(self): + async def test_happy_path(self): from EvoScientist.commands.implementation.skills import InstallSkill ctx, ui = _ctx() @@ -35,21 +33,21 @@ class TestInstallSkill: "path": "/tmp/demo", }, ): - _run(InstallSkill().execute(ctx, ["./some-path"])) + await InstallSkill().execute(ctx, ["./some-path"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Installed: demo-skill" in m for m in msgs) class TestUninstallSkill: - def test_usage_message_when_no_args(self): + async def test_usage_message_when_no_args(self): from EvoScientist.commands.implementation.skills import UninstallSkill ctx, ui = _ctx() - _run(UninstallSkill().execute(ctx, [])) + await UninstallSkill().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Usage:" in m for m in msgs) - def test_uninstall_success(self): + async def test_uninstall_success(self): from EvoScientist.commands.implementation.skills import UninstallSkill ctx, ui = _ctx() @@ -57,11 +55,11 @@ class TestUninstallSkill: "EvoScientist.tools.skills_manager.uninstall_skill", return_value={"success": True}, ): - _run(UninstallSkill().execute(ctx, ["demo-skill"])) + await UninstallSkill().execute(ctx, ["demo-skill"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Uninstalled: demo-skill" in m for m in msgs) - def test_uninstall_failure(self): + async def test_uninstall_failure(self): from EvoScientist.commands.implementation.skills import UninstallSkill ctx, ui = _ctx() @@ -69,6 +67,6 @@ class TestUninstallSkill: "EvoScientist.tools.skills_manager.uninstall_skill", return_value={"success": False, "error": "not found"}, ): - _run(UninstallSkill().execute(ctx, ["missing"])) + await UninstallSkill().execute(ctx, ["missing"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Failed: not found" in m for m in msgs) diff --git a/tests/test_json_sink.py b/tests/test_json_sink.py index 22fa7ab..36333a0 100644 --- a/tests/test_json_sink.py +++ b/tests/test_json_sink.py @@ -19,7 +19,7 @@ async def _agen(items): yield item -def test_writes_each_event_as_one_jsonl_line(run_async): +async def test_writes_each_event_as_one_jsonl_line(): """Each event dict is serialized to exactly one JSON line, in order.""" events = [ {"type": "thinking", "content": "hmm", "id": 0}, @@ -34,7 +34,7 @@ def test_writes_each_event_as_one_jsonl_line(run_async): ] out = io.StringIO() - run_async(write_events_as_json(_agen(events), out)) + await write_events_as_json(_agen(events), out) lines = out.getvalue().splitlines() assert len(lines) == len(events) @@ -43,7 +43,7 @@ def test_writes_each_event_as_one_jsonl_line(run_async): assert parsed[2]["args"] == {"path": "a.md"} -def test_returns_final_response_from_done_event(run_async): +async def test_returns_final_response_from_done_event(): """The sink returns the response text carried by the terminal `done` event.""" events = [ {"type": "text", "content": "partial"}, @@ -51,12 +51,12 @@ def test_returns_final_response_from_done_event(run_async): ] out = io.StringIO() - result = run_async(write_events_as_json(_agen(events), out)) + result = await write_events_as_json(_agen(events), out) assert result == "the answer" -def test_non_serializable_arg_does_not_crash_the_stream(run_async): +async def test_non_serializable_arg_does_not_crash_the_stream(): """A non-JSON-serializable value degrades to its str form instead of raising.""" class Weird: @@ -72,7 +72,7 @@ def test_non_serializable_arg_does_not_crash_the_stream(run_async): ] out = io.StringIO() - run_async(write_events_as_json(_agen(events), out)) + await write_events_as_json(_agen(events), out) lines = out.getvalue().splitlines() # Both lines must be valid JSON; the non-serializable value falls back to str. @@ -80,7 +80,7 @@ def test_non_serializable_arg_does_not_crash_the_stream(run_async): assert first["args"]["obj"] == "WEIRD" -def test_stream_json_sources_events_from_gateway(run_async): +async def test_stream_json_sources_events_from_gateway(): """stream_json pulls events from gateway.stream_events(request) and serializes them — it does not reach past the gateway abstraction.""" seen: dict[str, object] = {} @@ -98,7 +98,7 @@ def test_stream_json_sources_events_from_gateway(run_async): return _agen(events) out = io.StringIO() - result = run_async(stream_json(_FakeGateway(), object(), out=out)) + result = await stream_json(_FakeGateway(), object(), out=out) assert result == "hi" assert "request" in seen # the request was forwarded to the gateway @@ -106,7 +106,7 @@ def test_stream_json_sources_events_from_gateway(run_async): assert types == ["text", "done"] -def test_stream_json_propagates_gateway_errors(run_async): +async def test_stream_json_propagates_gateway_errors(): """An error from the gateway stream propagates out of stream_json so the CLI dispatch can turn it into a clean exit.""" @@ -124,4 +124,4 @@ def test_stream_json_propagates_gateway_errors(run_async): out = io.StringIO() with pytest.raises(RuntimeError, match="boom"): - run_async(stream_json(_FakeGateway(), object(), out=out)) + await stream_json(_FakeGateway(), object(), out=out) diff --git a/tests/test_langgraph_manager.py b/tests/test_langgraph_manager.py index 845064e..353013c 100644 --- a/tests/test_langgraph_manager.py +++ b/tests/test_langgraph_manager.py @@ -8,6 +8,7 @@ to be available. from __future__ import annotations import dataclasses +import sys from types import SimpleNamespace from unittest.mock import MagicMock, patch @@ -32,6 +33,77 @@ def reset_module_state(): manager._LOG_OFFSET_AT_START = 0 +# ============================================================================= +# langgraph CLI resolution +# ============================================================================= + + +class TestLanggraphCliResolution: + def _make_executable(self, path): + path.write_text("#!/bin/sh\n", encoding="utf-8") + path.chmod(0o755) + + def test_prefers_current_python_environment_over_path(self, tmp_path, monkeypatch): + local_bin = tmp_path / "local" / "bin" + local_bin.mkdir(parents=True) + local_langgraph = local_bin / "langgraph" + self._make_executable(local_langgraph) + + path_bin = tmp_path / "path" / "bin" + path_bin.mkdir(parents=True) + path_langgraph = path_bin / "langgraph" + self._make_executable(path_langgraph) + + monkeypatch.setattr(sys, "executable", str(local_bin / "python")) + monkeypatch.setattr( + manager.shutil, + "which", + lambda command: str(path_langgraph) if command == "langgraph" else None, + ) + + assert manager._langgraph_exe() == str(local_langgraph) + + def test_falls_back_to_path_when_environment_binary_missing( + self, tmp_path, monkeypatch + ): + path_bin = tmp_path / "path" / "bin" + path_bin.mkdir(parents=True) + path_langgraph = path_bin / "langgraph" + self._make_executable(path_langgraph) + + monkeypatch.setattr( + sys, "executable", str(tmp_path / "local" / "bin" / "python") + ) + monkeypatch.setattr( + manager.shutil, + "which", + lambda command: str(path_langgraph) if command == "langgraph" else None, + ) + + assert manager._langgraph_exe() == str(path_langgraph) + + def test_checks_windows_suffix_next_to_current_python(self, tmp_path, monkeypatch): + scripts_dir = tmp_path / "Scripts" + scripts_dir.mkdir() + local_langgraph = scripts_dir / "langgraph.exe" + self._make_executable(local_langgraph) + + path_bin = tmp_path / "path" / "bin" + path_bin.mkdir(parents=True) + path_langgraph = path_bin / "langgraph.exe" + self._make_executable(path_langgraph) + + monkeypatch.setattr(sys, "executable", str(scripts_dir / "python.exe")) + monkeypatch.setattr(manager.os, "name", "nt", raising=False) + monkeypatch.setattr( + manager.shutil, + "which", + lambda command: str(path_langgraph) if command == "langgraph" else None, + ) + + assert manager._langgraph_exe() == str(local_langgraph) + + # ============================================================================= # is_langgraph_dev_running # ============================================================================= diff --git a/tests/test_llm.py b/tests/test_llm.py index 22f83cc..23168d0 100644 --- a/tests/test_llm.py +++ b/tests/test_llm.py @@ -972,7 +972,6 @@ class TestPatchOpenAICompatContent: called_msgs = orig.call_args[0][0] assert called_msgs[0].content == "hello" - @pytest.mark.anyio async def test_agenerate_flattened(self): from langchain_core.messages import HumanMessage @@ -1003,7 +1002,6 @@ class TestPatchOpenAICompatContent: called_msgs = orig.call_args[0][0] assert called_msgs[0].content == "hello" - @pytest.mark.anyio async def test_astream_flattened(self): from langchain_core.messages import HumanMessage @@ -1044,7 +1042,6 @@ class TestPatchOpenAICompatContent: called_msgs = orig.call_args[0][0] assert called_msgs[0].content == [{"type": "text", "text": "see"}, img] - @pytest.mark.anyio async def test_agenerate_preserves_media(self): from langchain_core.messages import HumanMessage @@ -1077,7 +1074,6 @@ class TestPatchOpenAICompatContent: called_msgs = orig.call_args[0][0] assert called_msgs[0].content == [{"type": "text", "text": "see"}, img] - @pytest.mark.anyio async def test_astream_preserves_media(self): from langchain_core.messages import HumanMessage @@ -1550,7 +1546,6 @@ class TestNoVisionFallback: assert out == ["x", "y"] assert len(calls) == 2 - @pytest.mark.anyio async def test_astream_falls_back(self): from unittest.mock import MagicMock diff --git a/tests/test_mcp_client.py b/tests/test_mcp_client.py index 69297a2..5e1bb62 100644 --- a/tests/test_mcp_client.py +++ b/tests/test_mcp_client.py @@ -1399,9 +1399,7 @@ class TestLoadToolsProgressCallback: monkeypatch.setattr(lc_client, "MultiServerMCPClient", _FakeClient) - def test_success_emits_start_then_success_with_tool_count(self, monkeypatch): - import asyncio - + async def test_success_emits_start_then_success_with_tool_count(self, monkeypatch): from EvoScientist.mcp.client import _load_tools events: list[tuple[str, str, str]] = [] @@ -1415,16 +1413,14 @@ class TestLoadToolsProgressCallback: def record(event, name, detail): events.append((event, name, detail)) - asyncio.run(_load_tools(config, on_progress=record)) + await _load_tools(config, on_progress=record) assert events == [ ("start", "srv", ""), ("success", "srv", "3"), ] - def test_failure_emits_start_then_error_with_detail(self, monkeypatch): - import asyncio - + async def test_failure_emits_start_then_error_with_detail(self, monkeypatch): from EvoScientist.mcp.client import _load_tools events: list[tuple[str, str, str]] = [] @@ -1435,16 +1431,14 @@ class TestLoadToolsProgressCallback: def record(event, name, detail): events.append((event, name, detail)) - asyncio.run(_load_tools(config, on_progress=record)) + await _load_tools(config, on_progress=record) assert events == [ ("start", "srv", ""), ("error", "srv", "boom"), ] - def test_mixed_fleet_reports_each_server_independently(self, monkeypatch): - import asyncio - + async def test_mixed_fleet_reports_each_server_independently(self, monkeypatch): from EvoScientist.mcp.client import _load_tools events: list[tuple[str, str, str]] = [] @@ -1464,7 +1458,7 @@ class TestLoadToolsProgressCallback: def record(event, name, detail): events.append((event, name, detail)) - asyncio.run(_load_tools(config, on_progress=record)) + await _load_tools(config, on_progress=record) by_server = {} for ev, name, detail in events: @@ -1472,9 +1466,7 @@ class TestLoadToolsProgressCallback: assert by_server["ok_srv"] == [("start", ""), ("success", "1")] assert by_server["bad_srv"] == [("start", ""), ("error", "refused")] - def test_callback_errors_do_not_break_the_load(self, monkeypatch): - import asyncio - + async def test_callback_errors_do_not_break_the_load(self, monkeypatch): from EvoScientist.mcp.client import _load_tools self._patch_client(monkeypatch, {"srv": ["tool1"]}) @@ -1484,10 +1476,10 @@ class TestLoadToolsProgressCallback: def bad_callback(event, name, detail): raise RuntimeError("callback bug") - result = asyncio.run(_load_tools(config, on_progress=bad_callback)) + result = await _load_tools(config, on_progress=bad_callback) assert result == {"srv": ["tool1"]} - def test_semaphore_caps_concurrent_connections(self, monkeypatch): + async def test_semaphore_caps_concurrent_connections(self, monkeypatch): """Many configured servers must not all spawn at once.""" import asyncio @@ -1516,7 +1508,7 @@ class TestLoadToolsProgressCallback: config = { f"srv{i}": {"transport": "stdio", "command": "demo"} for i in range(10) } - asyncio.run(mcp_client._load_tools(config)) + await mcp_client._load_tools(config) assert inflight["peak"] <= 3 assert inflight["peak"] > 1 # sanity: we *are* parallelizing diff --git a/tests/test_mcp_command.py b/tests/test_mcp_command.py index e28321f..b65ef10 100644 --- a/tests/test_mcp_command.py +++ b/tests/test_mcp_command.py @@ -2,8 +2,6 @@ from unittest.mock import MagicMock, patch -from tests.conftest import run_async as _run - def _ctx(): from EvoScientist.commands.base import CommandContext @@ -14,16 +12,16 @@ def _ctx(): class TestMCPCommandDispatch: - def test_no_args_lists(self): + async def test_no_args_lists(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, ui = _ctx() with patch("EvoScientist.mcp.load_mcp_config", return_value={}): - _run(MCPCommand().execute(ctx, [])) + await MCPCommand().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("No MCP servers configured" in m for m in msgs) - def test_list_subcommand(self): + async def test_list_subcommand(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, ui = _ctx() @@ -31,10 +29,10 @@ class TestMCPCommandDispatch: "srv1": {"transport": "stdio", "tools": ["foo"], "expose_to": ["main"]}, } with patch("EvoScientist.mcp.load_mcp_config", return_value=cfg): - _run(MCPCommand().execute(ctx, ["list"])) + await MCPCommand().execute(ctx, ["list"]) ui.mount_renderable.assert_called_once() - def test_add_subcommand_dispatches(self): + async def test_add_subcommand_dispatches(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, _ui = _ctx() @@ -48,10 +46,10 @@ class TestMCPCommandDispatch: return_value={"transport": "stdio"}, ) as add_mock, ): - _run(MCPCommand().execute(ctx, ["add", "srv1", "python"])) + await MCPCommand().execute(ctx, ["add", "srv1", "python"]) add_mock.assert_called_once() - def test_edit_subcommand_dispatches(self): + async def test_edit_subcommand_dispatches(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, _ui = _ctx() @@ -64,28 +62,28 @@ class TestMCPCommandDispatch: "EvoScientist.mcp.edit_mcp_server", ) as edit_mock, ): - _run(MCPCommand().execute(ctx, ["edit", "srv1", "--tools", "bar"])) + await MCPCommand().execute(ctx, ["edit", "srv1", "--tools", "bar"]) edit_mock.assert_called_once_with("srv1", tools=["bar"]) - def test_remove_subcommand_success(self): + async def test_remove_subcommand_success(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, ui = _ctx() with patch("EvoScientist.mcp.remove_mcp_server", return_value=True): - _run(MCPCommand().execute(ctx, ["remove", "srv1"])) + await MCPCommand().execute(ctx, ["remove", "srv1"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Removed MCP server: srv1" in m for m in msgs) - def test_remove_subcommand_not_found(self): + async def test_remove_subcommand_not_found(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, ui = _ctx() with patch("EvoScientist.mcp.remove_mcp_server", return_value=False): - _run(MCPCommand().execute(ctx, ["remove", "missing"])) + await MCPCommand().execute(ctx, ["remove", "missing"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Server not found" in m for m in msgs) - def test_install_delegates_to_install_mcp_command(self): + async def test_install_delegates_to_install_mcp_command(self): """/mcp install should instantiate InstallMCPCommand and execute it.""" from EvoScientist.commands.implementation.mcp import MCPCommand @@ -101,13 +99,13 @@ class TestMCPCommandDispatch: instance.execute = fake_execute klass.return_value = instance - _run(MCPCommand().execute(ctx, ["install", "foo"])) + await MCPCommand().execute(ctx, ["install", "foo"]) klass.assert_called_once() - def test_unknown_subcommand_prints_help(self): + async def test_unknown_subcommand_prints_help(self): from EvoScientist.commands.implementation.mcp import MCPCommand ctx, ui = _ctx() - _run(MCPCommand().execute(ctx, ["bogus"])) + await MCPCommand().execute(ctx, ["bogus"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("MCP commands:" in m for m in msgs) diff --git a/tests/test_model_command.py b/tests/test_model_command.py index b2ba42d..35141f7 100644 --- a/tests/test_model_command.py +++ b/tests/test_model_command.py @@ -5,8 +5,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest -from tests.conftest import run_async as _run - class TestExtractModelAndProvider: """Unit tests for the argument parser helper.""" @@ -80,7 +78,7 @@ class TestExtractModelAndProvider: class TestModelCommandUnknownModel: """Verify error message for unknown models.""" - def test_unknown_model_shows_error(self): + async def test_unknown_model_shows_error(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -95,7 +93,7 @@ class TestModelCommandUnknownModel: "EvoScientist.EvoScientist._ensure_config", return_value=cfg, ): - _run(cmd.execute(ctx, ["nonexistent-model-xyz"])) + await cmd.execute(ctx, ["nonexistent-model-xyz"]) ui.append_system.assert_called_once() call_args = ui.append_system.call_args @@ -106,7 +104,7 @@ class TestModelCommandUnknownModel: class TestModelCommandPickerCancelled: """Verify no-op when the interactive picker is cancelled.""" - def test_picker_returns_none(self): + async def test_picker_returns_none(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -122,7 +120,7 @@ class TestModelCommandPickerCancelled: "EvoScientist.EvoScientist._ensure_config", return_value=cfg, ): - _run(cmd.execute(ctx, [])) + await cmd.execute(ctx, []) # No model switch should have happened ui.append_system.assert_not_called() @@ -131,7 +129,7 @@ class TestModelCommandPickerCancelled: class TestModelCommandSwitch: """Verify a successful model switch updates config and rebuilds agent.""" - def test_switch_known_model(self): + async def test_switch_known_model(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -158,7 +156,7 @@ class TestModelCommandSwitch: return_value=new_agent, ), ): - _run(cmd.execute(ctx, ["claude-opus-4-8"])) + await cmd.execute(ctx, ["claude-opus-4-8"]) # The switch is committed via set_active_config(temp_cfg), not by # mutating the original cfg object in place. @@ -176,7 +174,7 @@ class TestModelCommandSwitch: assert "claude-opus-4-8" in msg assert "anthropic" in msg - def test_switch_with_save_flag(self): + async def test_switch_with_save_flag(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -203,7 +201,7 @@ class TestModelCommandSwitch: ), patch("EvoScientist.config.settings.set_config_value") as mock_save, ): - _run(cmd.execute(ctx, ["claude-opus-4-8", "--save"])) + await cmd.execute(ctx, ["claude-opus-4-8", "--save"]) # Config file should be updated mock_save.assert_any_call("model", "claude-opus-4-8") @@ -213,7 +211,7 @@ class TestModelCommandSwitch: msg = ui.append_system.call_args[0][0] assert "saved to config" in msg - def test_switch_without_save_flag_does_not_persist(self): + async def test_switch_without_save_flag_does_not_persist(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -240,7 +238,7 @@ class TestModelCommandSwitch: ), patch("EvoScientist.config.settings.set_config_value") as mock_save, ): - _run(cmd.execute(ctx, ["claude-opus-4-8"])) + await cmd.execute(ctx, ["claude-opus-4-8"]) # Config file should NOT be updated mock_save.assert_not_called() @@ -253,7 +251,7 @@ class TestModelCommandSwitch: class TestModelCommandFailure: """Verify error handling when chat-model construction raises.""" - def test_build_chat_model_error(self): + async def test_build_chat_model_error(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -276,7 +274,7 @@ class TestModelCommandFailure: side_effect=RuntimeError("API key missing"), ) as mock_build, ): - _run(cmd.execute(ctx, ["claude-opus-4-8"])) + await cmd.execute(ctx, ["claude-opus-4-8"]) mock_build.assert_called_once() ui.append_system.assert_called_once() @@ -446,7 +444,7 @@ class TestApplyModelIntegration: pair so we can assert on identity. """ - def test_new_agent_is_bound_to_newly_selected_model(self, evo_module_state): + async def test_new_agent_is_bound_to_newly_selected_model(self, evo_module_state): from EvoScientist.commands.implementation.model import ModelCommand from EvoScientist.config.settings import EvoScientistConfig @@ -502,7 +500,7 @@ class TestApplyModelIntegration: ), ): cmd = ModelCommand() - _run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter")) + await cmd._apply_model(ctx, "minimax-m2.7", "openrouter") # The agent produced by _apply_model must be bound to the # NEWLY requested model, threaded in via chat_model=. @@ -539,7 +537,9 @@ class TestApplyModelPreservesConfigByReference: switch (the held object stops being the active ``_config`` after the first). """ - def test_held_config_reference_tracks_repeated_switches(self, evo_module_state): + async def test_held_config_reference_tracks_repeated_switches( + self, evo_module_state + ): from EvoScientist.commands.implementation.model import ModelCommand from EvoScientist.config.settings import EvoScientistConfig @@ -592,7 +592,7 @@ class TestApplyModelPreservesConfigByReference: ("minimax-m2.7", "openrouter"), ("claude-sonnet-4-6", "anthropic"), ]: - _run(cmd._apply_model(ctx, model, provider)) + await cmd._apply_model(ctx, model, provider) # The held reference must reflect the LATEST switch on every # iteration — not just the first — and stay the active config. assert agent_holder["config"].model == model @@ -610,7 +610,7 @@ class TestModelCommandLoadAgentFailure: the ordering could silently regress (e.g. if ``_apply_model`` were reordered to call ``set_chat_model`` first).""" - def test_load_agent_error_is_transactional(self): + async def test_load_agent_error_is_transactional(self): from EvoScientist.commands.implementation.model import ModelCommand cmd = ModelCommand() @@ -646,7 +646,7 @@ class TestModelCommandLoadAgentFailure: # Pass ``--save`` to strengthen the assertion: if the ordering # ever regresses, ``set_config_value`` would be called with # stale data. - _run(cmd.execute(ctx, ["claude-opus-4-8", "--save"])) + await cmd.execute(ctx, ["claude-opus-4-8", "--save"]) # _load_agent was attempted (transactional first step). mock_load.assert_called_once() @@ -677,7 +677,7 @@ class TestApplyModelLoadAgentFailureTransactional: downstream setters never run on failure. """ - def test_globals_unchanged_when_load_agent_raises(self, evo_module_state): + async def test_globals_unchanged_when_load_agent_raises(self, evo_module_state): from EvoScientist.commands.implementation.model import ModelCommand from EvoScientist.config.settings import EvoScientistConfig @@ -721,7 +721,7 @@ class TestApplyModelLoadAgentFailureTransactional: ), ): cmd = ModelCommand() - _run(cmd._apply_model(ctx, "minimax-m2.7", "openrouter")) + await cmd._apply_model(ctx, "minimax-m2.7", "openrouter") # All four globals are unchanged — nothing was committed. assert mod._config is cfg @@ -753,7 +753,7 @@ class TestModelCommandOllamaPicker: ctx.ui = ui return ctx, cfg, ui - def test_picker_entries_include_detected_ollama_models(self): + async def test_picker_entries_include_detected_ollama_models(self): """When Ollama is reachable, detected models appear in entries with provider='ollama' and the Custom sentinel is appended.""" from EvoScientist.commands.implementation.model import ModelCommand @@ -773,7 +773,7 @@ class TestModelCommandOllamaPicker: side_effect=fake_discover, ), ): - _run(ModelCommand().execute(ctx, [])) + await ModelCommand().execute(ctx, []) entries = ui.wait_for_model_pick.call_args[0][0] ollama_rows = [(n, mid, p) for (n, mid, p) in entries if p == "ollama"] @@ -785,7 +785,7 @@ class TestModelCommandOllamaPicker: "ollama", ) in ollama_rows - def test_picker_entries_include_sentinel_when_discovery_empty(self): + async def test_picker_entries_include_sentinel_when_discovery_empty(self): """Daemon unreachable / no models pulled — sentinel is the user's escape hatch and must always be present.""" from EvoScientist.commands.implementation.model import ModelCommand @@ -805,7 +805,7 @@ class TestModelCommandOllamaPicker: side_effect=fake_discover, ), ): - _run(ModelCommand().execute(ctx, [])) + await ModelCommand().execute(ctx, []) entries = ui.wait_for_model_pick.call_args[0][0] ollama_rows = [(n, mid, p) for (n, mid, p) in entries if p == "ollama"] @@ -813,7 +813,7 @@ class TestModelCommandOllamaPicker: ("Custom Ollama model...", "__custom_ollama__", "ollama") ] - def test_picker_skips_ollama_section_when_not_configured(self): + async def test_picker_skips_ollama_section_when_not_configured(self): """ollama_base_url unset → no discovery call, no ollama entries, no sentinel (issue non-goal: no implicit localhost detection).""" from EvoScientist.commands.implementation.model import ModelCommand @@ -832,13 +832,13 @@ class TestModelCommandOllamaPicker: discovery, ), ): - _run(ModelCommand().execute(ctx, [])) + await ModelCommand().execute(ctx, []) discovery.assert_not_called() entries = ui.wait_for_model_pick.call_args[0][0] assert not any(p == "ollama" for (_, _, p) in entries) - def test_picker_handles_cfg_without_ollama_base_url_attr(self): + async def test_picker_handles_cfg_without_ollama_base_url_attr(self): """getattr(cfg, 'ollama_base_url', None) fallback: old configs (or SimpleNamespace test fixtures) may not carry the attribute at all. Must not raise AttributeError, must not probe.""" @@ -864,13 +864,13 @@ class TestModelCommandOllamaPicker: discovery, ), ): - _run(ModelCommand().execute(ctx, [])) + await ModelCommand().execute(ctx, []) discovery.assert_not_called() entries = ui.wait_for_model_pick.call_args[0][0] assert not any(p == "ollama" for (_, _, p) in entries) - def test_picker_sentinel_result_is_treated_as_cancel(self): + async def test_picker_sentinel_result_is_treated_as_cancel(self): """Defense-in-depth: if the widget ever returns the sentinel name itself (shouldn't happen — it should substitute the typed name), dispatch treats it as a cancel and does NOT call _apply_model.""" @@ -893,12 +893,12 @@ class TestModelCommandOllamaPicker: ), patch("EvoScientist.cli.agent._load_agent") as load_agent, ): - _run(ModelCommand().execute(ctx, [])) + await ModelCommand().execute(ctx, []) load_agent.assert_not_called() assert cfg.model == "claude-sonnet-4-6" # unchanged - def test_picker_applies_detected_ollama_model(self): + async def test_picker_applies_detected_ollama_model(self): """User picks a live-detected Ollama model → _apply_model is invoked with (name, "ollama") and the agent is rebuilt.""" from EvoScientist.commands.implementation.model import ModelCommand @@ -928,7 +928,7 @@ class TestModelCommandOllamaPicker: return_value=MagicMock(), ), ): - _run(ModelCommand().execute(ctx, [])) + await ModelCommand().execute(ctx, []) # Committed via set_active_config(temp_cfg); original cfg untouched. set_cfg.assert_called_once() diff --git a/tests/test_model_fallback.py b/tests/test_model_fallback.py index 19f8f31..fa9e346 100644 --- a/tests/test_model_fallback.py +++ b/tests/test_model_fallback.py @@ -20,7 +20,6 @@ from EvoScientist.middleware.model_fallback import ( clear_fallbacks, set_ui_emit, ) -from tests.conftest import run_async as _run # ── Helpers ────────────────────────────────────────────────────── @@ -146,7 +145,7 @@ class TestIsNonFallbackable: class TestTryFallbacks: """End-to-end tests for the fallback chain traversal.""" - def test_first_fallback_succeeds(self): + async def test_first_fallback_succeeds(self): """When the first fallback model works, return its response.""" add_fallback("fb-model", "fb-provider") req = _fake_request() @@ -154,13 +153,13 @@ class TestTryFallbacks: with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() - result = _run(_try_fallbacks(req, invoke, Exception("503 boom"))) + result = await _try_fallbacks(req, invoke, Exception("503 boom")) assert result is AI_RESPONSE invoke.assert_awaited_once() mock_gcm.assert_called_once_with(model="fb-model", provider="fb-provider") - def test_skips_failing_fallback_tries_next(self): + async def test_skips_failing_fallback_tries_next(self): """When the first fallback fails, try the second.""" add_fallback("fb-bad", "prov-a") add_fallback("fb-good", "prov-b") @@ -177,12 +176,12 @@ class TestTryFallbacks: with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() - result = _run(_try_fallbacks(req, _invoke, Exception("503 boom"))) + result = await _try_fallbacks(req, _invoke, Exception("503 boom")) assert result is AI_RESPONSE assert call_count == 2 - def test_all_fallbacks_exhausted_raises_last(self): + async def test_all_fallbacks_exhausted_raises_last(self): """When every fallback fails, re-raise the last exception.""" add_fallback("fb-a", "prov-a") add_fallback("fb-b", "prov-b") @@ -202,11 +201,11 @@ class TestTryFallbacks: with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() with pytest.raises(Exception, match="429 from fb-b") as exc_info: - _run(_try_fallbacks(req, _invoke, Exception("503 primary"))) + await _try_fallbacks(req, _invoke, Exception("503 primary")) assert exc_info.value is last_error - def test_non_fallbackable_in_chain_aborts_immediately(self): + async def test_non_fallbackable_in_chain_aborts_immediately(self): """A non-fallbackable error from a fallback model aborts the chain.""" add_fallback("fb-a", "prov-a") add_fallback("fb-b", "prov-b") # should never be reached @@ -218,7 +217,7 @@ class TestTryFallbacks: with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() with pytest.raises(Exception, match="context_length_exceeded"): - _run(_try_fallbacks(req, _invoke, Exception("503 primary"))) + await _try_fallbacks(req, _invoke, Exception("503 primary")) # get_chat_model should only have been called once (for fb-a), # fb-b should never be reached. @@ -233,43 +232,41 @@ class TestTryFallbacks: class TestGuardAndFallback: """Verify that non-fallbackable errors are re-raised before trying the chain.""" - def test_context_overflow_raises_immediately(self): + async def test_context_overflow_raises_immediately(self): add_fallback("fb", "prov") req = _fake_request() invoke = AsyncMock() with pytest.raises(ContextOverflowError): - _run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke)) + await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke) invoke.assert_not_awaited() - def test_malformed_400_raises_immediately(self): + async def test_malformed_400_raises_immediately(self): add_fallback("fb", "prov") req = _fake_request() invoke = AsyncMock() with pytest.raises(Exception, match="invalid_request_error"): - _run( - _guard_and_fallback( - Exception("400: invalid_request_error"), req, invoke - ) + await _guard_and_fallback( + Exception("400: invalid_request_error"), req, invoke ) invoke.assert_not_awaited() - def test_server_error_proceeds_to_fallback(self): + async def test_server_error_proceeds_to_fallback(self): add_fallback("fb", "prov") req = _fake_request() invoke = AsyncMock(return_value=AI_RESPONSE) with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() - result = _run(_guard_and_fallback(Exception("503 overloaded"), req, invoke)) + result = await _guard_and_fallback(Exception("503 overloaded"), req, invoke) assert result is AI_RESPONSE invoke.assert_awaited_once() - def test_auth_error_proceeds_to_fallback(self): + async def test_auth_error_proceeds_to_fallback(self): """Auth errors should try the fallback chain (different provider).""" add_fallback("fb", "other-prov") req = _fake_request() @@ -277,10 +274,8 @@ class TestGuardAndFallback: with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() - result = _run( - _guard_and_fallback( - Exception("400 Bad Request: invalid_api_key"), req, invoke - ) + result = await _guard_and_fallback( + Exception("400 Bad Request: invalid_api_key"), req, invoke ) assert result is AI_RESPONSE @@ -295,7 +290,7 @@ class TestGuardAndFallback: class TestUiEmit: """Verify that fallback events are surfaced via the registered callback.""" - def test_emit_captures_messages(self): + async def test_emit_captures_messages(self): add_fallback("fb", "prov") req = _fake_request() invoke = AsyncMock(return_value=AI_RESPONSE) @@ -305,14 +300,14 @@ class TestUiEmit: with patch("EvoScientist.llm.models.get_chat_model") as mock_gcm: mock_gcm.return_value = MagicMock() - _run(_try_fallbacks(req, invoke, Exception("503 down"))) + await _try_fallbacks(req, invoke, Exception("503 down")) texts = [t for t, _ in messages] assert any("Primary model failed" in t for t in texts) assert any("Falling back to fb (prov)" in t for t in texts) assert any("succeeded" in t for t in texts) - def test_emit_shows_non_fallbackable_rejection(self): + async def test_emit_shows_non_fallbackable_rejection(self): add_fallback("fb", "prov") req = _fake_request() invoke = AsyncMock() @@ -321,7 +316,7 @@ class TestUiEmit: set_ui_emit(lambda text, style: messages.append((text, style))) with pytest.raises(ContextOverflowError): - _run(_guard_and_fallback(ContextOverflowError("overflow"), req, invoke)) + await _guard_and_fallback(ContextOverflowError("overflow"), req, invoke) texts = [t for t, _ in messages] assert any("not eligible for fallback" in t for t in texts) diff --git a/tests/test_model_passthrough_patch.py b/tests/test_model_passthrough_patch.py index f177ea2..cafc36c 100644 --- a/tests/test_model_passthrough_patch.py +++ b/tests/test_model_passthrough_patch.py @@ -16,7 +16,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from EvoScientist.llm import patches as patches_mod -from tests.conftest import run_async as _run # ============================================================================= # Helpers @@ -152,7 +151,7 @@ class TestStartAsyncTaskInjection: "configurable": {"model": "gpt-5", "model_provider": "openai"} } - def test_async_start_injects_config(self, restore_model_passthrough_patch): + async def test_async_start_injects_config(self, restore_model_passthrough_patch): try: from deepagents.middleware import async_subagents as ds_mod except ImportError: @@ -176,12 +175,10 @@ class TestStartAsyncTaskInjection: "EvoScientist.EvoScientist._ensure_config", return_value=_stub_cfg(model="claude-haiku-4-5", provider="anthropic"), ): - _run( - tool.coroutine( - description="hi", - subagent_type="writing-agent", - runtime=_runtime_stub(), - ) + await tool.coroutine( + description="hi", + subagent_type="writing-agent", + runtime=_runtime_stub(), ) runs_async.create.assert_awaited_once() @@ -267,7 +264,7 @@ class TestUpdateAsyncTaskInjection: "last_updated_at": "2026-05-07T00:00:00Z", } - def test_async_update_injects_config(self, restore_model_passthrough_patch): + async def test_async_update_injects_config(self, restore_model_passthrough_patch): """The async coroutine path must inject config too.""" try: from deepagents.middleware import async_subagents as ds_mod @@ -296,12 +293,10 @@ class TestUpdateAsyncTaskInjection: "EvoScientist.EvoScientist._ensure_config", return_value=_stub_cfg(model="gpt-5", provider="openai"), ): - _run( - tool.coroutine( - task_id="thread-001", - message="follow up async", - runtime=runtime, - ) + await tool.coroutine( + task_id="thread-001", + message="follow up async", + runtime=runtime, ) runs_async.create.assert_awaited_once() diff --git a/tests/test_new_command.py b/tests/test_new_command.py index c3620f1..dc67671 100644 --- a/tests/test_new_command.py +++ b/tests/test_new_command.py @@ -2,11 +2,9 @@ from unittest.mock import AsyncMock, MagicMock -from tests.conftest import run_async as _run - class TestNewCommand: - def test_execute_calls_start_new_session(self): + async def test_execute_calls_start_new_session(self): from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.session import NewCommand @@ -18,7 +16,7 @@ class TestNewCommand: ui=ui, workspace_dir="/old/ws", ) - _run(NewCommand().execute(ctx, [])) + await NewCommand().execute(ctx, []) ui.start_new_session.assert_awaited_once() def test_requires_agent_false(self): @@ -26,7 +24,7 @@ class TestNewCommand: assert NewCommand().requires_agent is False - def test_no_agent_access(self): + async def test_no_agent_access(self): """Command body must not touch ctx.agent (it's still loading).""" from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.session import NewCommand @@ -35,4 +33,4 @@ class TestNewCommand: ui.start_new_session = AsyncMock() ctx = CommandContext(agent=None, thread_id="tid", ui=ui) # No AttributeError even though ctx.agent is None - _run(NewCommand().execute(ctx, [])) + await NewCommand().execute(ctx, []) diff --git a/tests/test_observation_memory.py b/tests/test_observation_memory.py index cebef62..5989034 100644 --- a/tests/test_observation_memory.py +++ b/tests/test_observation_memory.py @@ -1655,9 +1655,7 @@ def test_turn_compaction_uses_latest_user_turn_only(): ] -def test_lifecycle_schedules_turn_worker_without_awaiting( - tmp_path, monkeypatch, run_async -): +async def test_lifecycle_schedules_turn_worker_without_awaiting(tmp_path, monkeypatch): memory_dir = tmp_path / "memories" workspace_dir = tmp_path / "workspace" calls = [] @@ -1682,21 +1680,18 @@ def test_lifecycle_schedules_turn_worker_without_awaiting( ) runtime = _runtime("thread-1") - async def run(): - state: AgentState[object] = { - "messages": [ - HumanMessage("previous turn"), - AIMessage("previous answer"), - HumanMessage("hi"), - AIMessage("done"), - ] - } - await middleware.aafter_agent( - state, - runtime, - ) - - run_async(run()) + state: AgentState[object] = { + "messages": [ + HumanMessage("previous turn"), + AIMessage("previous answer"), + HumanMessage("hi"), + AIMessage("done"), + ] + } + await middleware.aafter_agent( + state, + runtime, + ) assert len(calls) == 1 request, hooks = calls[0] @@ -2182,10 +2177,9 @@ def test_observation_linker_does_not_launch_when_observations_disabled( launch_call.assert_not_called() -def test_async_observation_linker_does_not_launch_when_observations_disabled( +async def test_async_observation_linker_does_not_launch_when_observations_disabled( tmp_path, monkeypatch, - run_async, ): context = _linker_context( memory_dir=tmp_path / "memories", @@ -2200,7 +2194,7 @@ def test_async_observation_linker_does_not_launch_when_observations_disabled( launch_call = MagicMock() monkeypatch.setattr(memory_launch, "alaunch_background_run", launch_call) - run = run_async(memory_launch.alaunch_observation_linker(context)) + run = await memory_launch.alaunch_observation_linker(context) assert run is None launch_call.assert_not_called() @@ -2335,8 +2329,8 @@ def test_sync_memory_worker_watcher_untracks_without_counting_on_poll_abort( assert status.observations_recorded == 0 -def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort( - tmp_path, monkeypatch, run_async +async def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort( + tmp_path, monkeypatch ): memory_dir = tmp_path / "memories" _mark_worker_started(memory_dir) @@ -2348,14 +2342,12 @@ def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort( async def get(self, **_kwargs): raise RuntimeError("poll failed") - run_async( - background_runs.awatch_background_run( - SimpleNamespace(runs=_Runs()), - thread_id="worker-thread", - run_id="run-1", - hooks=memory_launch._memory_worker_launch_hooks(memory_dir), - watcher_config=_fast_watcher_config(max_poll_failures=1), - ) + await background_runs.awatch_background_run( + SimpleNamespace(runs=_Runs()), + thread_id="worker-thread", + run_id="run-1", + hooks=memory_launch._memory_worker_launch_hooks(memory_dir), + watcher_config=_fast_watcher_config(max_poll_failures=1), ) status = worker_activity.memory_worker_status() assert status.is_running is False @@ -2363,8 +2355,8 @@ def test_async_memory_worker_watcher_untracks_without_counting_on_poll_abort( assert status.observations_recorded == 0 -def test_async_memory_worker_watcher_counts_completion_under_blockbuster( - tmp_path, run_async +async def test_async_memory_worker_watcher_counts_completion_under_blockbuster( + tmp_path, ): memory_dir = tmp_path / "memories" _mark_worker_started(memory_dir) @@ -2376,20 +2368,17 @@ def test_async_memory_worker_watcher_counts_completion_under_blockbuster( async def get(self, **_kwargs): return {"status": "success"} - async def run(): - blocker = BlockBuster(scanned_modules=[memory_worker, worker_activity]) - blocker.activate() - try: - await background_runs.awatch_background_run( - SimpleNamespace(runs=_Runs()), - thread_id="worker-thread", - run_id="run-1", - hooks=memory_launch._memory_worker_launch_hooks(memory_dir), - ) - finally: - blocker.deactivate() - - run_async(run()) + blocker = BlockBuster(scanned_modules=[memory_worker, worker_activity]) + blocker.activate() + try: + await background_runs.awatch_background_run( + SimpleNamespace(runs=_Runs()), + thread_id="worker-thread", + run_id="run-1", + hooks=memory_launch._memory_worker_launch_hooks(memory_dir), + ) + finally: + blocker.deactivate() status = worker_activity.memory_worker_status() assert status.is_running is False assert status.profile_updates == 1 @@ -2527,7 +2516,7 @@ def test_memory_worker_marks_active_status(tmp_path, monkeypatch): assert status.observations_recorded == 1 -def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch, run_async): +async def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch): monkeypatch.setattr( background_runs, "default_background_run_url", lambda: "http://x" ) @@ -2569,22 +2558,18 @@ def test_async_memory_worker_offloads_blocking_work(tmp_path, monkeypatch, run_a spawned: list[background_runs.BackgroundRun] = [] - async def run(): - event_loop_thread = threading.get_ident() - context = _memory_source_context( - memory_dir=tmp_path / "memories", - workspace_dir=tmp_path / "workspace", - trajectory=[{"role": "human", "content": "hi"}], - ) - request = memory_launch.memory_worker_launch_request(context) - await background_runs.alaunch_background_run( - request, - hooks=memory_launch._memory_worker_launch_hooks(tmp_path / "memories"), - spawn_status_watcher=spawned.append, - ) - return event_loop_thread - - event_loop_thread = run_async(run()) + event_loop_thread = threading.get_ident() + context = _memory_source_context( + memory_dir=tmp_path / "memories", + workspace_dir=tmp_path / "workspace", + trajectory=[{"role": "human", "content": "hi"}], + ) + request = memory_launch.memory_worker_launch_request(context) + await background_runs.alaunch_background_run( + request, + hooks=memory_launch._memory_worker_launch_hooks(tmp_path / "memories"), + spawn_status_watcher=spawned.append, + ) assert [name for name, _thread_id in call_threads] == ["health", "snapshot"] assert all(thread_id != event_loop_thread for _name, thread_id in call_threads) assert worker_activity.memory_worker_status().is_running is True diff --git a/tests/test_ollama_discovery.py b/tests/test_ollama_discovery.py index 1fb785f..083f576 100644 --- a/tests/test_ollama_discovery.py +++ b/tests/test_ollama_discovery.py @@ -17,7 +17,6 @@ from EvoScientist.llm.ollama_discovery import ( discover_ollama_models, validate_ollama_connection, ) -from tests.conftest import run_async as _run class TestValidateOllamaConnection: @@ -71,17 +70,17 @@ class TestValidateOllamaConnection: class TestDiscoverOllamaModels: """Async probe — contract: never raise, return list[str].""" - def test_empty_base_url_returns_empty_without_http(self): + async def test_empty_base_url_returns_empty_without_http(self): # No HTTP call should be made for an empty base_url — verified by # the fact that no mock is set up and the test completes. - names = _run(discover_ollama_models("")) + names = await discover_ollama_models("") assert names == [] - def test_none_base_url_returns_empty(self): - names = _run(discover_ollama_models(None)) + async def test_none_base_url_returns_empty(self): + names = await discover_ollama_models(None) assert names == [] - def test_200_returns_names(self): + async def test_200_returns_names(self): async def fake_get(self, url): resp = MagicMock() resp.status_code = 200 @@ -93,10 +92,10 @@ class TestDiscoverOllamaModels: return resp with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == ["llama3.3:latest", "qwen3:8b"] - def test_strips_entries_without_name(self): + async def test_strips_entries_without_name(self): async def fake_get(self, url): resp = MagicMock() resp.status_code = 200 @@ -112,36 +111,36 @@ class TestDiscoverOllamaModels: return resp with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == ["llama3.3"] - def test_timeout_returns_empty(self): + async def test_timeout_returns_empty(self): async def fake_get(self, url): raise httpx.TimeoutException("timed out") with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == [] - def test_connect_error_returns_empty(self): + async def test_connect_error_returns_empty(self): async def fake_get(self, url): raise httpx.ConnectError("refused") with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == [] - def test_non_200_returns_empty(self): + async def test_non_200_returns_empty(self): async def fake_get(self, url): resp = MagicMock() resp.status_code = 500 return resp with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == [] - def test_malformed_json_returns_empty(self): + async def test_malformed_json_returns_empty(self): async def fake_get(self, url): resp = MagicMock() resp.status_code = 200 @@ -149,10 +148,10 @@ class TestDiscoverOllamaModels: return resp with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == [] - def test_missing_models_key_returns_empty(self): + async def test_missing_models_key_returns_empty(self): async def fake_get(self, url): resp = MagicMock() resp.status_code = 200 @@ -160,10 +159,10 @@ class TestDiscoverOllamaModels: return resp with patch.object(httpx.AsyncClient, "get", fake_get): - names = _run(discover_ollama_models("http://localhost:11434")) + names = await discover_ollama_models("http://localhost:11434") assert names == [] - def test_trailing_slash_stripped_from_url(self): + async def test_trailing_slash_stripped_from_url(self): called = {} async def fake_get(self, url): @@ -174,7 +173,7 @@ class TestDiscoverOllamaModels: return resp with patch.object(httpx.AsyncClient, "get", fake_get): - _run(discover_ollama_models("http://localhost:11434/")) + await discover_ollama_models("http://localhost:11434/") assert called["url"] == "http://localhost:11434/api/tags" diff --git a/tests/test_onboard.py b/tests/test_onboard.py index 5762e78..2d7c233 100644 --- a/tests/test_onboard.py +++ b/tests/test_onboard.py @@ -184,6 +184,38 @@ class TestSharedConstantsAlignment: ) +class TestOAuthModeReconcile: + def test_reconcile_preserves_auxiliary_openai_oauth(self): + from EvoScientist.config.onboard.wizard import _reconcile_oauth_modes + + config = EvoScientistConfig( + provider="minimax", + auxiliary_provider="openai", + auxiliary_model="gpt-5.5", + openai_auth_mode="oauth", + anthropic_auth_mode="oauth", + ) + + _reconcile_oauth_modes(config) + + assert config.openai_auth_mode == "oauth" + assert config.anthropic_auth_mode == "api_key" + + def test_reconcile_preserves_auxiliary_provider_without_model(self): + from EvoScientist.config.onboard.wizard import _reconcile_oauth_modes + + config = EvoScientistConfig( + provider="minimax", + auxiliary_provider="openai", + auxiliary_model="", + openai_auth_mode="oauth", + ) + + _reconcile_oauth_modes(config) + + assert config.openai_auth_mode == "oauth" + + # ============================================================================= # Test render_progress # ============================================================================= @@ -386,6 +418,109 @@ class TestStepProvider: _step_provider(config) +class TestStepOAuthAuthMode: + @pytest.mark.parametrize( + ( + "step_name", + "config_attr", + "provider_label", + "oauth_choice_label", + "ccproxy_provider", + "status_label", + "question_label", + "login_prompt", + ), + [ + ( + "_step_anthropic_auth_mode", + "anthropic_auth_mode", + "Anthropic", + "Claude Code OAuth", + "claude_api", + "OAuth", + "Authentication mode", + "Log in to Claude now?", + ), + ( + "_step_openai_auth_mode", + "openai_auth_mode", + "OpenAI", + "Codex OAuth", + "codex", + "Codex OAuth", + "OpenAI authentication mode", + "Log in to Codex now?", + ), + ], + ) + def test_oauth_wrappers_use_provider_specific_ccproxy_flow( + self, + step_name, + config_attr, + provider_label, + oauth_choice_label, + ccproxy_provider, + status_label, + question_label, + login_prompt, + ): + """Anthropic/OpenAI wrappers share flow but keep provider-specific IDs.""" + from EvoScientist.config.onboard import steps as onboard_steps + + config = EvoScientistConfig(**{config_attr: "oauth"}) + select_question = MagicMock() + select_question.ask.return_value = "oauth" + confirm_question = MagicMock() + confirm_question.ask.return_value = True + + with ( + patch( + "EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True + ), + patch( + "EvoScientist.ccproxy_manager.check_ccproxy_auth", + return_value=(False, "not authenticated"), + ) as mock_check_auth, + patch( + "EvoScientist.config.onboard.prompter.install_navigation_keys" + ) as mock_nav, + patch( + "EvoScientist.config.onboard.steps.questionary.select", + return_value=select_question, + ) as mock_select, + patch( + "EvoScientist.config.onboard.steps.questionary.confirm", + return_value=confirm_question, + ) as mock_confirm, + patch( + "EvoScientist.config.onboard.steps._prompt_ccproxy_port" + ) as mock_port, + patch("EvoScientist.config.onboard.steps._run_ccproxy_login") as mock_login, + ): + result = getattr(onboard_steps, step_name)(config) + + assert result == "oauth" + mock_nav.assert_called_once_with(select_question, with_back=True) + mock_port.assert_called_once_with(config) + mock_check_auth.assert_called_once_with(ccproxy_provider) + mock_login.assert_called_once_with(ccproxy_provider, status_label) + + select_call = mock_select.call_args + assert select_call.args[0] == f"{question_label} [Esc/← to go back]:" + assert select_call.kwargs["default"] == "oauth" + choice_titles = [ + choice.title + for choice in select_call.kwargs["choices"] + if getattr(choice, "value", None) in {"api_key", "oauth"} + ] + assert choice_titles == [ + f"API Key (direct {provider_label} access)", + f"{oauth_choice_label} (via ccproxy — no API key needed)", + ] + mock_confirm.assert_called_once() + assert mock_confirm.call_args.args[0] == login_prompt + + class TestStepModel: def test_returns_selected_model(self): """Test that _step_model returns selected model.""" @@ -1266,6 +1401,7 @@ class TestRunOnboard: "claude-sonnet-4-6", # Model "assemble", # Auxiliary: Assemble "openai", # Auxiliary provider (a different company) + "api_key", # Auxiliary OpenAI auth mode "gpt-5.5", # Auxiliary model "daemon", # Workspace mode True, # Show thinking @@ -1289,12 +1425,185 @@ class TestRunOnboard: final_config = mock_save.call_args_list[-1].args[0] assert final_config.auxiliary_provider == "openai" assert final_config.auxiliary_model == "gpt-5.5" + assert final_config.openai_auth_mode == "api_key" # The auxiliary provider's key is stored in its per-provider field. assert final_config.openai_api_key == "sk-aux-openai" # Main agent is untouched. assert final_config.provider == "anthropic" assert final_config.model == "claude-sonnet-4-6" + def test_auxiliary_same_provider_reuses_main_credentials(self): + """Same-provider co-pilot should not imply separate credentials exist.""" + from EvoScientist.config.onboard.wizard import run_onboard + + mock_q = MagicMock() + with ( + _patch_all_questionary(mock_q), + patch("EvoScientist.config.onboard.wizard.load_config") as mock_load, + patch("EvoScientist.config.onboard.wizard.save_config") as mock_save, + patch("EvoScientist.config.onboard.wizard.console"), + patch("EvoScientist.config.onboard.steps.console"), + ): + mock_load.return_value = EvoScientistConfig( + provider="openai", + model="gpt-5.5", + openai_api_key="sk-main-openai", + ) + mock_q.select.return_value.ask.side_effect = [ + "assemble", # Auxiliary: Assemble + "openai", # Same provider as the main model + "gpt-5.5", # Auxiliary model + ] + mock_q.confirm.return_value.ask.side_effect = [True] # Save config + + result = run_onboard( + skip_validation=True, only_sections={"auxiliary_model"} + ) + + assert result is True + final_config = mock_save.call_args_list[-1].args[0] + assert final_config.auxiliary_provider == "openai" + assert final_config.auxiliary_model == "gpt-5.5" + assert final_config.openai_api_key == "sk-main-openai" + mock_q.password.assert_not_called() + assert mock_q.select.return_value.ask.call_count == 3 + + def test_auxiliary_same_provider_prompts_when_shared_key_missing(self): + """Same-provider reuse should not hide a missing shared API key.""" + from EvoScientist.config.onboard.wizard import run_onboard + + mock_q = MagicMock() + with ( + _patch_all_questionary(mock_q), + patch("EvoScientist.config.onboard.wizard.load_config") as mock_load, + patch("EvoScientist.config.onboard.wizard.save_config") as mock_save, + patch("EvoScientist.config.onboard.wizard.console"), + patch("EvoScientist.config.onboard.steps.console"), + patch("EvoScientist.config.onboard.helpers.console"), + patch( + "EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True + ), + ): + mock_load.return_value = EvoScientistConfig( + provider="openai", + model="gpt-5.5", + openai_auth_mode="api_key", + openai_api_key="", + ) + mock_q.select.return_value.ask.side_effect = [ + "assemble", # Auxiliary: Assemble + "openai", # Same provider as the main model + "api_key", # Shared OpenAI auth mode + "gpt-5.5", # Auxiliary model + ] + mock_q.password.return_value.ask.side_effect = [ + "sk-shared-openai", + ] + mock_q.confirm.return_value.ask.side_effect = [True] # Save config + + result = run_onboard( + skip_validation=True, only_sections={"auxiliary_model"} + ) + + assert result is True + final_config = mock_save.call_args_list[-1].args[0] + assert final_config.auxiliary_provider == "openai" + assert final_config.auxiliary_model == "gpt-5.5" + assert final_config.openai_auth_mode == "api_key" + assert final_config.openai_api_key == "sk-shared-openai" + mock_q.password.assert_called_once() + assert mock_q.select.return_value.ask.call_count == 4 + + def test_auxiliary_openai_oauth_skips_api_key(self): + """Auxiliary OpenAI now uses the shared auth flow and skips keys on OAuth.""" + from EvoScientist.config.onboard.wizard import run_onboard + + mock_q = MagicMock() + with ( + _patch_all_questionary(mock_q), + patch("EvoScientist.config.onboard.wizard.load_config") as mock_load, + patch("EvoScientist.config.onboard.wizard.save_config") as mock_save, + patch("EvoScientist.config.onboard.wizard.console"), + patch("EvoScientist.config.onboard.steps.console"), + patch("EvoScientist.config.onboard.helpers.console"), + patch( + "EvoScientist.ccproxy_manager.is_ccproxy_available", return_value=True + ), + patch( + "EvoScientist.ccproxy_manager.check_ccproxy_auth", + return_value=(False, "not authenticated"), + ) as mock_auth, + ): + mock_load.return_value = EvoScientistConfig() + mock_q.select.return_value.ask.side_effect = [ + "assemble", # Auxiliary: Assemble + "openai", # Auxiliary provider + "oauth", # OpenAI auth mode + "gpt-5.5", # Auxiliary model + ] + mock_q.text.return_value.ask.side_effect = [ + "", # ccproxy port (keep default) + ] + mock_q.confirm.return_value.ask.side_effect = [ + False, # Do not log in to Codex now + True, # Save config + ] + + result = run_onboard( + skip_validation=True, only_sections={"auxiliary_model"} + ) + + assert result is True + final_config = mock_save.call_args_list[-1].args[0] + assert final_config.auxiliary_provider == "openai" + assert final_config.auxiliary_model == "gpt-5.5" + assert final_config.openai_auth_mode == "oauth" + assert final_config.openai_api_key == "" + mock_q.password.assert_not_called() + mock_auth.assert_called_once_with("codex") + + def test_auxiliary_reconfigure_clears_unused_openai_oauth(self): + """Switching co-pilot away from OpenAI clears stale OpenAI OAuth mode.""" + from EvoScientist.config.onboard.wizard import run_onboard + + mock_q = MagicMock() + with ( + _patch_all_questionary(mock_q), + patch("EvoScientist.config.onboard.wizard.load_config") as mock_load, + patch("EvoScientist.config.onboard.wizard.save_config") as mock_save, + patch("EvoScientist.config.onboard.wizard.console"), + patch("EvoScientist.config.onboard.steps.console"), + patch("EvoScientist.config.onboard.helpers.console"), + ): + mock_load.return_value = EvoScientistConfig( + provider="anthropic", + model="claude-sonnet-4-6", + anthropic_auth_mode="oauth", + auxiliary_provider="openai", + auxiliary_model="gpt-5.5", + openai_auth_mode="oauth", + ) + mock_q.select.return_value.ask.side_effect = [ + "assemble", # Auxiliary: Assemble + "minimax", # Auxiliary provider no longer uses OpenAI + "global", # MiniMax region + "minimax-m2", # Auxiliary model + ] + mock_q.password.return_value.ask.side_effect = [ + "sk-minimax", # MiniMax API key + ] + mock_q.confirm.return_value.ask.side_effect = [True] # Save config + + result = run_onboard( + skip_validation=True, only_sections={"auxiliary_model"} + ) + + assert result is True + final_config = mock_save.call_args_list[-1].args[0] + assert final_config.auxiliary_provider == "minimax" + assert final_config.openai_auth_mode == "api_key" + assert final_config.anthropic_auth_mode == "oauth" + def test_auxiliary_custom_provider_collects_base_url(self): """Regression for the custom-provider fix: a custom auxiliary provider collects its base URL (provider -> base URL -> key -> model order).""" diff --git a/tests/test_pick_skills_interactive.py b/tests/test_pick_skills_interactive.py index 66b658b..2580bb1 100644 --- a/tests/test_pick_skills_interactive.py +++ b/tests/test_pick_skills_interactive.py @@ -79,14 +79,13 @@ class TestPickSkillsInteractive: class TestInstallSkillsHandlesEmpty: """InstallSkills.execute must distinguish None vs [] from the picker.""" - def test_empty_list_suppresses_cancel_message(self): + async def test_empty_list_suppresses_cancel_message(self): """When picker returns [], user should NOT see "Browse cancelled" (the picker already printed its own message).""" from unittest.mock import AsyncMock from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.skills import InstallSkills - from tests.conftest import run_async as _run ui = MagicMock() ui.supports_interactive = True @@ -97,18 +96,17 @@ class TestInstallSkillsHandlesEmpty: "EvoScientist.tools.skills_manager.fetch_remote_skill_index", return_value=_INDEX, ): - _run(InstallSkills().execute(ctx, [])) + await InstallSkills().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert not any("Browse cancelled" in m for m in msgs) - def test_none_shows_cancel_message(self): + async def test_none_shows_cancel_message(self): """When picker returns None (actual cancel), user sees the message.""" from unittest.mock import AsyncMock from EvoScientist.commands.base import CommandContext from EvoScientist.commands.implementation.skills import InstallSkills - from tests.conftest import run_async as _run ui = MagicMock() ui.supports_interactive = True @@ -119,7 +117,7 @@ class TestInstallSkillsHandlesEmpty: "EvoScientist.tools.skills_manager.fetch_remote_skill_index", return_value=_INDEX, ): - _run(InstallSkills().execute(ctx, [])) + await InstallSkills().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Browse cancelled" in m for m in msgs) diff --git a/tests/test_profile_memory_middleware.py b/tests/test_profile_memory_middleware.py index 6101424..a5b3c5d 100644 --- a/tests/test_profile_memory_middleware.py +++ b/tests/test_profile_memory_middleware.py @@ -344,9 +344,7 @@ def test_profile_memory_uses_path_pointers_when_profiles_exceed_budget( ) -def test_profile_memory_async_path_bootstraps_and_injects( - tmp_path, monkeypatch, run_async -): +async def test_profile_memory_async_path_bootstraps_and_injects(tmp_path, monkeypatch): memories = tmp_path / "memories" workspace = tmp_path / "workspace" workspace.mkdir() @@ -356,7 +354,7 @@ def test_profile_memory_async_path_bootstraps_and_injects( return request middleware = memory_module.create_memory_middleware(str(memories)) - run_async(middleware.awrap_model_call(_request(), _handler)) + await middleware.awrap_model_call(_request(), _handler) assert (memories / "profile" / "USER_PROFILE.md").exists() @@ -399,8 +397,8 @@ def test_profile_memory_read_failure_uses_path_pointers_without_overwriting( assert soul_path.read_bytes() == original_bytes -def test_profile_memory_async_path_inlines_content_under_blockbuster( - tmp_path, monkeypatch, run_async +async def test_profile_memory_async_path_inlines_content_under_blockbuster( + tmp_path, monkeypatch ): memories = tmp_path / "memories" workspace = tmp_path / "workspace" @@ -425,17 +423,13 @@ def test_profile_memory_async_path_inlines_content_under_blockbuster( monkeypatch.setattr(middleware, "_read_profile_memory", tracked_read_profile_memory) - async def run(): - event_loop_thread = threading.get_ident() - blocker = BlockBuster(scanned_modules=memory_module) - blocker.activate() - try: - modified = await middleware.amodify_request(_request()) - finally: - blocker.deactivate() - return event_loop_thread, modified - - event_loop_thread, modified = run_async(run()) + event_loop_thread = threading.get_ident() + blocker = BlockBuster(scanned_modules=memory_module) + blocker.activate() + try: + modified = await middleware.amodify_request(_request()) + finally: + blocker.deactivate() assert call_threads assert all(thread_id != event_loop_thread for thread_id in call_threads) @@ -534,8 +528,8 @@ def test_profile_memory_uses_explicit_workspace_for_project_profile( ).exists() -def test_profile_memory_resolves_project_id_once_per_middleware( - tmp_path, monkeypatch, run_async +async def test_profile_memory_resolves_project_id_once_per_middleware( + tmp_path, monkeypatch ): memories = tmp_path / "memories" workspace = tmp_path / "workspace" @@ -552,7 +546,7 @@ def test_profile_memory_resolves_project_id_once_per_middleware( str(memories), workspace_dir=str(workspace), max_inline_profile_chars=10 ) middleware.modify_request(_request()) - run_async(middleware.amodify_request(_request())) + await middleware.amodify_request(_request()) assert calls == [workspace] assert middleware.project_id == "P-cached-project" diff --git a/tests/test_qq_channel.py b/tests/test_qq_channel.py index 5c9fa6e..8960fd6 100644 --- a/tests/test_qq_channel.py +++ b/tests/test_qq_channel.py @@ -8,7 +8,6 @@ from EvoScientist.channels.qq.channel import ( QQConfig, _build_qq_keyboard, ) -from tests.conftest import run_async as _run class TestQQChannelSend: @@ -22,7 +21,7 @@ class TestQQChannelSend: channel._client.api.post_group_message = AsyncMock() return channel - def test_send_prefers_native_markdown_for_c2c(self): + async def test_send_prefers_native_markdown_for_c2c(self): channel = self._make_ready_channel() msg = OutboundMessage( channel="qq", @@ -35,7 +34,7 @@ class TestQQChannelSend: }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True channel._client.api.post_c2c_message.assert_awaited_once() sent = channel._client.api.post_c2c_message.await_args.kwargs @@ -46,7 +45,7 @@ class TestQQChannelSend: assert sent["msg_seq"] == 1 assert "content" not in sent - def test_send_falls_back_to_plain_text_when_markdown_send_fails(self): + async def test_send_falls_back_to_plain_text_when_markdown_send_fails(self): channel = self._make_ready_channel() channel._trace_event = MagicMock(side_effect=RuntimeError("trace failed")) channel._client.api.post_c2c_message = AsyncMock( @@ -63,7 +62,7 @@ class TestQQChannelSend: }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True assert channel._client.api.post_c2c_message.await_count == 2 first = channel._client.api.post_c2c_message.await_args_list[0].kwargs @@ -80,7 +79,7 @@ class TestQQChannelSend: # trigger "duplicate msg_seq". assert second["msg_seq"] == 2 - def test_send_does_not_fallback_on_transport_error(self): + async def test_send_does_not_fallback_on_transport_error(self): channel = self._make_ready_channel() async def _send_once(coro_factory, max_retries=3): @@ -101,13 +100,13 @@ class TestQQChannelSend: }, ) - assert _run(channel.send(msg)) is False + assert await channel.send(msg) is False channel._client.api.post_c2c_message.assert_awaited_once() sent = channel._client.api.post_c2c_message.await_args.kwargs assert sent["msg_type"] == 2 assert "content" not in sent - def test_send_does_not_fallback_when_transport_error_mentions_markdown(self): + async def test_send_does_not_fallback_when_transport_error_mentions_markdown(self): """A transport-layer error whose message incidentally contains the word "markdown" must NOT be reclassified as a markdown compatibility failure, otherwise genuine send failures get silently swallowed as plain-text.""" @@ -133,10 +132,10 @@ class TestQQChannelSend: }, ) - assert _run(channel.send(msg)) is False + assert await channel.send(msg) is False channel._client.api.post_c2c_message.assert_awaited_once() - def test_send_falls_back_on_qq_server_error_code(self): + async def test_send_falls_back_on_qq_server_error_code(self): """QQ server-side markdown errors (e.g. 304014 template not configured) should trigger plain-text fallback with a fresh msg_seq.""" channel = self._make_ready_channel() @@ -159,7 +158,7 @@ class TestQQChannelSend: }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True assert channel._client.api.post_c2c_message.await_count == 2 first = channel._client.api.post_c2c_message.await_args_list[0].kwargs @@ -231,7 +230,7 @@ class TestQQSendWithButtons: channel._client.api.post_group_message = AsyncMock() return channel - def test_c2c_send_attaches_keyboard(self): + async def test_c2c_send_attaches_keyboard(self): channel = self._make_channel() msg = OutboundMessage( channel="qq", @@ -247,7 +246,7 @@ class TestQQSendWithButtons: ], }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True sent = channel._client.api.post_c2c_message.await_args.kwargs assert sent["msg_type"] == 2 @@ -256,7 +255,7 @@ class TestQQSendWithButtons: assert rows[0]["buttons"][0]["action"]["data"] == "1" assert rows[1]["buttons"][0]["action"]["data"] == "2" - def test_group_send_does_not_attach_keyboard(self): + async def test_group_send_does_not_attach_keyboard(self): """Group keyboards are out of scope — silently dropped.""" channel = self._make_channel() msg = OutboundMessage( @@ -270,11 +269,11 @@ class TestQQSendWithButtons: "buttons": [{"text": "Approve", "value": "1"}], }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True sent = channel._client.api.post_group_message.await_args.kwargs assert "keyboard" not in sent - def test_fallback_appends_button_hint_when_keyboard_present(self): + async def test_fallback_appends_button_hint_when_keyboard_present(self): """If markdown send fails and we fall back to plain text, the keyboard is lost — append a textual hint so the user still has a way to reply (the values still pass `_parse_approval_reply`). @@ -302,7 +301,7 @@ class TestQQSendWithButtons: ], }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True plain_call = channel._client.api.post_c2c_message.await_args_list[1].kwargs assert plain_call["msg_type"] == 0 @@ -311,7 +310,7 @@ class TestQQSendWithButtons: assert "1=Approve" in plain_call["content"] assert "2=Reject" in plain_call["content"] - def test_fallback_hint_handles_non_string_button_value(self): + async def test_fallback_hint_handles_non_string_button_value(self): """Regression: integer/None button values must not crash the plain-text fallback (the keyboard builder already coerces them).""" channel = self._make_channel() @@ -335,7 +334,7 @@ class TestQQSendWithButtons: ], }, ) - assert _run(channel.send(msg)) is True + assert await channel.send(msg) is True plain_call = channel._client.api.post_c2c_message.await_args_list[1].kwargs assert "42=OK" in plain_call["content"] assert "Cancel=Cancel" in plain_call["content"] @@ -392,9 +391,9 @@ class TestQQInteractionCallback: ) return interaction - def test_click_publishes_to_bus_with_button_data(self): + async def test_click_publishes_to_bus_with_button_data(self): channel = self._make_channel() - _run(channel._on_interaction(self._make_interaction("1"))) + await channel._on_interaction(self._make_interaction("1")) channel._bus.publish_inbound.assert_awaited_once() inbound = channel._bus.publish_inbound.await_args[0][0] @@ -405,63 +404,61 @@ class TestQQInteractionCallback: assert inbound.metadata["button_value"] == "1" assert inbound.metadata["msg_type"] == "c2c" - def test_click_acks_interaction(self): + async def test_click_acks_interaction(self): channel = self._make_channel() - _run(channel._on_interaction(self._make_interaction("1"))) + await channel._on_interaction(self._make_interaction("1")) channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0) - def test_click_bypasses_debounce(self): + async def test_click_bypasses_debounce(self): """Click never hits queue_message (debounce buffer).""" channel = self._make_channel() channel.queue_message = AsyncMock() - _run(channel._on_interaction(self._make_interaction("3"))) + await channel._on_interaction(self._make_interaction("3")) channel.queue_message.assert_not_called() channel._bus.publish_inbound.assert_awaited_once() - def test_group_interaction_ignored(self): + async def test_group_interaction_ignored(self): """No user_openid → group/guild click → don't publish.""" channel = self._make_channel() intr = self._make_interaction(user_openid="") intr.group_openid = "group_xxx" - _run(channel._on_interaction(intr)) + await channel._on_interaction(intr) channel._bus.publish_inbound.assert_not_called() # ACK still fires — it runs first, before the group-skip return. channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0) - def test_click_dropped_when_middleware_rejects(self): + async def test_click_dropped_when_middleware_rejects(self): channel = self._make_channel() channel._build_inbound_async = AsyncMock(return_value=None) - _run(channel._on_interaction(self._make_interaction("1"))) + await channel._on_interaction(self._make_interaction("1")) channel._bus.publish_inbound.assert_not_called() # ACK still fires (we don't want the user staring at a stuck button) channel._client.api.on_interaction_result.assert_awaited_once() - def test_empty_button_data_falls_back_to_button_id(self): + async def test_empty_button_data_falls_back_to_button_id(self): channel = self._make_channel() - _run( - channel._on_interaction( - self._make_interaction(button_data="", button_id="btn_3") - ) + await channel._on_interaction( + self._make_interaction(button_data="", button_id="btn_3") ) inbound = channel._bus.publish_inbound.await_args[0][0] assert inbound.content == "btn_3" - def test_ack_fires_even_when_handler_throws(self): + async def test_ack_fires_even_when_handler_throws(self): """ACK must run before downstream processing so the QQ button UI stays responsive even if middleware/bus crashes.""" channel = self._make_channel() channel._build_inbound_async = AsyncMock(side_effect=RuntimeError("boom")) # Should not raise — handler swallows downstream errors. - _run(channel._on_interaction(self._make_interaction("1"))) + await channel._on_interaction(self._make_interaction("1")) channel._client.api.on_interaction_result.assert_awaited_once_with("intr_1", 0) - def test_button_value_metadata_is_string_coerced(self): + async def test_button_value_metadata_is_string_coerced(self): """Regression: metadata['button_value'] must be a string (was raw).""" channel = self._make_channel() resolved = MagicMock(button_id="btn_0", button_data=42, message_id="msg_orig") data = MagicMock(type=None, resolved=resolved) intr = MagicMock(id="intr_1", user_openid="u_x", group_openid=None, data=data) - _run(channel._on_interaction(intr)) + await channel._on_interaction(intr) inbound = channel._bus.publish_inbound.await_args[0][0] assert inbound.content == "42" assert inbound.metadata["button_value"] == "42" diff --git a/tests/test_resume_command.py b/tests/test_resume_command.py index 6dde6df..b777ee2 100644 --- a/tests/test_resume_command.py +++ b/tests/test_resume_command.py @@ -2,7 +2,6 @@ from unittest.mock import AsyncMock, MagicMock -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -24,7 +23,7 @@ def _ctx(thread_id="current", workspace_dir="/ws", thread_store=None): class TestResumeCommand: - def test_with_arg_resolves_and_calls_ui(self): + async def test_with_arg_resolves_and_calls_ui(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx( @@ -33,23 +32,23 @@ class TestResumeCommand: metadata={"workspace_dir": "/restored"}, ) ) - _run(ResumeCommand().execute(ctx, ["target-tid"])) + await ResumeCommand().execute(ctx, ["target-tid"]) ui.handle_session_resume.assert_awaited_once_with("target-tid", "/restored") # ctx mutations assert ctx.thread_id == "target-tid" assert ctx.workspace_dir == "/restored" - def test_no_arg_empty_threads_prints_message(self): + async def test_no_arg_empty_threads_prints_message(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx() - _run(ResumeCommand().execute(ctx, [])) + await ResumeCommand().execute(ctx, []) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("No sessions to resume" in m for m in msgs) ui.wait_for_thread_pick.assert_not_called() ui.handle_session_resume.assert_not_called() - def test_no_arg_calls_picker(self): + async def test_no_arg_calls_picker(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx() @@ -60,11 +59,11 @@ class TestResumeCommand: resolved_thread_id="picked-tid", ) ctx.graph_gateway = FakeGraphGateway(thread_store=store) - _run(ResumeCommand().execute(ctx, [])) + await ResumeCommand().execute(ctx, []) ui.wait_for_thread_pick.assert_awaited_once() ui.handle_session_resume.assert_awaited_once() - def test_picker_cancel_returns(self): + async def test_picker_cancel_returns(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx() @@ -72,28 +71,28 @@ class TestResumeCommand: threads = [{"thread_id": "t1", "preview": "", "message_count": 0}] store = FakeThreadStore(threads=threads) ctx.graph_gateway = FakeGraphGateway(thread_store=store) - _run(ResumeCommand().execute(ctx, [])) + await ResumeCommand().execute(ctx, []) ui.handle_session_resume.assert_not_called() - def test_ambiguous_prefix(self): + async def test_ambiguous_prefix(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx(thread_store=FakeThreadStore(matches=["abc-one", "abc-two"])) - _run(ResumeCommand().execute(ctx, ["abc"])) + await ResumeCommand().execute(ctx, ["abc"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Ambiguous" in m for m in msgs) ui.handle_session_resume.assert_not_called() - def test_not_found(self): + async def test_not_found(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx() - _run(ResumeCommand().execute(ctx, ["missing"])) + await ResumeCommand().execute(ctx, ["missing"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("not found" in m for m in msgs) ui.handle_session_resume.assert_not_called() - def test_prefix_resolves_to_unique_match(self): + async def test_prefix_resolves_to_unique_match(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx( @@ -102,18 +101,18 @@ class TestResumeCommand: metadata={"workspace_dir": "/ws1"}, ) ) - _run(ResumeCommand().execute(ctx, ["abc"])) + await ResumeCommand().execute(ctx, ["abc"]) ui.handle_session_resume.assert_awaited_once_with("abc-one", "/ws1") assert ctx.thread_id == "abc-one" - def test_empty_workspace_metadata_preserves_ctx_workspace(self): + async def test_empty_workspace_metadata_preserves_ctx_workspace(self): from EvoScientist.commands.implementation.session import ResumeCommand ctx, ui = _ctx( workspace_dir="/keep", thread_store=FakeThreadStore(resolved_thread_id="tid", metadata={}), ) - _run(ResumeCommand().execute(ctx, ["tid"])) + await ResumeCommand().execute(ctx, ["tid"]) # ResumeCommand only overwrites ctx.workspace_dir if metadata has one assert ctx.workspace_dir == "/keep" # Callback still fires with the metadata value (empty string) diff --git a/tests/test_rich_command_ui.py b/tests/test_rich_command_ui.py index f540286..09cfcf9 100644 --- a/tests/test_rich_command_ui.py +++ b/tests/test_rich_command_ui.py @@ -5,8 +5,6 @@ from unittest.mock import MagicMock from rich.console import Console from rich.table import Table -from tests.conftest import run_async as _run - def _make_ui(**kwargs): """Build a RichCLICommandUI backed by a MagicMock console.""" @@ -40,9 +38,9 @@ class TestBasicIO: ui.mount_renderable(table) console.print.assert_called_once_with(table) - def test_flush_is_async_noop(self): + async def test_flush_is_async_noop(self): ui, console = _make_ui() - _run(ui.flush()) + await ui.flush() # flush should not print anything console.print.assert_not_called() @@ -50,33 +48,29 @@ class TestBasicIO: class TestWaitForModelPick: """CLI model picker fallback: print table + return None.""" - def test_returns_none(self): + async def test_returns_none(self): ui, _ = _make_ui() entries = [ ("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"), ("gpt-4o", "openai/gpt-4o", "openai"), ] - result = _run( - ui.wait_for_model_pick( - entries, - current_model="claude-sonnet-4-6", - current_provider="anthropic", - ) + result = await ui.wait_for_model_pick( + entries, + current_model="claude-sonnet-4-6", + current_provider="anthropic", ) assert result is None - def test_prints_table_with_current_model_marker(self): + async def test_prints_table_with_current_model_marker(self): ui, console = _make_ui() entries = [ ("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic"), ("gpt-4o", "openai/gpt-4o", "openai"), ] - _run( - ui.wait_for_model_pick( - entries, - current_model="claude-sonnet-4-6", - current_provider="anthropic", - ) + await ui.wait_for_model_pick( + entries, + current_model="claude-sonnet-4-6", + current_provider="anthropic", ) # First call renders the Table (Rich renderable), second prints usage. assert console.print.call_count == 2 @@ -87,27 +81,23 @@ class TestWaitForModelPick: assert "Usage: /model" in usage_arg assert "--save" in usage_arg - def test_no_current_model_no_marker(self): + async def test_no_current_model_no_marker(self): ui, console = _make_ui() entries = [("claude-sonnet-4-6", "anthropic/claude-sonnet", "anthropic")] - _run( - ui.wait_for_model_pick( - entries, - current_model=None, - current_provider=None, - ) + await ui.wait_for_model_pick( + entries, + current_model=None, + current_provider=None, ) # Just asserts the coroutine runs without marker-branch issues. assert console.print.call_count == 2 - def test_empty_entries_still_prints_header_and_usage(self): + async def test_empty_entries_still_prints_header_and_usage(self): ui, console = _make_ui() - result = _run( - ui.wait_for_model_pick( - [], - current_model=None, - current_provider=None, - ) + result = await ui.wait_for_model_pick( + [], + current_model=None, + current_provider=None, ) assert result is None # Header table + usage hint should still be printed even with @@ -213,7 +203,7 @@ class TestWaitForThreadPick: }, ] - def test_returns_selected_thread_id(self, monkeypatch): + async def test_returns_selected_thread_id(self, monkeypatch): import EvoScientist.cli.rich_command_ui as mod ui, _ = _make_ui() @@ -226,7 +216,7 @@ class TestWaitForThreadPick: return prompt monkeypatch.setattr("questionary.select", fake_select) - result = _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")) + result = await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:") assert result == "abc123" assert called["title"] == "pick:" # _build_items prepends a workspace header — choices has headers + @@ -235,14 +225,14 @@ class TestWaitForThreadPick: # Table import removed; this test no longer depends on console output. assert mod.RichCLICommandUI is not None # sanity - def test_cancel_returns_none(self, monkeypatch): + async def test_cancel_returns_none(self, monkeypatch): ui, _ = _make_ui() prompt = self._fake_prompt(None) monkeypatch.setattr("questionary.select", lambda *a, **k: prompt) - result = _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")) + result = await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:") assert result is None - def test_current_thread_marker_in_label(self, monkeypatch): + async def test_current_thread_marker_in_label(self, monkeypatch): ui, _ = _make_ui() prompt = self._fake_prompt(None) captured_choices: list = [] @@ -252,7 +242,7 @@ class TestWaitForThreadPick: return prompt monkeypatch.setattr("questionary.select", fake_select) - _run(ui.wait_for_thread_pick(self._threads(), "abc123", "pick:")) + await ui.wait_for_thread_pick(self._threads(), "abc123", "pick:") # At least one Choice title contains "abc123 *" (current marker) choice_titles = [getattr(c, "title", "") for c in captured_choices] assert any("abc123 *" in t for t in choice_titles) @@ -282,38 +272,38 @@ class TestCompactIndicator: class TestPhaseBMigrated: """Session lifecycle callbacks (start/resume) filled in Phase B.""" - def test_start_new_session_fires_callback(self): + async def test_start_new_session_fires_callback(self): from unittest.mock import AsyncMock cb = AsyncMock() ui, _ = _make_ui(on_start_new_session=cb) - _run(ui.start_new_session()) + await ui.start_new_session() cb.assert_awaited_once() - def test_start_new_session_without_callback_is_noop(self): + async def test_start_new_session_without_callback_is_noop(self): ui, console = _make_ui() - _run(ui.start_new_session()) + await ui.start_new_session() console.print.assert_not_called() - def test_handle_session_resume_awaits_callback(self): + async def test_handle_session_resume_awaits_callback(self): from unittest.mock import AsyncMock cb = AsyncMock() ui, _ = _make_ui(on_handle_session_resume=cb) - _run(ui.handle_session_resume("tid-x", "/workspace")) + await ui.handle_session_resume("tid-x", "/workspace") cb.assert_awaited_once_with("tid-x", "/workspace") - def test_handle_session_resume_without_callback_is_noop(self): + async def test_handle_session_resume_without_callback_is_noop(self): ui, _ = _make_ui() # Should not raise - _run(ui.handle_session_resume("tid-x")) + await ui.handle_session_resume("tid-x") - def test_handle_session_resume_workspace_defaults_none(self): + async def test_handle_session_resume_workspace_defaults_none(self): from unittest.mock import AsyncMock cb = AsyncMock() ui, _ = _make_ui(on_handle_session_resume=cb) - _run(ui.handle_session_resume("tid-x")) + await ui.handle_session_resume("tid-x") cb.assert_awaited_once_with("tid-x", None) @@ -321,7 +311,7 @@ class TestPhaseCMigrated: """Skill/MCP browse pickers delegate to questionary helpers via ``asyncio.to_thread`` since questionary blocks the event loop.""" - def test_skill_browse_delegates_to_picker(self, monkeypatch): + async def test_skill_browse_delegates_to_picker(self, monkeypatch): from unittest.mock import MagicMock picker = MagicMock(return_value=["skill-a", "skill-b"]) @@ -330,11 +320,11 @@ class TestPhaseCMigrated: picker, ) ui, _ = _make_ui() - result = _run(ui.wait_for_skill_browse([{"name": "a"}], {"installed"}, "core")) + result = await ui.wait_for_skill_browse([{"name": "a"}], {"installed"}, "core") assert result == ["skill-a", "skill-b"] picker.assert_called_once_with([{"name": "a"}], {"installed"}, "core") - def test_skill_browse_cancel_returns_none(self, monkeypatch): + async def test_skill_browse_cancel_returns_none(self, monkeypatch): from unittest.mock import MagicMock monkeypatch.setattr( @@ -342,10 +332,10 @@ class TestPhaseCMigrated: MagicMock(return_value=None), ) ui, _ = _make_ui() - result = _run(ui.wait_for_skill_browse([], set(), "")) + result = await ui.wait_for_skill_browse([], set(), "") assert result is None - def test_mcp_browse_delegates_to_picker(self, monkeypatch): + async def test_mcp_browse_delegates_to_picker(self, monkeypatch): from unittest.mock import MagicMock sentinel_entries = [MagicMock(name="entry1"), MagicMock(name="entry2")] @@ -355,11 +345,11 @@ class TestPhaseCMigrated: picker, ) ui, _ = _make_ui() - result = _run(ui.wait_for_mcp_browse([MagicMock()], {"configured"}, "")) + result = await ui.wait_for_mcp_browse([MagicMock()], {"configured"}, "") assert result is sentinel_entries picker.assert_called_once() - def test_mcp_browse_cancel_returns_none(self, monkeypatch): + async def test_mcp_browse_cancel_returns_none(self, monkeypatch): from unittest.mock import MagicMock monkeypatch.setattr( @@ -367,5 +357,5 @@ class TestPhaseCMigrated: MagicMock(return_value=None), ) ui, _ = _make_ui() - result = _run(ui.wait_for_mcp_browse([], set(), "")) + result = await ui.wait_for_mcp_browse([], set(), "") assert result is None diff --git a/tests/test_schedule_command.py b/tests/test_schedule_command.py index d3ea9ac..7daddc2 100644 --- a/tests/test_schedule_command.py +++ b/tests/test_schedule_command.py @@ -2,8 +2,6 @@ from unittest.mock import MagicMock, patch -from tests.conftest import run_async as _run - def _ctx(): from EvoScientist.commands.base import CommandContext @@ -12,17 +10,17 @@ def _ctx(): return CommandContext(agent=None, thread_id="tid", ui=ui), ui -def test_list_when_backend_down(): +async def test_list_when_backend_down(): from EvoScientist.commands.implementation.schedule import ScheduleCommand ctx, ui = _ctx() with patch("EvoScientist.cron.schedule.is_available", return_value=False): - _run(ScheduleCommand().execute(ctx, ["list"])) + await ScheduleCommand().execute(ctx, ["list"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("unavailable" in m.lower() for m in msgs) -def test_add_parses_five_field_cron_and_prompt(): +async def test_add_parses_five_field_cron_and_prompt(): from EvoScientist.commands.implementation.schedule import ScheduleCommand ctx, _ui = _ctx() @@ -33,17 +31,15 @@ def test_add_parses_five_field_cron_and_prompt(): return_value={"cron_id": "c-9"}, ) as mk, ): - _run( - ScheduleCommand().execute( - ctx, ["add", "*/10", "*", "*", "*", "*", "search", "uk", "weather"] - ) + await ScheduleCommand().execute( + ctx, ["add", "*/10", "*", "*", "*", "*", "search", "uk", "weather"] ) kw = mk.call_args.kwargs assert kw["schedule"] == "*/10 * * * *" assert kw["prompt"] == "search uk weather" -def test_list_renders_table(): +async def test_list_renders_table(): from EvoScientist.commands.implementation.schedule import ScheduleCommand ctx, ui = _ctx() @@ -60,11 +56,11 @@ def test_list_renders_table(): patch("EvoScientist.cron.schedule.is_available", return_value=True), patch("EvoScientist.cron.schedule.list_schedules", return_value=rows), ): - _run(ScheduleCommand().execute(ctx, ["list"])) + await ScheduleCommand().execute(ctx, ["list"]) ui.mount_renderable.assert_called_once() -def test_add_parses_quoted_cron_and_prompt(): +async def test_add_parses_quoted_cron_and_prompt(): from EvoScientist.commands.implementation.schedule import ScheduleCommand ctx, _ui = _ctx() @@ -75,15 +71,15 @@ def test_add_parses_quoted_cron_and_prompt(): return_value={"cron_id": "c-9"}, ) as mk, ): - _run( - ScheduleCommand().execute(ctx, ["add", "*/10 * * * *", "search uk weather"]) + await ScheduleCommand().execute( + ctx, ["add", "*/10 * * * *", "search uk weather"] ) kw = mk.call_args.kwargs assert kw["schedule"] == "*/10 * * * *" assert kw["prompt"] == "search uk weather" -def test_run_with_matching_prefix_fires_matched_prompt(): +async def test_run_with_matching_prefix_fires_matched_prompt(): from EvoScientist.commands.implementation.schedule import ScheduleCommand ctx, _ui = _ctx() @@ -96,11 +92,11 @@ def test_run_with_matching_prefix_fires_matched_prompt(): return_value={"run_id": "r-1"}, ) as rn, ): - _run(ScheduleCommand().execute(ctx, ["run", "c-123"])) + await ScheduleCommand().execute(ctx, ["run", "c-123"]) rn.assert_called_once_with("do the thing") -def test_run_with_no_match_reports(): +async def test_run_with_no_match_reports(): from EvoScientist.commands.implementation.schedule import ScheduleCommand ctx, ui = _ctx() @@ -109,13 +105,13 @@ def test_run_with_no_match_reports(): patch("EvoScientist.cron.schedule.list_schedules", return_value=[]), patch("EvoScientist.cron.schedule.run_now") as rn, ): - _run(ScheduleCommand().execute(ctx, ["run", "nope"])) + await ScheduleCommand().execute(ctx, ["run", "nope"]) rn.assert_not_called() msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("No schedule matching" in m for m in msgs) -def test_pause_resume_set_enabled_with_resolved_id(): +async def test_pause_resume_set_enabled_with_resolved_id(): from EvoScientist.commands.implementation.schedule import ScheduleCommand rows = [{"cron_id": "c-abcdef", "metadata": {"name": "t"}}] @@ -126,7 +122,7 @@ def test_pause_resume_set_enabled_with_resolved_id(): patch("EvoScientist.cron.schedule.list_schedules", return_value=rows), patch("EvoScientist.cron.schedule.set_enabled") as se, ): - _run(ScheduleCommand().execute(ctx, [sub, "c-abc"])) + await ScheduleCommand().execute(ctx, [sub, "c-abc"]) se.assert_called_once_with("c-abcdef", expected) @@ -135,7 +131,7 @@ def test_pause_resume_set_enabled_with_resolved_id(): # --------------------------------------------------------------------------- -def test_list_error_shows_red_message_no_exception(): +async def test_list_error_shows_red_message_no_exception(): """B1: list_schedules raising after is_available() shows a red error, not a traceback.""" from EvoScientist.commands.implementation.schedule import ScheduleCommand @@ -147,7 +143,7 @@ def test_list_error_shows_red_message_no_exception(): side_effect=RuntimeError("backend gone"), ), ): - _run(ScheduleCommand().execute(ctx, ["list"])) + await ScheduleCommand().execute(ctx, ["list"]) msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Error:" in m for m in msgs) # Verify no exception escaped (test would have raised above otherwise) @@ -158,7 +154,7 @@ def test_list_error_shows_red_message_no_exception(): # --------------------------------------------------------------------------- -def test_remove_ambiguous_prefix_aborts_without_deleting(): +async def test_remove_ambiguous_prefix_aborts_without_deleting(): """B2: two crons sharing a prefix → ambiguity message, delete NOT called.""" from EvoScientist.commands.implementation.schedule import ScheduleCommand @@ -172,7 +168,7 @@ def test_remove_ambiguous_prefix_aborts_without_deleting(): patch("EvoScientist.cron.schedule.list_schedules", return_value=rows), patch("EvoScientist.cron.schedule.delete_schedule") as mk, ): - _run(ScheduleCommand().execute(ctx, ["remove", "abc"])) + await ScheduleCommand().execute(ctx, ["remove", "abc"]) mk.assert_not_called() msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Multiple" in m for m in msgs) @@ -183,7 +179,7 @@ def test_remove_ambiguous_prefix_aborts_without_deleting(): # --------------------------------------------------------------------------- -def test_remove_backend_error_shows_red_error_not_no_match(): +async def test_remove_backend_error_shows_red_error_not_no_match(): """FIX 1: list_schedules() crashing in _resolve → red 'Error:' message, not 'No schedule matching'.""" from EvoScientist.commands.implementation.schedule import ScheduleCommand @@ -196,7 +192,7 @@ def test_remove_backend_error_shows_red_error_not_no_match(): ), patch("EvoScientist.cron.schedule.delete_schedule") as mk, ): - _run(ScheduleCommand().execute(ctx, ["remove", "abc"])) + await ScheduleCommand().execute(ctx, ["remove", "abc"]) mk.assert_not_called() msgs = [c.args[0] for c in ui.append_system.call_args_list] assert any("Error:" in m for m in msgs), f"Expected red Error: message, got: {msgs}" @@ -209,7 +205,7 @@ def test_remove_backend_error_shows_red_error_not_no_match(): ) -def test_add_name_sanitized_from_nasty_prompt(): +async def test_add_name_sanitized_from_nasty_prompt(): """B3: prompt with newline / slashes / special chars → clean kebab-case name.""" import re @@ -225,7 +221,7 @@ def test_add_name_sanitized_from_nasty_prompt(): return_value={"cron_id": "c-x"}, ) as mk, ): - _run(ScheduleCommand().execute(ctx, ["add", "*/5 * * * *", nasty_prompt])) + await ScheduleCommand().execute(ctx, ["add", "*/5 * * * *", nasty_prompt]) name = mk.call_args.kwargs["name"] # Must be non-empty, no spaces, no newlines, no slashes assert name diff --git a/tests/test_serve_agent_holder.py b/tests/test_serve_agent_holder.py index b94a273..37d71c7 100644 --- a/tests/test_serve_agent_holder.py +++ b/tests/test_serve_agent_holder.py @@ -27,7 +27,6 @@ from EvoScientist.cli.commands import ( from EvoScientist.commands.base import ChannelRuntime from EvoScientist.config import EvoScientistConfig from EvoScientist.gateway import RuntimeGateways, ThreadStore -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -71,7 +70,7 @@ def _runtime_state( ) -def test_hook_updates_runtime_state_on_agent_swap(): +async def test_hook_updates_runtime_state_on_agent_swap(): """``/model`` mutates ``ctx.agent`` to a new handle — the hook must push that handle into the shared runtime state so the outer poll loop sees it on the next message.""" @@ -86,12 +85,12 @@ def test_hook_updates_runtime_state_on_agent_swap(): cmd = MagicMock() cmd.name = "/model" - _run(hook(ctx, original_agent, cmd)) + await hook(ctx, original_agent, cmd) assert state.agent is new_agent -def test_hook_syncs_channel_runtime(): +async def test_hook_syncs_channel_runtime(): """Other readers (the bus) look at ``ChannelRuntime.agent``; the hook keeps the runtime in sync with the runtime state update.""" original_agent = _agent("original-agent") @@ -109,13 +108,13 @@ def test_hook_syncs_channel_runtime(): cmd = MagicMock() cmd.name = "/model" - _run(hook(ctx, original_agent, cmd)) + await hook(ctx, original_agent, cmd) assert runtime.agent is new_agent assert runtime.thread_id == "t" -def test_hook_noop_when_agent_unchanged(): +async def test_hook_noop_when_agent_unchanged(): """Commands like ``/evoskills`` don't touch ``ctx.agent`` — the runtime state must stay put.""" original_agent = _agent("original-agent") @@ -128,12 +127,12 @@ def test_hook_noop_when_agent_unchanged(): cmd = MagicMock() cmd.name = "/evoskills" - _run(hook(ctx, original_agent, cmd)) + await hook(ctx, original_agent, cmd) assert state.agent is original_agent -def test_hook_noop_when_ctx_agent_is_none(): +async def test_hook_noop_when_ctx_agent_is_none(): """Guard against commands that reset ``ctx.agent`` to ``None`` — we never want to write ``None`` into runtime state.""" original_agent = _agent("original-agent") @@ -146,12 +145,12 @@ def test_hook_noop_when_ctx_agent_is_none(): cmd = MagicMock() cmd.name = "/whatever" - _run(hook(ctx, original_agent, cmd)) + await hook(ctx, original_agent, cmd) assert state.agent is original_agent -def test_hook_updates_thread_id_on_resume(): +async def test_hook_updates_thread_id_on_resume(): """``/resume`` mutates ``ctx.thread_id`` — the hook must push the new id into runtime state so the outer poll loop runs subsequent messages on the resumed thread.""" @@ -166,12 +165,12 @@ def test_hook_updates_thread_id_on_resume(): cmd = MagicMock() cmd.name = "/resume" - _run(hook(ctx, agent, cmd)) + await hook(ctx, agent, cmd) assert state.thread_id == "new-tid" -def test_hook_updates_workspace_dir_on_resume(): +async def test_hook_updates_workspace_dir_on_resume(): """`/resume` can restore a different workspace; serve must reload for it.""" cfg = _config() old_agent = _agent("old-agent") @@ -201,7 +200,7 @@ def test_hook_updates_workspace_dir_on_resume(): return_value=reloaded_agent, ) as load_agent, ): - _run(hook(ctx, old_agent, cmd)) + await hook(ctx, old_agent, cmd) sync_server.assert_awaited_once_with(cfg, workspace_dir="/restored-ws") load_agent.assert_called_once_with(workspace_dir="/restored-ws", config=cfg) @@ -209,7 +208,7 @@ def test_hook_updates_workspace_dir_on_resume(): assert state.agent is reloaded_agent -def test_hook_syncs_channel_runtime_thread_id(): +async def test_hook_syncs_channel_runtime_thread_id(): """The bus reads ``ChannelRuntime.thread_id``; hook must sync it alongside the runtime state update.""" agent = _agent("a") @@ -224,12 +223,12 @@ def test_hook_syncs_channel_runtime_thread_id(): cmd = MagicMock() cmd.name = "/resume" - _run(hook(ctx, agent, cmd)) + await hook(ctx, agent, cmd) assert runtime.thread_id == "new-tid" -def test_hook_noop_when_thread_id_unchanged(): +async def test_hook_noop_when_thread_id_unchanged(): """Most commands don't touch thread_id — runtime state stays put.""" agent = _agent("a") state = _runtime_state(agent=agent, thread_id="same-tid") @@ -241,12 +240,12 @@ def test_hook_noop_when_thread_id_unchanged(): cmd = MagicMock() cmd.name = "/evoskills" - _run(hook(ctx, agent, cmd)) + await hook(ctx, agent, cmd) assert state.thread_id == "same-tid" -def test_hook_skips_resume_warning_when_thread_unchanged(): +async def test_hook_skips_resume_warning_when_thread_unchanged(): """Bare ``/resume`` with no argument prints usage but leaves ``ctx.thread_id`` unchanged — the in-memory-state warning must NOT fire because no resume actually happened.""" @@ -261,13 +260,13 @@ def test_hook_skips_resume_warning_when_thread_unchanged(): cmd = MagicMock() cmd.name = "/resume" - _run(hook(ctx, agent, cmd)) + await hook(ctx, agent, cmd) ctx.ui.append_system.assert_not_called() ctx.ui.flush.assert_not_called() -def test_hook_emits_resume_warning_when_thread_changed(): +async def test_hook_emits_resume_warning_when_thread_changed(): """``/resume `` that actually changes thread_id must surface the in-memory-state warning via ``ctx.ui``.""" agent = _agent("a") @@ -283,7 +282,7 @@ def test_hook_emits_resume_warning_when_thread_changed(): cmd = MagicMock() cmd.name = "/resume" - _run(hook(ctx, agent, cmd)) + await hook(ctx, agent, cmd) ctx.ui.append_system.assert_called_once() warn_text, warn_kwargs = ( @@ -296,7 +295,7 @@ def test_hook_emits_resume_warning_when_thread_changed(): ctx.ui.flush.assert_awaited_once() -def test_start_new_session_cb_rotates_thread_id(): +async def test_start_new_session_cb_rotates_thread_id(): """``/new`` via channel calls this callback — must generate a new thread id, push into runtime state, and sync the channel runtime.""" agent = _agent("a") @@ -311,13 +310,13 @@ def test_start_new_session_cb_rotates_thread_id(): state, runtime, ) - _run(cb()) + await cb() assert state.thread_id == "freshly-generated-tid" assert runtime.thread_id == "freshly-generated-tid" -def test_start_new_session_cb_leaves_agent_alone(): +async def test_start_new_session_cb_leaves_agent_alone(): """``/new`` rotates thread only — agent handle must stay put (serve's agent is a single pre-loaded instance, not per-thread).""" agent = _agent("a") @@ -328,12 +327,12 @@ def test_start_new_session_cb_leaves_agent_alone(): ) cb = _make_serve_start_new_session_cb(state) - _run(cb()) + await cb() assert state.agent is agent -def test_serve_resume_callback_syncs_reloads_and_adopts_workspace(): +async def test_serve_resume_callback_syncs_reloads_and_adopts_workspace(): cfg = _config() old_agent = _agent("old-agent") reloaded_agent = _agent("reloaded-agent") @@ -364,7 +363,7 @@ def test_serve_resume_callback_syncs_reloads_and_adopts_workspace(): side_effect=_load_agent, ) as load_agent, ): - _run(cb("new-tid", "/new-ws")) + await cb("new-tid", "/new-ws") sync_server.assert_awaited_once_with(cfg, workspace_dir="/new-ws") load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg) @@ -376,7 +375,7 @@ def test_serve_resume_callback_syncs_reloads_and_adopts_workspace(): assert runtime.agent is reloaded_agent -def test_hook_emits_resume_warning_after_resume_callback_adopts_thread(): +async def test_hook_emits_resume_warning_after_resume_callback_adopts_thread(): cfg = _config() old_agent = _agent("old-agent") reloaded_agent = _agent("reloaded-agent") @@ -399,7 +398,7 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread(): return_value=reloaded_agent, ), ): - _run(cb("abc12345-resumed-tid", "/new-ws")) + await cb("abc12345-resumed-tid", "/new-ws") hook = _make_serve_cmd_completed_hook(state, runtime, config=cfg) ctx = MagicMock() @@ -410,14 +409,14 @@ def test_hook_emits_resume_warning_after_resume_callback_adopts_thread(): cmd = MagicMock() cmd.name = "/resume" - _run(hook(ctx, reloaded_agent, cmd)) + await hook(ctx, reloaded_agent, cmd) ctx.ui.append_system.assert_called_once() assert "in-memory state" in ctx.ui.append_system.call_args.args[0] ctx.ui.flush.assert_awaited_once() -def test_serve_resume_callback_preserves_state_when_sync_fails(): +async def test_serve_resume_callback_preserves_state_when_sync_fails(): cfg = _config() old_agent = _agent("old-agent") loaded_but_not_adopted = _agent("loaded-but-not-adopted") @@ -442,7 +441,7 @@ def test_serve_resume_callback_preserves_state_when_sync_fails(): patch("EvoScientist.cli.commands.set_active_workspace") as set_active, pytest.raises(RuntimeError, match="workspace conflict"), ): - _run(cb("new-tid", "/new-ws")) + await cb("new-tid", "/new-ws") load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg) set_active.assert_called_once_with("/old-ws") @@ -455,7 +454,7 @@ def test_serve_resume_callback_preserves_state_when_sync_fails(): assert runtime.thread_id == "old-tid" -def test_serve_resume_callback_load_failure_does_not_sync_or_adopt(): +async def test_serve_resume_callback_load_failure_does_not_sync_or_adopt(): cfg = _config() old_agent = _agent("old-agent") state = _runtime_state( @@ -479,7 +478,7 @@ def test_serve_resume_callback_load_failure_does_not_sync_or_adopt(): ) as sync_server, pytest.raises(RuntimeError, match="load failed"), ): - _run(cb("new-tid", "/new-ws")) + await cb("new-tid", "/new-ws") load_agent.assert_called_once_with(workspace_dir="/new-ws", config=cfg) set_active.assert_called_once_with("/old-ws") @@ -493,7 +492,7 @@ def test_serve_resume_callback_load_failure_does_not_sync_or_adopt(): assert runtime.thread_id == "old-tid" -def test_hook_handles_both_agent_and_thread_swap(): +async def test_hook_handles_both_agent_and_thread_swap(): """Edge case: a command that changes both (hypothetical). Both updates must land in runtime state.""" old_agent = _agent("old-agent") @@ -506,7 +505,7 @@ def test_hook_handles_both_agent_and_thread_swap(): ctx.thread_id = "new-tid" cmd = MagicMock() - _run(hook(ctx, old_agent, cmd)) + await hook(ctx, old_agent, cmd) assert state.agent is new_agent assert state.thread_id == "new-tid" diff --git a/tests/test_sessions.py b/tests/test_sessions.py index a430cb2..60839e1 100644 --- a/tests/test_sessions.py +++ b/tests/test_sessions.py @@ -29,7 +29,6 @@ from EvoScientist.sessions import ( resolve_thread_id_prefix, thread_exists, ) -from tests.conftest import run_async as _run def _mock_path(db_path: str): @@ -120,7 +119,7 @@ class TestFormatRelativeTime(unittest.TestCase): assert "month" in result -class TestThreadFunctions(unittest.TestCase): +class TestThreadFunctions(unittest.IsolatedAsyncioTestCase): """Tests using a real temporary SQLite database.""" @classmethod @@ -194,7 +193,10 @@ class TestThreadFunctions(unittest.TestCase): ) await conn.commit() - _run(_setup()) + # setUpClass is a sync classmethod with no running loop, and + # IsolatedAsyncioTestCase offers no async class-level hook — + # asyncio.run() is the standard one-shot runner here. + asyncio.run(_setup()) # Patch get_db_path to point to our temp DB cls._patcher = patch( @@ -219,76 +221,76 @@ class TestThreadFunctions(unittest.TestCase): except OSError: pass - def test_list_threads(self): - threads = _run(list_threads(limit=10)) + async def test_list_threads(self): + threads = await list_threads(limit=10) # Should only contain EvoScientist threads assert len(threads) == 3 # Most recent first assert threads[0]["thread_id"] == "def00001" - def test_list_threads_with_message_count(self): - threads = _run(list_threads(limit=10, include_message_count=True)) + async def test_list_threads_with_message_count(self): + threads = await list_threads(limit=10, include_message_count=True) assert "message_count" in threads[0] - def test_thread_exists_true(self): - assert _run(thread_exists("abc12345")) + async def test_thread_exists_true(self): + assert await thread_exists("abc12345") - def test_thread_exists_false(self): - assert not _run(thread_exists("nonexist")) + async def test_thread_exists_false(self): + assert not await thread_exists("nonexist") - def test_find_similar(self): - similar = _run(find_similar_threads("abc1")) + async def test_find_similar(self): + similar = await find_similar_threads("abc1") assert len(similar) == 2 assert "abc12345" in similar assert "abc12399" in similar - def test_find_similar_no_match(self): - similar = _run(find_similar_threads("xyz")) + async def test_find_similar_no_match(self): + similar = await find_similar_threads("xyz") assert len(similar) == 0 - def test_resolve_prefix_exact_match(self): - resolved, matches = _run(resolve_thread_id_prefix("abc12345")) + async def test_resolve_prefix_exact_match(self): + resolved, matches = await resolve_thread_id_prefix("abc12345") assert resolved == "abc12345" assert matches == [] - def test_resolve_prefix_unique_prefix(self): - resolved, matches = _run(resolve_thread_id_prefix("def00")) + async def test_resolve_prefix_unique_prefix(self): + resolved, matches = await resolve_thread_id_prefix("def00") assert resolved == "def00001" assert matches == [] - def test_resolve_prefix_ambiguous(self): - resolved, matches = _run(resolve_thread_id_prefix("abc1")) + async def test_resolve_prefix_ambiguous(self): + resolved, matches = await resolve_thread_id_prefix("abc1") assert resolved is None assert set(matches) == {"abc12345", "abc12399"} - def test_resolve_prefix_not_found(self): - resolved, matches = _run(resolve_thread_id_prefix("zzz")) + async def test_resolve_prefix_not_found(self): + resolved, matches = await resolve_thread_id_prefix("zzz") assert resolved is None assert matches == [] - def test_find_similar_escapes_sql_wildcards(self): + async def test_find_similar_escapes_sql_wildcards(self): # '%' / '_' must be treated as literal characters, not SQL LIKE # wildcards, so a prefix that doesn't occur verbatim returns nothing # (prior buggy behavior: '%' matched every thread). - assert _run(find_similar_threads("%")) == [] - assert _run(find_similar_threads("_")) == [] + assert await find_similar_threads("%") == [] + assert await find_similar_threads("_") == [] - def test_get_most_recent(self): - recent = _run(get_most_recent()) + async def test_get_most_recent(self): + recent = await get_most_recent() assert recent is not None assert recent == "def00001" - def test_get_thread_metadata(self): - meta = _run(get_thread_metadata("abc12345")) + async def test_get_thread_metadata(self): + meta = await get_thread_metadata("abc12345") assert meta is not None assert meta["workspace_dir"] == "/tmp/ws_abc12345" assert meta["model"] == "claude-sonnet-4-6" - def test_get_thread_metadata_missing(self): - meta = _run(get_thread_metadata("nonexist")) + async def test_get_thread_metadata_missing(self): + meta = await get_thread_metadata("nonexist") assert meta is None - def test_delete_thread(self): + async def test_delete_thread(self): # Insert a thread to delete async def _insert(): import aiosqlite @@ -306,16 +308,16 @@ class TestThreadFunctions(unittest.TestCase): ) await conn.commit() - _run(_insert()) + await _insert() - assert _run(thread_exists("todelete")) - assert _run(delete_thread("todelete")) - assert not _run(thread_exists("todelete")) + assert await thread_exists("todelete") + assert await delete_thread("todelete") + assert not await thread_exists("todelete") - def test_delete_nonexistent(self): - assert not _run(delete_thread("nope1234")) + async def test_delete_nonexistent(self): + assert not await delete_thread("nope1234") - def test_get_thread_messages_applies_summarization_event(self): + async def test_get_thread_messages_applies_summarization_event(self): async def _insert(): import aiosqlite @@ -369,18 +371,18 @@ class TestThreadFunctions(unittest.TestCase): ) await conn.commit() - _run(_insert()) + await _insert() try: - messages = _run(get_thread_messages("sum12345")) + messages = await get_thread_messages("sum12345") assert len(messages) == 2 assert isinstance(messages[0], AIMessage) assert messages[0].content == "summary" assert isinstance(messages[1], HumanMessage) assert messages[1].content == "third" finally: - _run(_cleanup()) + await _cleanup() - def test_get_thread_messages_reconstructs_multi_delta_chain(self): + async def test_get_thread_messages_reconstructs_multi_delta_chain(self): """3-checkpoint chain with ``_DeltaSnapshot`` seed + pending writes. Exercises the upstream ``aget_delta_channel_history`` walk: the @@ -477,19 +479,19 @@ class TestThreadFunctions(unittest.TestCase): "shortcut." ) - _run(_insert()) + await _insert() try: - _run(_assert_walk_branch_active()) - messages = _run(get_thread_messages("chain12345")) + await _assert_walk_branch_active() + messages = await get_thread_messages("chain12345") assert [m.content for m in messages] == ["m1", "m2", "m3", "m4"] assert isinstance(messages[0], HumanMessage) assert isinstance(messages[1], AIMessage) assert isinstance(messages[2], HumanMessage) assert isinstance(messages[3], AIMessage) finally: - _run(_cleanup()) + await _cleanup() - def test_get_thread_messages_handles_overwrite_bare_message(self): + async def test_get_thread_messages_handles_overwrite_bare_message(self): """``Overwrite(value=)`` wraps to a single-element list. The ``Overwrite`` reset branch in ``_load_checkpoint_messages`` @@ -543,9 +545,9 @@ class TestThreadFunctions(unittest.TestCase): ) await conn.commit() - _run(_insert()) + await _insert() try: - messages = _run(get_thread_messages("ow_bare01")) + messages = await get_thread_messages("ow_bare01") # Overwrite replaced the seed completely; bare message wrapped # in a 1-element list. assert len(messages) == 1 @@ -553,9 +555,9 @@ class TestThreadFunctions(unittest.TestCase): assert messages[0].content == "replaced" assert messages[0].id == "repl" finally: - _run(_cleanup()) + await _cleanup() - def test_get_thread_messages_ignores_colliding_other_agent(self): + async def test_get_thread_messages_ignores_colliding_other_agent(self): """Multi-agent DB with thread_id collision: must surface only ours. Without the agent_name filter on the head-checkpoint lookup, @@ -617,35 +619,35 @@ class TestThreadFunctions(unittest.TestCase): ) await conn.commit() - _run(_insert()) + await _insert() try: - messages = _run(get_thread_messages("collide01")) + messages = await get_thread_messages("collide01") assert [m.content for m in messages] == ["ours_1", "ours_2"] # Defense-in-depth: explicitly forbid leakage of the other # agent's content. for msg in messages: assert not msg.content.startswith("theirs_") finally: - _run(_cleanup()) + await _cleanup() # -- Agent isolation: OtherAgent data should never be visible -- - def test_thread_exists_ignores_other_agent(self): - assert not _run(thread_exists("zzz99999")) + async def test_thread_exists_ignores_other_agent(self): + assert not await thread_exists("zzz99999") - def test_find_similar_ignores_other_agent(self): - similar = _run(find_similar_threads("zzz")) + async def test_find_similar_ignores_other_agent(self): + similar = await find_similar_threads("zzz") assert len(similar) == 0 - def test_get_metadata_ignores_other_agent(self): - meta = _run(get_thread_metadata("zzz99999")) + async def test_get_metadata_ignores_other_agent(self): + meta = await get_thread_metadata("zzz99999") assert meta is None - def test_delete_ignores_other_agent(self): + async def test_delete_ignores_other_agent(self): # Should not delete OtherAgent's data - assert not _run(delete_thread("zzz99999")) + assert not await delete_thread("zzz99999") - def test_delete_thread_preserves_other_agent_writes(self): + async def test_delete_thread_preserves_other_agent_writes(self): """Deleting a shared thread_id must only remove writes linked to EvoScientist checkpoints, leaving OtherAgent's writes intact.""" @@ -690,10 +692,10 @@ class TestThreadFunctions(unittest.TestCase): ) await conn.commit() - _run(_insert()) + await _insert() # Delete — should only affect EvoScientist's data - _run(delete_thread(shared_tid)) + await delete_thread(shared_tid) # Verify OtherAgent's writes survive async def _check(): @@ -707,12 +709,12 @@ class TestThreadFunctions(unittest.TestCase): rows = await cur.fetchall() return [r[0] for r in rows] - remaining = _run(_check()) + remaining = await _check() assert "cp_other_shared" in remaining assert "cp_evo_shared" not in remaining -class TestPruningCheckpointer(unittest.TestCase): +class TestPruningCheckpointer(unittest.IsolatedAsyncioTestCase): """Integration tests for ``PruningCheckpointer`` against a real ``AsyncSqliteSaver`` backed by a temp SQLite file. """ @@ -731,14 +733,14 @@ class TestPruningCheckpointer(unittest.TestCase): except OSError: pass - def _run_with_wrapper(self, keep: int, body): + async def _run_with_wrapper(self, keep: int, body): """Open ``PruningCheckpointer`` against the temp DB on a single loop, invoke ``body(saver)`` (an async callable), then close cleanly. Required because ``aiosqlite.Connection`` is bound to the event - loop it was opened on; reusing it across separate ``run_async`` - calls raises ``ValueError("no active connection")``. + loop it was opened on; reusing it across separate event loops + raises ``ValueError("no active connection")``. """ from EvoScientist.sessions import PruningCheckpointer @@ -749,7 +751,7 @@ class TestPruningCheckpointer(unittest.TestCase): await saver.setup() return await body(saver) - return _run(_go()) + return await _go() @staticmethod def _config(thread_id: str, ns: str = "") -> dict: @@ -778,7 +780,7 @@ class TestPruningCheckpointer(unittest.TestCase): def _metadata() -> dict: return {"agent_name": AGENT_NAME, "step": 0, "writes": {}, "parents": {}} - def _row_count(self, thread_id: str, ns: str = "") -> int: + async def _row_count(self, thread_id: str, ns: str = "") -> int: async def _count(): import aiosqlite @@ -790,9 +792,9 @@ class TestPruningCheckpointer(unittest.TestCase): row = await cur.fetchone() return int(row[0]) if row else 0 - return _run(_count()) + return await _count() - def test_aput_prunes_after_insert(self): + async def test_aput_prunes_after_insert(self): tid = "tprune01" async def _body(wrapper): @@ -804,10 +806,10 @@ class TestPruningCheckpointer(unittest.TestCase): {}, ) - self._run_with_wrapper(keep=3, body=_body) - assert self._row_count(tid) == 3 + await self._run_with_wrapper(keep=3, body=_body) + assert await self._row_count(tid) == 3 - def test_aput_keeps_latest_for_resume(self): + async def test_aput_keeps_latest_for_resume(self): """After pruning, ``aget_tuple`` must return the just-written checkpoint.""" tid = "tresume1" @@ -825,12 +827,12 @@ class TestPruningCheckpointer(unittest.TestCase): ) return last_cfg, tuple_ - last_cfg, tuple_ = self._run_with_wrapper(keep=2, body=_body) + last_cfg, tuple_ = await self._run_with_wrapper(keep=2, body=_body) assert last_cfg["configurable"]["checkpoint_id"] == "cpr_0004" assert tuple_ is not None assert tuple_.checkpoint["id"] == "cpr_0004" - def test_aput_writes_against_kept_checkpoint(self): + async def test_aput_writes_against_kept_checkpoint(self): """HITL safety: ``aput_writes`` after prune still attaches successfully.""" tid = "twrites1" @@ -848,7 +850,7 @@ class TestPruningCheckpointer(unittest.TestCase): await wrapper.aput_writes(last, [("__interrupt__", "v")], "task1") return last - last_cfg = self._run_with_wrapper(keep=2, body=_body) + last_cfg = await self._run_with_wrapper(keep=2, body=_body) async def _check(): import aiosqlite @@ -861,9 +863,9 @@ class TestPruningCheckpointer(unittest.TestCase): row = await cur.fetchone() return int(row[0]) if row else 0 - assert _run(_check()) == 1 + assert await _check() == 1 - def test_aput_partitions_by_ns(self): + async def test_aput_partitions_by_ns(self): """Two checkpoint namespaces are pruned independently.""" tid = "tns01" @@ -882,11 +884,11 @@ class TestPruningCheckpointer(unittest.TestCase): {}, ) - self._run_with_wrapper(keep=2, body=_body) - assert self._row_count(tid, ns="") == 2 - assert self._row_count(tid, ns="sub:1") == 2 + await self._run_with_wrapper(keep=2, body=_body) + assert await self._row_count(tid, ns="") == 2 + assert await self._row_count(tid, ns="sub:1") == 2 - def test_inherits_base_checkpoint_saver(self): + async def test_inherits_base_checkpoint_saver(self): """LangGraph's ``compile()`` requires ``isinstance(saver, BaseCheckpointSaver)``. Inheriting from ``AsyncSqliteSaver`` (which inherits from @@ -908,9 +910,9 @@ class TestPruningCheckpointer(unittest.TestCase): assert callable(saver.aget_tuple) assert callable(saver.aput_writes) - self._run_with_wrapper(keep=2, body=_body) + await self._run_with_wrapper(keep=2, body=_body) - def test_prune_failure_does_not_break_aput(self): + async def test_prune_failure_does_not_break_aput(self): """If pruning raises, ``aput`` still returns successfully.""" tid = "tfail01" @@ -926,10 +928,10 @@ class TestPruningCheckpointer(unittest.TestCase): {}, ) - result = self._run_with_wrapper(keep=2, body=_body) + result = await self._run_with_wrapper(keep=2, body=_body) assert result["configurable"]["checkpoint_id"] == "cpf_0001" - def test_prune_keep_zero_disables(self): + async def test_prune_keep_zero_disables(self): """``keep_per_ns=0`` is a no-op — all rows survive.""" tid = "tzero01" @@ -942,10 +944,10 @@ class TestPruningCheckpointer(unittest.TestCase): {}, ) - self._run_with_wrapper(keep=0, body=_body) - assert self._row_count(tid) == 4 + await self._run_with_wrapper(keep=0, body=_body) + assert await self._row_count(tid) == 4 - def test_prune_preserves_other_agent(self): + async def test_prune_preserves_other_agent(self): """A row with a different ``agent_name`` is never deleted.""" tid = "tother1" @@ -969,10 +971,10 @@ class TestPruningCheckpointer(unittest.TestCase): {}, ) - self._run_with_wrapper(keep=2, body=_body) + await self._run_with_wrapper(keep=2, body=_body) # OtherAgent's row + 2 EvoScientist rows = 3 total - assert self._row_count(tid) == 3 + assert await self._row_count(tid) == 3 async def _check_other(): import aiosqlite @@ -984,9 +986,9 @@ class TestPruningCheckpointer(unittest.TestCase): ) as cur: return (await cur.fetchone()) is not None - assert _run(_check_other()) + assert await _check_other() - def test_keep_one_boundary(self): + async def test_keep_one_boundary(self): """``keep_per_ns=1`` keeps only the latest row, deletes the rest.""" tid = "tk1_001" @@ -999,8 +1001,8 @@ class TestPruningCheckpointer(unittest.TestCase): {}, ) - self._run_with_wrapper(keep=1, body=_body) - assert self._row_count(tid) == 1 + await self._run_with_wrapper(keep=1, body=_body) + assert await self._row_count(tid) == 1 async def _which(): import aiosqlite @@ -1014,9 +1016,9 @@ class TestPruningCheckpointer(unittest.TestCase): return row[0] if row else None # The newest write (highest checkpoint_id) is the one kept. - assert _run(_which()) == "k1_0001" + assert await _which() == "k1_0001" - def test_concurrent_same_thread_aput_invariant(self): + async def test_concurrent_same_thread_aput_invariant(self): """Concurrent ``aput()`` calls cannot squeeze either caller's just-written row out of the top-N retention window. @@ -1055,11 +1057,11 @@ class TestPruningCheckpointer(unittest.TestCase): results = await asyncio.gather(t1, t2) return results - results = self._run_with_wrapper(keep=1, body=_body) + results = await self._run_with_wrapper(keep=1, body=_body) # Whichever caller landed last is the one survivor; importantly, # the row count is exactly 1 (no torn state where both rows # disappeared or both survived). - assert self._row_count(tid) == 1 + assert await self._row_count(tid) == 1 async def _winner(): import aiosqlite @@ -1072,7 +1074,7 @@ class TestPruningCheckpointer(unittest.TestCase): row = await cur.fetchone() return row[0] if row else None - survivor = _run(_winner()) + survivor = await _winner() # The survivor must be one of the two we wrote, not some torn ID. assert survivor in {"cc_a", "cc_b"} # And both aput results must report a valid checkpoint_id (neither @@ -1080,7 +1082,7 @@ class TestPruningCheckpointer(unittest.TestCase): for r in results: assert r["configurable"]["checkpoint_id"] in {"cc_a", "cc_b"} - def test_uuid_ordering_keeps_latest(self): + async def test_uuid_ordering_keeps_latest(self): """Uses langgraph's actual UUIDv6-shaped checkpoint IDs to confirm ``ORDER BY checkpoint_id DESC`` keeps the chronologically latest. @@ -1114,8 +1116,8 @@ class TestPruningCheckpointer(unittest.TestCase): await saver.aput(self._config(tid), cp, self._metadata(), {}) return ids - ids = self._run_with_wrapper(keep=2, body=_body) - assert self._row_count(tid) == 2 + ids = await self._run_with_wrapper(keep=2, body=_body) + assert await self._row_count(tid) == 2 async def _check(): import aiosqlite @@ -1127,13 +1129,13 @@ class TestPruningCheckpointer(unittest.TestCase): ) as cur: return [r[0] for r in await cur.fetchall()] - survivors = _run(_check()) + survivors = await _check() # The two latest UUIDv6 ids — by chronological generation — # must be the survivors. Lexicographic DESC ordering must match. assert survivors == [ids[4], ids[3]] -class TestPruningCheckpointerDeltaChannel(unittest.TestCase): +class TestPruningCheckpointerDeltaChannel(unittest.IsolatedAsyncioTestCase): """Tests for DeltaChannel-aware pruning. The naive ``keep_latest`` pruner can sever the ``_DeltaSnapshot`` @@ -1192,7 +1194,7 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): ) as cur: return [r[0] for r in await cur.fetchall()] - def test_preserves_snapshot_ancestor(self): + async def test_preserves_snapshot_ancestor(self): """Snapshot lives outside the anchor window → walk reaches and stops.""" from langgraph.checkpoint.serde.types import _DeltaSnapshot @@ -1221,10 +1223,10 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): await conn.commit() return await self._surviving_ids(conn, tid) - survivors = _run(_go()) + survivors = await _go() assert survivors == [f"cp_{i:03d}" for i in range(3, 11)] - def test_preserves_full_chain_when_no_snapshot(self): + async def test_preserves_full_chain_when_no_snapshot(self): """No snapshot anywhere → walk reaches root, preserves everything.""" from EvoScientist.sessions import PruningCheckpointer @@ -1248,10 +1250,10 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): await conn.commit() return await self._surviving_ids(conn, tid) - survivors = _run(_go()) + survivors = await _go() assert survivors == [f"cp_{i:03d}" for i in range(1, 11)] - def test_plain_list_seed_also_terminates_walk(self): + async def test_plain_list_seed_also_terminates_walk(self): """Pre-DeltaChannel format (plain list in channel_values) also counts as seed.""" from EvoScientist.sessions import PruningCheckpointer @@ -1279,10 +1281,10 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): await conn.commit() return await self._surviving_ids(conn, tid) - survivors = _run(_go()) + survivors = await _go() assert survivors == ["cp_002", "cp_003", "cp_004", "cp_005", "cp_006"] - def test_chain_break_stops_walk_cleanly(self): + async def test_chain_break_stops_walk_cleanly(self): """Missing ancestor row breaks the chain; walk stops without raising.""" from EvoScientist.sessions import PruningCheckpointer @@ -1316,13 +1318,13 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): await conn.commit() return await self._surviving_ids(conn, tid) - survivors = _run(_go()) + survivors = await _go() # anchors = cp_004, cp_005. Walk visits cp_003 (preserved), # then cp_002 → None → break. cp_001 pruned. cp_002 already # absent. Survivors: cp_003, cp_004, cp_005. assert survivors == ["cp_003", "cp_004", "cp_005"] - def test_deserialization_failure_safe_side_over_preserves(self): + async def test_deserialization_failure_safe_side_over_preserves(self): """Corrupt blob mid-walk: pruner preserves what it visited so far.""" from EvoScientist.sessions import PruningCheckpointer @@ -1357,13 +1359,13 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): await conn.commit() return await self._surviving_ids(conn, tid) - survivors = _run(_go()) + survivors = await _go() # anchors = cp_004, cp_005. Walk visits cp_003 (added to # extra_preserve before deserialize fails). cp_001, cp_002 # pruned. Survivors: cp_003, cp_004, cp_005. assert survivors == ["cp_003", "cp_004", "cp_005"] - def test_anchor_count_below_keep_is_noop(self): + async def test_anchor_count_below_keep_is_noop(self): """When checkpoint count < keep_per_ns, prune returns early without DELETE.""" from EvoScientist.sessions import PruningCheckpointer @@ -1387,11 +1389,11 @@ class TestPruningCheckpointerDeltaChannel(unittest.TestCase): await conn.commit() return await self._surviving_ids(conn, tid) - survivors = _run(_go()) + survivors = await _go() assert survivors == ["cp_001", "cp_002", "cp_003"] -class TestMigrationSweep(unittest.TestCase): +class TestMigrationSweep(unittest.IsolatedAsyncioTestCase): """Tests for the legacy-bloat migration sweep.""" def setUp(self): @@ -1426,7 +1428,7 @@ class TestMigrationSweep(unittest.TestCase): except OSError: pass - def _seed(self, threads_x_ns_x_count: list[tuple[str, str, int]]): + async def _seed(self, threads_x_ns_x_count: list[tuple[str, str, int]]): async def _go(): import aiosqlite @@ -1470,9 +1472,9 @@ class TestMigrationSweep(unittest.TestCase): ) await conn.commit() - _run(_go()) + await _go() - def _user_version(self) -> int: + async def _user_version(self) -> int: async def _go(): import aiosqlite @@ -1481,9 +1483,9 @@ class TestMigrationSweep(unittest.TestCase): row = await cur.fetchone() return int(row[0]) if row else 0 - return _run(_go()) + return await _go() - def _row_count(self, thread_id: str, ns: str) -> int: + async def _row_count(self, thread_id: str, ns: str) -> int: async def _go(): import aiosqlite @@ -1495,12 +1497,12 @@ class TestMigrationSweep(unittest.TestCase): row = await cur.fetchone() return int(row[0]) if row else 0 - return _run(_go()) + return await _go() - def test_sweep_partitions_threads_and_ns(self): + async def test_sweep_partitions_threads_and_ns(self): from EvoScientist.sessions import _run_migration_sweep - self._seed( + await self._seed( [ ("t1", "", 8), ("t1", "sub:1", 6), @@ -1508,29 +1510,29 @@ class TestMigrationSweep(unittest.TestCase): ] ) - pairs = _run(_run_migration_sweep(keep=3)) + pairs = await _run_migration_sweep(keep=3) assert pairs == 3 - assert self._row_count("t1", "") == 3 - assert self._row_count("t1", "sub:1") == 3 - assert self._row_count("t2", "") == 3 + assert await self._row_count("t1", "") == 3 + assert await self._row_count("t1", "sub:1") == 3 + assert await self._row_count("t2", "") == 3 - def test_sweep_sets_user_version(self): + async def test_sweep_sets_user_version(self): from EvoScientist.sessions import _MIGRATION_VERSION, _run_migration_sweep - self._seed([("ta", "", 5)]) - assert self._user_version() == 0 - _run(_run_migration_sweep(keep=2)) - assert self._user_version() == _MIGRATION_VERSION + await self._seed([("ta", "", 5)]) + assert await self._user_version() == 0 + await _run_migration_sweep(keep=2) + assert await self._user_version() == _MIGRATION_VERSION - def test_sweep_skipped_when_marker_set(self): + async def test_sweep_skipped_when_marker_set(self): from EvoScientist.sessions import ( _MIGRATION_VERSION, _run_migration_sweep, _set_user_version, ) - self._seed([("tb", "", 5)]) + await self._seed([("tb", "", 5)]) async def _bump(): import aiosqlite @@ -1538,39 +1540,39 @@ class TestMigrationSweep(unittest.TestCase): async with aiosqlite.connect(self._db_path) as conn: await _set_user_version(conn, _MIGRATION_VERSION) - _run(_bump()) + await _bump() # Already at marker → sweep is a no-op even though many rows exist. - pairs = _run(_run_migration_sweep(keep=2)) + pairs = await _run_migration_sweep(keep=2) assert pairs == 0 - assert self._row_count("tb", "") == 5 + assert await self._row_count("tb", "") == 5 - def test_needs_migration_below_threshold(self): + async def test_needs_migration_below_threshold(self): from EvoScientist.sessions import _needs_migration # Empty DB (file doesn't exist yet) → False - assert not _run(_needs_migration()) + assert not await _needs_migration() # Tiny DB → False - self._seed([("tc", "", 1)]) - assert not _run(_needs_migration()) + await self._seed([("tc", "", 1)]) + assert not await _needs_migration() - def test_needs_migration_above_threshold(self): + async def test_needs_migration_above_threshold(self): """Use monkeypatch on the threshold constant so tests stay fast.""" from EvoScientist import sessions as sessions_module - self._seed([("td", "", 3)]) + await self._seed([("td", "", 3)]) with patch.object(sessions_module, "_MIGRATION_THRESHOLD_BYTES", 1): # Tiny DB exceeds the 1-byte threshold → marker check kicks in. - assert _run(sessions_module._needs_migration()) + assert await sessions_module._needs_migration() - def test_keep_zero_short_circuits_sweep(self): + async def test_keep_zero_short_circuits_sweep(self): from EvoScientist.sessions import _run_migration_sweep - self._seed([("te", "", 4)]) - pairs = _run(_run_migration_sweep(keep=0)) + await self._seed([("te", "", 4)]) + pairs = await _run_migration_sweep(keep=0) assert pairs == 0 - assert self._row_count("te", "") == 4 + assert await self._row_count("te", "") == 4 - def test_sweep_handles_missing_writes_table(self): + async def test_sweep_handles_missing_writes_table(self): """Legacy DB with only ``checkpoints`` (no ``writes``) must still prune. Regression test: the sweep used to unconditionally @@ -1580,7 +1582,7 @@ class TestMigrationSweep(unittest.TestCase): from EvoScientist.sessions import _run_migration_sweep # Seed creates both tables; drop ``writes`` to simulate legacy. - self._seed([("tw", "", 5)]) + await self._seed([("tw", "", 5)]) async def _drop_writes(): import aiosqlite @@ -1589,13 +1591,13 @@ class TestMigrationSweep(unittest.TestCase): await conn.execute("DROP TABLE writes") await conn.commit() - _run(_drop_writes()) + await _drop_writes() - pairs = _run(_run_migration_sweep(keep=2)) + pairs = await _run_migration_sweep(keep=2) assert pairs == 1 - assert self._row_count("tw", "") == 2 + assert await self._row_count("tw", "") == 2 - def test_get_checkpointer_blocks_on_sweep_then_idempotent(self): + async def test_get_checkpointer_blocks_on_sweep_then_idempotent(self): """End-to-end: ``get_checkpointer()`` must run the sweep BEFORE yielding the saver so a concurrent ``aput()`` can't race the DELETEs. After the first call sets ``user_version=1``, subsequent @@ -1607,7 +1609,7 @@ class TestMigrationSweep(unittest.TestCase): get_checkpointer, ) - self._seed([("ge", "", 6)]) + await self._seed([("ge", "", 6)]) # Force the sweep to be needed regardless of file size. with patch.object(sessions_module, "_MIGRATION_THRESHOLD_BYTES", 1): @@ -1618,8 +1620,8 @@ class TestMigrationSweep(unittest.TestCase): async with get_checkpointer() as saver: return saver is not None - assert _run(_first()) - assert self._user_version() == _MIGRATION_VERSION + assert await _first() + assert await self._user_version() == _MIGRATION_VERSION # Second entry: sweep must be skipped — patch _run_migration_sweep # to raise so any accidental re-invocation fails the test loudly. @@ -1634,9 +1636,9 @@ class TestMigrationSweep(unittest.TestCase): async with get_checkpointer() as saver: return saver is not None - assert _run(_second()) + assert await _second() - def test_sweep_preserves_snapshot_ancestor(self): + async def test_sweep_preserves_snapshot_ancestor(self): """Migration sweep must apply the same DeltaChannel walk as steady-state. Without this, legacy users upgrading to PR #231 would hit a @@ -1713,8 +1715,8 @@ class TestMigrationSweep(unittest.TestCase): ) await conn.commit() - _run(_seed()) - pairs = _run(_run_migration_sweep(keep=5)) + await _seed() + pairs = await _run_migration_sweep(keep=5) assert pairs == 1 async def _survivors(): @@ -1728,7 +1730,7 @@ class TestMigrationSweep(unittest.TestCase): ) as cur: return [r[0] for r in await cur.fetchall()] - survivors = _run(_survivors()) + survivors = await _survivors() # cp_001, cp_002 pruned. cp_003 (snapshot) + walk-through (cp_004, # cp_005) + anchors (cp_006..cp_010) survive. assert survivors == [f"cp_{i:03d}" for i in range(3, 11)] @@ -1738,7 +1740,7 @@ class TestMigrationSweep(unittest.TestCase): assert "cp_002" not in survivors -class TestDbStats(unittest.TestCase): +class TestDbStats(unittest.IsolatedAsyncioTestCase): """Tests for the read-only ``db_stats`` diagnostic helper.""" def setUp(self): @@ -1761,7 +1763,7 @@ class TestDbStats(unittest.TestCase): except OSError: pass - def _seed(self): + async def _seed(self): async def _go(): import aiosqlite @@ -1835,9 +1837,9 @@ class TestDbStats(unittest.TestCase): ) await conn.commit() - _run(_go()) + await _go() - def test_stats_returns_evo_only_counts(self): + async def test_stats_returns_evo_only_counts(self): """All counts (incl. ``write_count``) must scope to EvoScientist rows. Regression for the previous bare ``COUNT(*) FROM writes`` which @@ -1847,31 +1849,31 @@ class TestDbStats(unittest.TestCase): """ from EvoScientist.sessions import db_stats - self._seed() - stats = _run(db_stats()) + await self._seed() + stats = await db_stats() assert stats["thread_count"] == 2 assert stats["checkpoint_count"] == 8 # OtherAgent's 1 row excluded assert stats["write_count"] == 4 # 2 OtherAgent writes excluded assert stats["size_bytes"] > 0 assert stats["db_path"].endswith("stats.db") - def test_stats_top_threads_ordered_desc(self): + async def test_stats_top_threads_ordered_desc(self): from EvoScientist.sessions import db_stats - self._seed() - stats = _run(db_stats(top_n=5)) + await self._seed() + stats = await db_stats(top_n=5) ids = [row["thread_id"] for row in stats["top_threads"]] counts = [row["count"] for row in stats["top_threads"]] # Sorted desc by count: evo01 (5) before evo02 (3); OtherAgent excluded assert ids == ["evo01", "evo02"] assert counts == [5, 3] - def test_stats_missing_db(self): + async def test_stats_missing_db(self): """No DB on disk → returns zeroed stats, never raises.""" from EvoScientist.sessions import db_stats # Don't seed — file doesn't exist. - stats = _run(db_stats()) + stats = await db_stats() assert stats["thread_count"] == 0 assert stats["checkpoint_count"] == 0 assert stats["write_count"] == 0 @@ -2002,11 +2004,11 @@ class TestReduceMessagesDeltaUpstreamParity: assert _signature(out) == [("HumanMessage", "d1", "x")] -class TestCreateCheckpointerForLanggraphApi(unittest.TestCase): +class TestCreateCheckpointerForLanggraphApi(unittest.IsolatedAsyncioTestCase): """Tests for ``create_checkpointer_for_langgraph_api`` — the WebUI/deploy SQLite checkpointer factory that replaces the default ``InMemorySaver``.""" - def test_yields_pruning_checkpointer(self): + async def test_yields_pruning_checkpointer(self): """Factory yields a ``PruningCheckpointer`` instance.""" from EvoScientist.sessions import ( PruningCheckpointer, @@ -2024,9 +2026,9 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase): async with create_checkpointer_for_langgraph_api() as cp: assert isinstance(cp, PruningCheckpointer) - _run(_run_inner()) + await _run_inner() - def test_checkpointer_is_set_up(self): + async def test_checkpointer_is_set_up(self): """Factory calls ``setup()`` so tables exist before yielding.""" import aiosqlite @@ -2050,9 +2052,9 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase): "checkpoints table must exist after setup()" ) - _run(_run_inner()) + await _run_inner() - def test_checkpointer_persists_across_contexts(self): + async def test_checkpointer_persists_across_contexts(self): """Data written in one context manager is readable in a new one. This is the core regression test: verifies that session data @@ -2117,7 +2119,7 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase): ) assert result.config["configurable"]["thread_id"] == thread_id - _run(_run_inner()) + await _run_inner() def test_capability_surface_matches_langgraph_api_probe(self): """Document the REAL capability surface langgraph-api will detect. @@ -2149,7 +2151,7 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase): "docstring in create_checkpointer_for_langgraph_api" ) - def test_aput_stamps_workspace_metadata_for_graph_rows(self): + async def test_aput_stamps_workspace_metadata_for_graph_rows(self): """Graph rows get workspace metadata; only main rows get agent_name.""" import json @@ -2229,10 +2231,10 @@ class TestCreateCheckpointerForLanggraphApi(unittest.TestCase): return_value="/tmp/test-workspace", ), ): - _run(_run_inner(db)) + await _run_inner(db) -class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): +class TestRestoreWebuiThreadsToGlobalStore(unittest.IsolatedAsyncioTestCase): """Tests for ``_restore_webui_threads_to_global_store``. Verifies that UUID-format threads written to SQLite by ``langgraph dev`` @@ -2288,7 +2290,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): return patch("EvoScientist.sessions._api_workspace_dir", return_value=self._WS) - def test_restores_uuid_threads_into_global_store(self): + async def test_restores_uuid_threads_into_global_store(self): """UUID-format thread IDs from SQLite are injected into GlobalStore.""" import sys from unittest.mock import MagicMock, patch @@ -2323,7 +2325,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ), self._patch_workspace(), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() # Only the UUID thread should have been added; the short-hex CLI thread # should not appear because it doesn't match the UUID LIKE pattern. @@ -2359,7 +2361,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ) assert isinstance(added[0]["updated_at"], _dt) - def test_fixes_existing_string_thread_ids_in_place(self): + async def test_fixes_existing_string_thread_ids_in_place(self): """Threads already in GlobalStore with string thread_id get fixed in-place. When .pckl loads successfully, threads are already in the store but @@ -2401,7 +2403,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ), self._patch_workspace(), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() # No duplicate: still exactly one entry. assert len(mock_store["threads"]) == 1, ( @@ -2420,7 +2422,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): assert t["metadata"].get("workspace_dir") == self._WS assert t["metadata"].get("model") == "test-model" - def test_restore_includes_current_workspace_graph_threads_only(self): + async def test_restore_includes_current_workspace_graph_threads_only(self): """Restore includes current-workspace graph threads only. Threads from other workspaces and pre-stamping rows without @@ -2468,7 +2470,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ), self._patch_workspace(), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() added = mock_store["threads"] restored = {entry["thread_id"]: entry for entry in added} @@ -2490,7 +2492,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): self._WS ) - def test_purge_removes_only_evomemory_rows(self): + async def test_purge_removes_only_evomemory_rows(self): """Startup purge drops evomemory-* residue, leaves everything else.""" import sqlite3 from unittest.mock import patch @@ -2513,9 +2515,9 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): "EvoScientist.sessions.get_db_path", return_value=_mock_path(db), ): - _run(_purge_internal_worker_threads()) + await _purge_internal_worker_threads() # Idempotent: second run is a no-op, not an error. - _run(_purge_internal_worker_threads()) + await _purge_internal_worker_threads() con = sqlite3.connect(db) remaining = { @@ -2525,7 +2527,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): assert remaining == {keep_main, keep_cli, keep_subagent} - def test_cli_session_filters_exclude_non_main_graph_rows(self): + async def test_cli_session_filters_exclude_non_main_graph_rows(self): from unittest.mock import patch from EvoScientist.sessions import ( @@ -2547,14 +2549,14 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): "EvoScientist.sessions.get_db_path", return_value=_mock_path(db), ): - assert [row["thread_id"] for row in _run(list_threads())] == [ + assert [row["thread_id"] for row in await list_threads()] == [ main_thread ] - assert _run(thread_exists(main_thread)) - assert not _run(thread_exists(worker_thread)) - assert _run(resolve_thread_id_prefix(worker_thread[:8])) == (None, []) + assert await thread_exists(main_thread) + assert not await thread_exists(worker_thread) + assert await resolve_thread_id_prefix(worker_thread[:8]) == (None, []) - def test_restores_cli_rows_and_excludes_worker_residue(self): + async def test_restores_cli_rows_and_excludes_worker_residue(self): """CLI rows (agent_name, no graph_id) are restored with graph_id backfilled; crashed-worker residue (agent_name AND graph_id= evomemory-*) stays excluded — graph_id wins over agent_name.""" @@ -2597,7 +2599,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ), self._patch_workspace(), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() added = mock_store["threads"] assert len(added) == 1 @@ -2607,7 +2609,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): assert added[0]["metadata"].get("workspace_dir") == self._WS assert added[0]["metadata"].get("model") == "test-model" - def test_mixed_cli_webui_rows_keep_assistant_and_graph_id(self): + async def test_mixed_cli_webui_rows_keep_assistant_and_graph_id(self): """Interop thread (CLI rows + WebUI rows under one UUID): bare columns under GROUP BY let SQLite pick an arbitrary row's NULL — all metadata fields must be MAX-aggregated (Codex F2).""" @@ -2649,7 +2651,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ), self._patch_workspace(), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() added = mock_store["threads"] assert len(added) == 1, f"expected 1 thread, got {added}" @@ -2659,7 +2661,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): assert added[0]["metadata"].get("workspace_dir") == self._WS assert added[0]["metadata"].get("model") == "test-model" - def test_restored_stub_gets_title_from_first_human_message(self): + async def test_restored_stub_gets_title_from_first_human_message(self): """Stubs carry metadata.title derived from the thread's first human message, so the WebUI sidebar doesn't show "Untitled Thread".""" import sys @@ -2719,13 +2721,13 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): return_value=self._WS, ), ): - _run(_write_then_restore()) + await _write_then_restore() added = mock_store["threads"] assert len(added) == 1, f"expected 1 restored thread, got {added}" assert added[0]["metadata"].get("title") == "hello title test" - def test_removes_preloaded_uuid_entries_outside_restore_scope(self): + async def test_removes_preloaded_uuid_entries_outside_restore_scope(self): """Stale and out-of-scope .pckl UUID entries are dropped. Stale UUID entries point at deleted/lost state and render as empty @@ -2773,7 +2775,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): ), self._patch_workspace(), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() ids = [t["thread_id"] for t in mock_store["threads"]] assert _uuid_mod.UUID(ghost) not in ids, f"ghost must be removed, got {ids}" @@ -2784,7 +2786,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): # In-scope thread restored as usual. assert _uuid_mod.UUID(in_scope) in ids - def test_no_op_when_langgraph_runtime_inmem_absent(self): + async def test_no_op_when_langgraph_runtime_inmem_absent(self): """ImportError for langgraph_runtime_inmem is silently swallowed.""" import sys from unittest.mock import patch @@ -2793,9 +2795,9 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): with patch.dict(sys.modules, {"langgraph_runtime_inmem.database": None}): # Must not raise. - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() - def test_no_op_when_db_has_no_checkpoints_table(self): + async def test_no_op_when_db_has_no_checkpoints_table(self): """Missing checkpoints table is handled gracefully.""" import sys from unittest.mock import MagicMock, patch @@ -2825,12 +2827,12 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): sys.modules, {"langgraph_runtime_inmem.database": fake_module} ), ): - _run(_restore_webui_threads_to_global_store()) + await _restore_webui_threads_to_global_store() # threads list untouched. assert mock_store["threads"] == [] - def test_create_checkpointer_calls_restore(self): + async def test_create_checkpointer_calls_restore(self): """create_checkpointer_for_langgraph_api calls _restore_webui_threads_to_global_store.""" from unittest.mock import patch @@ -2858,7 +2860,7 @@ class TestRestoreWebuiThreadsToGlobalStore(unittest.TestCase): async with create_checkpointer_for_langgraph_api(): pass - _run(_run_inner()) + await _run_inner() assert restore_called, "_restore_webui_threads_to_global_store must be called" diff --git a/tests/test_slack_channel.py b/tests/test_slack_channel.py index c168f7f..2c291b1 100644 --- a/tests/test_slack_channel.py +++ b/tests/test_slack_channel.py @@ -4,7 +4,6 @@ import pytest from EvoScientist.channels.base import ChannelError from EvoScientist.channels.slack.channel import SlackChannel, SlackConfig -from tests.conftest import run_async as _run class TestSlackConfig: @@ -38,24 +37,24 @@ class TestSlackChannel: assert channel.config is config assert channel._running is False - def test_start_raises_without_bot_token(self): + async def test_start_raises_without_bot_token(self): config = SlackConfig(bot_token="", app_token="xapp-test") channel = SlackChannel(config) with pytest.raises(ChannelError, match="bot token"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_app_token(self): + async def test_start_raises_without_app_token(self): config = SlackConfig(bot_token="xoxb-test", app_token="") channel = SlackChannel(config) with pytest.raises(ChannelError, match="app token"): - _run(channel.start()) + await channel.start() - def test_stop_when_not_running(self): + async def test_stop_when_not_running(self): config = SlackConfig(bot_token="xoxb-test", app_token="xapp-test") channel = SlackChannel(config) - _run(channel.stop()) + await channel.stop() - def test_send_returns_false_without_client(self): + async def test_send_returns_false_without_client(self): from EvoScientist.channels.base import OutboundMessage config = SlackConfig(bot_token="xoxb-test", app_token="xapp-test") @@ -66,7 +65,7 @@ class TestSlackChannel: content="hello", metadata={"chat_id": "C123"}, ) - result = _run(channel.send(msg)) + result = await channel.send(msg) assert result is False diff --git a/tests/test_status_bar.py b/tests/test_status_bar.py index 271afb1..4b19b98 100644 --- a/tests/test_status_bar.py +++ b/tests/test_status_bar.py @@ -2,7 +2,6 @@ from __future__ import annotations -import asyncio from datetime import datetime, timedelta from typing import ClassVar @@ -243,7 +242,7 @@ def test_build_status_text_uses_rich_styles(): assert text.spans -def test_build_session_status_snapshot_uses_fallback_window(monkeypatch): +async def test_build_session_status_snapshot_uses_fallback_window(monkeypatch): class _FakeModel: model_name: ClassVar[str] = "provider/demo-model" profile: ClassVar[dict[str, object]] = {} @@ -262,16 +261,12 @@ def test_build_session_status_snapshot_uses_fallback_window(monkeypatch): _fake_count, ) - snapshot = asyncio.run( - build_session_status_snapshot( - "thread-1", - pending_user_text="pending", - graph_gateway=FakeGraphGateway( - thread_store=FakeThreadStore( - messages=[HumanMessage(content="existing")] - ) - ), - ) + snapshot = await build_session_status_snapshot( + "thread-1", + pending_user_text="pending", + graph_gateway=FakeGraphGateway( + thread_store=FakeThreadStore(messages=[HumanMessage(content="existing")]) + ), ) assert snapshot.model_full == "provider/demo-model" diff --git a/tests/test_stream_events.py b/tests/test_stream_events.py index d3489ac..c1d3491 100644 --- a/tests/test_stream_events.py +++ b/tests/test_stream_events.py @@ -20,7 +20,6 @@ from EvoScientist.stream.tool_results import ( _extract_command_tool_content, _extract_tool_content, ) -from tests.conftest import run_async from tests.stream_v3_fakes import ( ErroringV3Agent, FakeSubagent, @@ -144,10 +143,10 @@ class TestExtractToolContent: class TestV3ProtocolStreaming: """Test stream_agent_events against v3 protocol events.""" - def test_message_delta_emits_text(self): + async def test_message_delta_emits_text(self): """v3 content-block text deltas are processed.""" agent = FakeV3Agent([message_delta("hello world")]) - events = collect_events(agent) + events = await collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert len(text_events) == 1 assert text_events[0]["content"] == "hello world" @@ -156,7 +155,7 @@ class TestV3ProtocolStreaming: assert "stream_mode" not in kwargs assert "subgraphs" not in kwargs - def test_streamed_non_selector_json_is_replayed(self): + async def test_streamed_non_selector_json_is_replayed(self): """Normal JSON answers are not swallowed by selector JSON buffering.""" agent = FakeV3Agent( [ @@ -165,13 +164,13 @@ class TestV3ProtocolStreaming: message_delta(": 1}"), ] ) - events = collect_events(agent) + events = await collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert "".join(e["content"] for e in text_events) == '{"answer": 1}' assert events[-1]["type"] == "done" assert events[-1]["response"] == '{"answer": 1}' - def test_incomplete_non_selector_json_flushes_on_message_finish(self): + async def test_incomplete_non_selector_json_flushes_on_message_finish(self): """Buffered non-selector text is not lost if the message ends mid-object.""" agent = FakeV3Agent( [ @@ -180,40 +179,42 @@ class TestV3ProtocolStreaming: message_finish(), ] ) - events = collect_events(agent) + events = await collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert "".join(e["content"] for e in text_events) == '{"answer":' assert events[-1]["response"] == '{"answer":' - def test_json_answer_with_tools_key_is_replayed_without_selector_context(self): + async def test_json_answer_with_tools_key_is_replayed_without_selector_context( + self, + ): """Normal answers may legitimately contain a top-level tools key.""" agent = FakeV3Agent( [ message_delta('{"tools":["hammer"],"answer":"use safely"}'), ] ) - events = collect_events(agent) + events = await collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert len(text_events) == 1 assert text_events[0]["content"] == '{"tools":["hammer"],"answer":"use safely"}' assert events[-1]["response"] == '{"tools":["hammer"],"answer":"use safely"}' - def test_text_delta_strips_legacy_thinking_tags(self): + async def test_text_delta_strips_legacy_thinking_tags(self): """Legacy tags are still removed on the v3 text path.""" agent = FakeV3Agent( [message_delta("some reasoningThe answer is 42.")] ) - events = collect_events(agent) + events = await collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert len(text_events) == 1 assert text_events[0]["content"] == "The answer is 42." - def test_text_delta_with_only_legacy_thinking_tags_is_skipped(self): + async def test_text_delta_with_only_legacy_thinking_tags_is_skipped(self): agent = FakeV3Agent([message_delta("just reasoning")]) - events = collect_events(agent) + events = await collect_events(agent) assert [e for e in events if e.get("type") == "text"] == [] - def test_updates_event_without_summary_is_skipped(self): + async def test_updates_event_without_summary_is_skipped(self): """Non-summary updates are skipped without error.""" agent = FakeV3Agent( [ @@ -221,12 +222,14 @@ class TestV3ProtocolStreaming: message_delta("should appear"), ] ) - events = collect_events(agent) + events = await collect_events(agent) text_events = [e for e in events if e.get("type") == "text"] assert len(text_events) == 1 assert text_events[0]["content"] == "should appear" - def test_user_message_clears_completed_memory_activity_counts(self, monkeypatch): + async def test_user_message_clears_completed_memory_activity_counts( + self, monkeypatch + ): calls = [] monkeypatch.setattr( "EvoScientist.stream.events.clear_completed_memory_activity_counts", @@ -234,11 +237,13 @@ class TestV3ProtocolStreaming: ) agent = FakeV3Agent([]) - collect_events(agent, message="new user turn") + await collect_events(agent, message="new user turn") assert calls == [True] - def test_command_message_clears_completed_memory_activity_counts(self, monkeypatch): + async def test_command_message_clears_completed_memory_activity_counts( + self, monkeypatch + ): calls = [] monkeypatch.setattr( "EvoScientist.stream.events.clear_completed_memory_activity_counts", @@ -247,12 +252,12 @@ class TestV3ProtocolStreaming: agent = FakeV3Agent([]) resume_command = Command(resume={"decisions": [{"type": "approve"}]}) - collect_events(agent, message=resume_command) + await collect_events(agent, message=resume_command) assert calls == [True] assert agent.astream_events.call_args.args[0] is resume_command - def test_summarization_filtered(self): + async def test_summarization_filtered(self): """v3 messages with lc_source=summarization emit summarization events.""" agent = FakeV3Agent( [ @@ -260,7 +265,7 @@ class TestV3ProtocolStreaming: message_delta("real content"), ] ) - events = collect_events(agent) + events = await collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -272,7 +277,7 @@ class TestV3ProtocolStreaming: assert len(text_events) == 1 assert text_events[0]["content"] == "real content" - def test_updates_mode_summarization_event_emitted(self): + async def test_updates_mode_summarization_event_emitted(self): """_summarization_event updates should emit a summarization event.""" summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts", @@ -294,7 +299,7 @@ class TestV3ProtocolStreaming: message_delta("real content"), ] ) - events = collect_events(agent) + events = await collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -303,7 +308,7 @@ class TestV3ProtocolStreaming: assert len(summary_events) == 1 assert summary_events[0]["content"] == "Key facts" - def test_updates_mode_does_not_duplicate_streamed_summarization(self): + async def test_updates_mode_does_not_duplicate_streamed_summarization(self): """If streamed summarization already emitted, updates fallback should not duplicate it.""" summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts" @@ -324,7 +329,7 @@ class TestV3ProtocolStreaming: message_delta("real content"), ] ) - events = collect_events(agent) + events = await collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -333,7 +338,7 @@ class TestV3ProtocolStreaming: assert len(summary_events) == 1 assert summary_events[0]["content"] == "synthetic summary" - def test_updates_mode_does_not_reemit_existing_summarization_event(self): + async def test_updates_mode_does_not_reemit_existing_summarization_event(self): """Persisted _summarization_event from a prior turn should not be replayed.""" summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts", @@ -352,7 +357,7 @@ class TestV3ProtocolStreaming: ], state_values=summary_event, ) - events = collect_events(agent) + events = await collect_events(agent) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" ] @@ -360,7 +365,7 @@ class TestV3ProtocolStreaming: summary_events = [e for e in events if e.get("type") == "summarization"] assert summary_events == [] - def test_direct_stream_loads_existing_summarization_event_when_omitted(self): + async def test_direct_stream_loads_existing_summarization_event_when_omitted(self): """Public stream_agent_events() suppresses persisted summary replays.""" summary_message = HumanMessage( content="Here is a summary of the conversation to date:\n\nKey facts", @@ -380,13 +385,9 @@ class TestV3ProtocolStreaming: state_values=summary_event, ) - async def _collect(): - events = [] - async for event in stream_agent_events(agent, "hi", "t1"): - events.append(event) - return events - - events = run_async(_collect()) + events = [] + async for event in stream_agent_events(agent, "hi", "t1"): + events.append(event) summary_start_events = [ e for e in events if e.get("type") == "summarization_start" @@ -395,19 +396,19 @@ class TestV3ProtocolStreaming: summary_events = [e for e in events if e.get("type") == "summarization"] assert summary_events == [] - def test_whole_message_reasoning_is_not_duplicated(self): + async def test_whole_message_reasoning_is_not_duplicated(self): """Providers can expose the same reasoning in kwargs and content blocks.""" message = AIMessage( additional_kwargs={"reasoning_content": "Think once."}, content=[{"type": "reasoning", "reasoning": "Think once."}], ) agent = FakeV3Agent([protocol_event("messages", (message, {}))]) - events = collect_events(agent) + events = await collect_events(agent) thinking_events = [e for e in events if e.get("type") == "thinking"] assert len(thinking_events) == 1 assert thinking_events[0]["content"] == "Think once." - def test_tool_selector_reasoning_delta_is_suppressed(self): + async def test_tool_selector_reasoning_delta_is_suppressed(self): """Selector reasoning must not appear as main-agent thinking.""" import EvoScientist.middleware.tool_selector as selector_mod @@ -432,7 +433,7 @@ class TestV3ProtocolStreaming: ) ] ) - events = collect_events(agent) + events = await collect_events(agent) finally: selector_mod._selector_active = original_active @@ -441,7 +442,7 @@ class TestV3ProtocolStreaming: for e in events ) - def test_tool_selector_whole_message_reasoning_is_suppressed(self): + async def test_tool_selector_whole_message_reasoning_is_suppressed(self): """Selector reasoning in whole-message payloads is also hidden.""" import EvoScientist.middleware.tool_selector as selector_mod @@ -453,7 +454,7 @@ class TestV3ProtocolStreaming: content="", ) agent = FakeV3Agent([protocol_event("messages", (message, {}))]) - events = collect_events(agent) + events = await collect_events(agent) finally: selector_mod._selector_active = original_active @@ -462,7 +463,7 @@ class TestV3ProtocolStreaming: for e in events ) - def test_tool_events_emit_call_and_result(self): + async def test_tool_events_emit_call_and_result(self): """v3 tool projection events become UI tool call/result events.""" output = ToolMessage( name="read_file", @@ -475,7 +476,7 @@ class TestV3ProtocolStreaming: tool_finished(output), ] ) - events = collect_events(agent) + events = await collect_events(agent) tool_call = next(e for e in events if e.get("type") == "tool_call") tool_result = next(e for e in events if e.get("type") == "tool_result") assert tool_call["name"] == "read_file" @@ -489,7 +490,7 @@ class TestV3ProtocolStreaming: @pytest.mark.filterwarnings( "ignore:The v3 streaming protocol on Pregel is experimental" ) - def test_live_deepagents_v3_tool_result_preserves_tool_call_id(self): + async def test_live_deepagents_v3_tool_result_preserves_tool_call_id(self): """DeepAgents v3 emits tool_call_id on started and finished tool events.""" @tool @@ -518,17 +519,14 @@ class TestV3ProtocolStreaming: system_prompt="Use tools when requested.", ) - async def _collect_events(): - return [ - event - async for event in stream_agent_events( - agent, - "run probe", - "live-deepagents-tool-id", - ) - ] - - events = run_async(_collect_events()) + events = [ + event + async for event in stream_agent_events( + agent, + "run probe", + "live-deepagents-tool-id", + ) + ] tool_call = next(e for e in events if e.get("type") == "tool_call") tool_result = next(e for e in events if e.get("type") == "tool_result") @@ -551,7 +549,7 @@ class TestV3ProtocolStreaming: @pytest.mark.filterwarnings( "ignore:The v3 streaming protocol on Pregel is experimental" ) - def test_live_deepagents_v3_hitl_emits_tool_call_and_single_interrupt(self): + async def test_live_deepagents_v3_hitl_emits_tool_call_and_single_interrupt(self): """Live HITL streams the model tool call once before one interrupt.""" @tool @@ -581,17 +579,14 @@ class TestV3ProtocolStreaming: checkpointer=InMemorySaver(), ) - async def _collect_events(): - return [ - event - async for event in stream_agent_events( - agent, - "run echo", - "live-deepagents-hitl", - ) - ] - - events = run_async(_collect_events()) + events = [ + event + async for event in stream_agent_events( + agent, + "run echo", + "live-deepagents-hitl", + ) + ] tool_calls = [e for e in events if e.get("type") == "tool_call"] interrupts = [e for e in events if e.get("type") == "interrupt"] @@ -611,7 +606,7 @@ class TestV3ProtocolStreaming: @pytest.mark.filterwarnings( "ignore:The v3 streaming protocol on Pregel is experimental" ) - def test_live_deepagents_v3_ask_user_suppresses_interrupt_tool_result(self): + async def test_live_deepagents_v3_ask_user_suppresses_interrupt_tool_result(self): """ask_user pause markers are not displayed as failed tool results.""" model = _ToolCallingFakeModel( @@ -654,15 +649,15 @@ class TestV3ProtocolStreaming: ) ] - first_events = run_async(_collect("ask")) + first_events = await _collect("ask") first_types = [event.get("type") for event in first_events] assert first_types == ["tool_call", "ask_user", "done"] ask_event = next(e for e in first_events if e.get("type") == "ask_user") assert ask_event["tool_call_id"] == "call_ask_1" assert ask_event["questions"] == [{"question": "What dataset?", "type": "text"}] - resumed_events = run_async( - _collect(Command(resume={"answers": ["CIFAR-10"], "status": "answered"})) + resumed_events = await _collect( + Command(resume={"answers": ["CIFAR-10"], "status": "answered"}) ) tool_result = next(e for e in resumed_events if e.get("type") == "tool_result") assert tool_result == { @@ -678,7 +673,9 @@ class TestV3ProtocolStreaming: @pytest.mark.filterwarnings( "ignore:The v3 streaming protocol on Pregel is experimental" ) - def test_live_deepagents_v3_task_result_uses_subagent_tool_message_content(self): + async def test_live_deepagents_v3_task_result_uses_subagent_tool_message_content( + self, + ): """Live task results should display the subagent ToolMessage content.""" root_model = _ToolCallingFakeModel( @@ -717,17 +714,14 @@ class TestV3ProtocolStreaming: ], ) - async def _collect_events(): - return [ - event - async for event in stream_agent_events( - agent, - "delegate", - "live-deepagents-subagent", - ) - ] - - events = run_async(_collect_events()) + events = [ + event + async for event in stream_agent_events( + agent, + "delegate", + "live-deepagents-subagent", + ) + ] subagent_start = next(e for e in events if e.get("type") == "subagent_start") subagent_end = next(e for e in events if e.get("type") == "subagent_end") @@ -745,7 +739,7 @@ class TestV3ProtocolStreaming: assert task_result["content"] == "subagent final" assert "Command(" not in task_result["content"] - def test_message_tool_call_block_emits_pre_execution_tool_call(self): + async def test_message_tool_call_block_emits_pre_execution_tool_call(self): """Model-declared tool calls remain visible before execution starts.""" agent = FakeV3Agent( [ @@ -776,14 +770,14 @@ class TestV3ProtocolStreaming: ), ] ) - events = collect_events(agent) + events = await collect_events(agent) event_types = [e["type"] for e in events] assert event_types.index("tool_call") < event_types.index("interrupt") tool_call = next(e for e in events if e.get("type") == "tool_call") assert tool_call["id"] == "tc-msg" assert tool_call["args"] == {"command": "ls"} - def test_tool_selection_flushes_before_tool_only_step(self): + async def test_tool_selection_flushes_before_tool_only_step(self): """Selector UI event is emitted even when selection is followed only by a tool.""" import EvoScientist.middleware.tool_selector as selector_mod @@ -806,7 +800,7 @@ class TestV3ProtocolStreaming: tool_finished(output), ] ) - events = collect_events(agent) + events = await collect_events(agent) finally: selector_mod._current_selected_tools = original_selected selector_mod._total_tools_count = original_total @@ -817,7 +811,7 @@ class TestV3ProtocolStreaming: selection = next(e for e in events if e.get("type") == "tool_selection") assert selection["tools"] == ["read_file"] - def test_subagent_projection_routes_namespaced_events(self): + async def test_subagent_projection_routes_namespaced_events(self): """DeepAgents subagent projection supplies identity for namespaced events.""" namespace = ("task", "abc") output = ToolMessage( @@ -838,7 +832,7 @@ class TestV3ProtocolStreaming: ], subagents=[FakeSubagent(namespace, "research-agent")], ) - events = collect_events(agent) + events = await collect_events(agent) assert any(e.get("type") == "subagent_start" for e in events) assert any(e.get("type") == "subagent_end" for e in events) @@ -865,7 +859,7 @@ class TestV3ProtocolStreaming: ) assert event_types.index("subagent_end") < event_types.index("done") - def test_namespaced_events_wait_for_delayed_subagent_registration(self): + async def test_namespaced_events_wait_for_delayed_subagent_registration(self): """Subagent events are not dropped if protocol events arrive first.""" namespace = ("task", "late") @@ -900,7 +894,7 @@ class TestV3ProtocolStreaming: return Snapshot() - events = collect_events(Agent()) + events = await collect_events(Agent()) event_types = [e["type"] for e in events] text = next(e for e in events if e.get("type") == "subagent_text") @@ -908,7 +902,7 @@ class TestV3ProtocolStreaming: assert text["instance_id"] == "task:late" assert event_types.index("subagent_start") < event_types.index("subagent_text") - def test_subagent_tool_dedupe_uses_resolved_path(self): + async def test_subagent_tool_dedupe_uses_resolved_path(self): """Tool call/result events can arrive on namespace suffixes for one subagent.""" subagent_path = ("task", "abc") call_namespace = (*subagent_path, "agent") @@ -936,7 +930,7 @@ class TestV3ProtocolStreaming: ], subagents=[FakeSubagent(subagent_path, "research-agent")], ) - events = collect_events(agent) + events = await collect_events(agent) calls = [e for e in events if e.get("type") == "subagent_tool_call"] results = [e for e in events if e.get("type") == "subagent_tool_result"] @@ -948,7 +942,7 @@ class TestV3ProtocolStreaming: assert results[0]["instance_id"] == "task:abc" assert results[0]["id"] == "sa-tc" - def test_subagent_end_is_emitted_before_later_root_text(self): + async def test_subagent_end_is_emitted_before_later_root_text(self): """Finished subagents stop showing as active while root streaming continues.""" output_returned = asyncio.Event() @@ -996,19 +990,19 @@ class TestV3ProtocolStreaming: return Snapshot() - events = collect_events(Agent()) + events = await collect_events(Agent()) event_types = [e["type"] for e in events] assert event_types.index("subagent_end") < event_types.index("text") - def test_subagent_projection_is_subscribed_before_protocol_pump(self): + async def test_subagent_projection_is_subscribed_before_protocol_pump(self): """Subagent handles are not dropped by lazy projection subscription.""" namespace = ("task", "early") agent = SubscriptionSensitiveV3Agent( [message_delta("Sub-agent finding.", namespace=namespace)], [FakeSubagent(namespace, "research-agent")], ) - events = collect_events(agent) + events = await collect_events(agent) assert any(e.get("type") == "subagent_start" for e in events) assert any(e.get("type") == "subagent_end" for e in events) assert [e for e in events if e.get("type") == "text"] == [] @@ -1018,7 +1012,7 @@ class TestV3ProtocolStreaming: assert text["content"] == "Sub-agent finding." assert text["instance_id"] == "task:early" - def test_parallel_same_name_subagent_events_carry_instance_ids(self): + async def test_parallel_same_name_subagent_events_carry_instance_ids(self): """Lifecycle and tool events distinguish same-name parallel subagents.""" ns1 = ("task", "one") ns2 = ("task", "two") @@ -1048,7 +1042,7 @@ class TestV3ProtocolStreaming: FakeSubagent(ns2, "research-agent"), ], ) - events = collect_events(agent) + events = await collect_events(agent) starts = [e for e in events if e.get("type") == "subagent_start"] calls = [e for e in events if e.get("type") == "subagent_tool_call"] @@ -1060,7 +1054,7 @@ class TestV3ProtocolStreaming: assert {e["instance_id"] for e in results} == {"task:one", "task:two"} assert {e["instance_id"] for e in ends} == {"task:one", "task:two"} - def test_stream_construction_error_emits_error_before_reraising(self): + async def test_stream_construction_error_emits_error_before_reraising(self): """astream_events construction failures preserve the UI error event contract.""" events = [] @@ -1073,10 +1067,10 @@ class TestV3ProtocolStreaming: events.append(ev) with pytest.raises(RuntimeError, match="boom"): - run_async(collect()) + await collect() assert events == [{"type": "error", "message": "boom"}] - def test_generator_close_aborts_underlying_v3_stream(self): + async def test_generator_close_aborts_underlying_v3_stream(self): """Early consumer exit should abort the caller-driven v3 run.""" async def consume_one_and_close(): @@ -1090,7 +1084,7 @@ class TestV3ProtocolStreaming: await stream.aclose() return first, agent.aborted - first, aborted = run_async(consume_one_and_close()) + first, aborted = await consume_one_and_close() assert first["type"] == "text" assert first["content"] == "hi" assert aborted is True @@ -1099,7 +1093,7 @@ class TestV3ProtocolStreaming: class TestUsageStatsExtraction: """Test token usage extraction from v3 message-finish events.""" - def test_usage_metadata_emitted(self): + async def test_usage_metadata_emitted(self): """v3 message-finish usage emits usage_stats event.""" agent = FakeV3Agent( [ @@ -1113,16 +1107,16 @@ class TestUsageStatsExtraction: ), ] ) - events = collect_events(agent) + events = await collect_events(agent) usage_events = [e for e in events if e.get("type") == "usage_stats"] assert len(usage_events) == 1 assert usage_events[0]["input_tokens"] == 100 assert usage_events[0]["output_tokens"] == 50 - def test_no_usage_metadata_no_event(self): + async def test_no_usage_metadata_no_event(self): """message-finish without usage does not emit usage_stats.""" agent = FakeV3Agent([message_delta("hi"), message_finish()]) - events = collect_events(agent) + events = await collect_events(agent) usage_events = [e for e in events if e.get("type") == "usage_stats"] assert len(usage_events) == 0 @@ -1158,7 +1152,7 @@ class TestSummarizationHelpers: assert isinstance(summary_message, HumanMessage) assert summary_message.content == "Summary body" - def test_zero_tokens_not_emitted(self): + async def test_zero_tokens_not_emitted(self): """Zero input and output tokens should not emit usage_stats.""" agent = FakeV3Agent( [ @@ -1168,6 +1162,6 @@ class TestSummarizationHelpers: ), ] ) - events = collect_events(agent) + events = await collect_events(agent) usage_events = [e for e in events if e.get("type") == "usage_stats"] assert len(usage_events) == 0 diff --git a/tests/test_stream_recovery.py b/tests/test_stream_recovery.py index 10e311d..de57eda 100644 --- a/tests/test_stream_recovery.py +++ b/tests/test_stream_recovery.py @@ -16,7 +16,6 @@ from langgraph.graph import END, START, StateGraph from langgraph.types import interrupt from EvoScientist.stream.events import _clear_interrupted_graph_state -from tests.conftest import run_async as _run class _S(TypedDict): @@ -58,7 +57,7 @@ def _interrupting_app(): return g.compile(checkpointer=InMemorySaver()) -def test_recovery_clears_stuck_state_after_crash(): +async def test_recovery_clears_stuck_state_after_crash(): app = _crashing_app() cfg = {"configurable": {"thread_id": "t1"}} try: @@ -68,7 +67,7 @@ def test_recovery_clears_stuck_state_after_crash(): # The crash left the graph frozen at node 'b'. assert app.get_state(cfg).next == ("b",) - _run(_clear_interrupted_graph_state(app, cfg)) + await _clear_interrupted_graph_state(app, cfg) snap = app.get_state(cfg) assert snap.next == () # stuck state actually cleared @@ -79,7 +78,7 @@ def test_recovery_clears_stuck_state_after_crash(): assert app.invoke({"x": 41}, cfg)["x"] == 142 -def test_recovery_preserves_pending_hitl_interrupt(): +async def test_recovery_preserves_pending_hitl_interrupt(): app = _interrupting_app() cfg = {"configurable": {"thread_id": "t1"}} app.invoke({"x": 0}, cfg) # parks at interrupt() @@ -87,7 +86,7 @@ def test_recovery_preserves_pending_hitl_interrupt(): assert before.next == ("ask",) assert before.interrupts - _run(_clear_interrupted_graph_state(app, cfg)) + await _clear_interrupted_graph_state(app, cfg) after = app.get_state(cfg) assert after.next == ("ask",) # interrupt left intact, still resumable diff --git a/tests/test_stt.py b/tests/test_stt.py index 3245019..1e8b731 100644 --- a/tests/test_stt.py +++ b/tests/test_stt.py @@ -12,7 +12,6 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest from EvoScientist.stt import STT_MODELS, is_audio_file, transcribe_file -from tests.conftest import run_async # ── is_audio_file ───────────────────────────────────────────────────── @@ -50,8 +49,8 @@ def test_stt_models_keys(): # ── transcribe_file: non-audio path ────────────────────────────────── -def test_transcribe_non_audio_returns_none(): - result = run_async(transcribe_file("photo.jpg", language="auto")) +async def test_transcribe_non_audio_returns_none(): + result = await transcribe_file("photo.jpg", language="auto") assert result is None @@ -78,37 +77,37 @@ def _patch_whisper(whisper_model): ) -def test_transcribe_en_uses_whisper(): +async def test_transcribe_en_uses_whisper(): import EvoScientist.stt as stt_mod stt_mod._engine = None with _patch_whisper(_make_whisper_mock("Hello world")): - result = run_async(transcribe_file("voice.mp3", language="en")) + result = await transcribe_file("voice.mp3", language="en") stt_mod._engine = None assert result == "Hello world" -def test_transcribe_auto_uses_whisper(): +async def test_transcribe_auto_uses_whisper(): import EvoScientist.stt as stt_mod stt_mod._engine = None with _patch_whisper(_make_whisper_mock("Bonjour monde")): - result = run_async(transcribe_file("voice.ogg", language="auto")) + result = await transcribe_file("voice.ogg", language="auto") stt_mod._engine = None assert result == "Bonjour monde" -def test_transcribe_zh_uses_whisper(): +async def test_transcribe_zh_uses_whisper(): import EvoScientist.stt as stt_mod stt_mod._engine = None with _patch_whisper(_make_whisper_mock("你好世界")): - result = run_async(transcribe_file("voice.ogg", language="zh")) + result = await transcribe_file("voice.ogg", language="zh") stt_mod._engine = None assert result == "你好世界" -def test_transcribe_custom_model_override(): +async def test_transcribe_custom_model_override(): """stt_model config overrides the default model mapping.""" import EvoScientist.stt as stt_mod @@ -121,10 +120,8 @@ def test_transcribe_custom_model_override(): self._model = _make_whisper_mock("test") with patch.object(stt_mod._WhisperEngine, "__init__", patched_init): - run_async( - transcribe_file( - "voice.ogg", language="auto", model="openai/whisper-large-v3" - ) + await transcribe_file( + "voice.ogg", language="auto", model="openai/whisper-large-v3" ) stt_mod._engine = None assert captured_model_id == ["openai/whisper-large-v3"] @@ -133,7 +130,7 @@ def test_transcribe_custom_model_override(): # ── transcribe_file: missing dependency ────────────────────────────── -def test_transcribe_missing_dep_returns_none(): +async def test_transcribe_missing_dep_returns_none(): import sys import EvoScientist.stt as stt_mod @@ -142,7 +139,7 @@ def test_transcribe_missing_dep_returns_none(): saved = sys.modules.pop("faster_whisper", None) try: with patch.dict("sys.modules", {"faster_whisper": None}): - result = run_async(transcribe_file("voice.mp3", language="auto")) + result = await transcribe_file("voice.mp3", language="auto") finally: if saved is not None: sys.modules["faster_whisper"] = saved @@ -169,7 +166,7 @@ def _make_channel(): return ch, captured -def test_enqueue_raw_stt_prepends_transcript(): +async def test_enqueue_raw_stt_prepends_transcript(): """_enqueue_raw prepends STT transcript to raw.text when stt_enabled.""" from EvoScientist.channels.base import RawIncoming @@ -189,22 +186,18 @@ def test_enqueue_raw_stt_prepends_transcript(): timestamp=datetime.now(), ) - async def _run(): - with ( - patch( - "EvoScientist.stt.transcribe_file", new=AsyncMock(return_value="你好") - ), - patch("EvoScientist.stt.is_audio_file", return_value=True), - ): - await ch._enqueue_raw(raw) + with ( + patch("EvoScientist.stt.transcribe_file", new=AsyncMock(return_value="你好")), + patch("EvoScientist.stt.is_audio_file", return_value=True), + ): + await ch._enqueue_raw(raw) - run_async(_run()) assert captured[0].text == "你好" # annotation should be removed after transcription assert captured[0].content_annotations == [] -def test_enqueue_raw_stt_disabled_skips_transcription(): +async def test_enqueue_raw_stt_disabled_skips_transcription(): """When stt_enabled=False, transcription is not called.""" from EvoScientist.channels.base import RawIncoming @@ -221,16 +214,14 @@ def test_enqueue_raw_stt_disabled_skips_transcription(): mock_transcribe = AsyncMock() - async def _run(): - with patch("EvoScientist.stt.transcribe_file", mock_transcribe): - await ch._enqueue_raw(raw) + with patch("EvoScientist.stt.transcribe_file", mock_transcribe): + await ch._enqueue_raw(raw) - run_async(_run()) mock_transcribe.assert_not_called() assert captured[0].text == "" -def test_enqueue_raw_stt_appends_to_existing_text(): +async def test_enqueue_raw_stt_appends_to_existing_text(): """Transcript is prepended before any existing caption text.""" from EvoScientist.channels.base import RawIncoming @@ -249,17 +240,15 @@ def test_enqueue_raw_stt_appends_to_existing_text(): timestamp=datetime.now(), ) - async def _run(): - with ( - patch( - "EvoScientist.stt.transcribe_file", - new=AsyncMock(return_value="hello world"), - ), - patch("EvoScientist.stt.is_audio_file", return_value=True), - ): - await ch._enqueue_raw(raw) + with ( + patch( + "EvoScientist.stt.transcribe_file", + new=AsyncMock(return_value="hello world"), + ), + patch("EvoScientist.stt.is_audio_file", return_value=True), + ): + await ch._enqueue_raw(raw) - run_async(_run()) assert captured[0].text == "hello world\ncaption text" diff --git a/tests/test_subagent_summarize.py b/tests/test_subagent_summarize.py index 331c0d5..d10d42e 100644 --- a/tests/test_subagent_summarize.py +++ b/tests/test_subagent_summarize.py @@ -17,7 +17,6 @@ from EvoScientist.channels.bus.message_bus import MessageBus from EvoScientist.channels.channel_manager import ChannelManager from EvoScientist.channels.consumer import InboundConsumer, _join_subagent_text from EvoScientist.stream.emitter import StreamEvent, StreamEventEmitter -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway from tests.fakes import StubChannel as _StubChannel from tests.stream_v3_fakes import ( @@ -82,7 +81,7 @@ class TestSubagentTextEmitter: class TestStreamAgentEventsSubagentText: """Verify sub-agent text chunks yield subagent_text events.""" - def test_subagent_text_emitted_for_subagent_chunks(self): + async def test_subagent_text_emitted_for_subagent_chunks(self): """When a sub-agent produces text, subagent_text events should appear.""" namespace = ("sub", "research") agent = FakeV3Agent( @@ -93,23 +92,23 @@ class TestStreamAgentEventsSubagentText: ], subagents=[FakeSubagent(namespace, "research-agent")], ) - events = collect_events(agent) + events = await collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] assert len(sa_text) == 1 assert "Sub-agent finding" in sa_text[0]["content"] # instance_id must be present and non-empty assert sa_text[0].get("instance_id"), "instance_id must be a non-empty string" - def test_subagent_text_not_emitted_for_main_agent(self): + async def test_subagent_text_not_emitted_for_main_agent(self): """Main agent text should produce 'text' events, not 'subagent_text'.""" agent = FakeV3Agent([message_delta("Main agent reply.")]) - events = collect_events(agent) + events = await collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] text_events = [e for e in events if e.get("type") == "text"] assert len(sa_text) == 0 assert len(text_events) == 1 - def test_multiple_subagent_text_chunks_all_emitted(self): + async def test_multiple_subagent_text_chunks_all_emitted(self): """Multiple text chunks from a sub-agent all yield subagent_text events.""" namespace = ("sub", "a") agent = FakeV3Agent( @@ -120,7 +119,7 @@ class TestStreamAgentEventsSubagentText: ], subagents=[FakeSubagent(namespace, "research-agent")], ) - events = collect_events(agent) + events = await collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] assert len(sa_text) == 3 combined = "".join(e["content"] for e in sa_text) @@ -132,7 +131,7 @@ class TestStreamAgentEventsSubagentText: assert len(ids) == 1, f"Expected 1 unique instance_id, got {ids}" assert all(e.get("instance_id") for e in sa_text) - def test_parallel_same_name_agents_get_distinct_instance_ids(self): + async def test_parallel_same_name_agents_get_distinct_instance_ids(self): """Two sub-agents with the same display name but different namespaces produce subagent_text events with different instance_id values. @@ -154,7 +153,7 @@ class TestStreamAgentEventsSubagentText: FakeSubagent(ns2, "research-agent"), ], ) - events = collect_events(agent) + events = await collect_events(agent) sa_text = [e for e in events if e.get("type") == "subagent_text"] assert len(sa_text) == 3 @@ -211,7 +210,7 @@ def _make_consumer(stream_events: list[dict], **kw): class TestConsumerSubagentTextFallback: """InboundConsumer should use sub-agent text as fallback when main agent is silent.""" - def test_subagent_text_used_when_no_final_content(self): + async def test_subagent_text_used_when_no_final_content(self): """When the main agent produces no text, sub-agent text becomes the response.""" events = [ { @@ -230,27 +229,24 @@ class TestConsumerSubagentTextFallback: ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="analyze papers", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="analyze papers", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "Found 3 relevant papers. Key insight: X is Y." - assert outbound.channel == "stub" + assert outbound.content == "Found 3 relevant papers. Key insight: X is Y." + assert outbound.channel == "stub" - await consumer.stop() - await task + await consumer.stop() + await task - _run(_test()) - - def test_final_content_takes_priority_over_subagent_text(self): + async def test_final_content_takes_priority_over_subagent_text(self): """When the main agent produces text, sub-agent text is ignored.""" events = [ { @@ -264,26 +260,23 @@ class TestConsumerSubagentTextFallback: ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="test", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "Here is my summary." + assert outbound.content == "Here is my summary." - await consumer.stop() - await task + await consumer.stop() + await task - _run(_test()) - - def test_duplicate_thinking_not_relayed_across_resume_rounds(self): + async def test_duplicate_thinking_not_relayed_across_resume_rounds(self): """Repeated thinking from resumed rounds should only be sent once.""" bus = MessageBus() mgr = ChannelManager(bus) @@ -329,30 +322,27 @@ class TestConsumerSubagentTextFallback: return_value={"answers": ["yes"], "status": "answered"} ) - async def _test(): - await bus.publish_inbound( - BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="analyze papers", - ) + await bus.publish_inbound( + BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="analyze papers", ) + ) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "final answer" - assert channel.send_thinking_message.await_count == 1 - call = channel.send_thinking_message.await_args_list[0] - assert call.args[1] == thinking.rstrip() + assert outbound.content == "final answer" + assert channel.send_thinking_message.await_count == 1 + call = channel.send_thinking_message.await_args_list[0] + assert call.args[1] == thinking.rstrip() - await consumer.stop() - await task + await consumer.stop() + await task - _run(_test()) - - def test_new_thinking_relayed_after_resume(self): + async def test_new_thinking_relayed_after_resume(self): """Genuinely different thinking in round 2 should be sent.""" bus = MessageBus() mgr = ChannelManager(bus) @@ -399,58 +389,52 @@ class TestConsumerSubagentTextFallback: return_value={"answers": ["yes"], "status": "answered"} ) - async def _test(): - await bus.publish_inbound( - BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="analyze papers", - ) + await bus.publish_inbound( + BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="analyze papers", ) + ) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "final answer" - assert channel.send_thinking_message.await_count == 2 - call1 = channel.send_thinking_message.await_args_list[0] - call2 = channel.send_thinking_message.await_args_list[1] - assert call1.args[1] == thinking_r1.rstrip() - assert call2.args[1] == thinking_r2.rstrip() + assert outbound.content == "final answer" + assert channel.send_thinking_message.await_count == 2 + call1 = channel.send_thinking_message.await_args_list[0] + call2 = channel.send_thinking_message.await_args_list[1] + assert call1.args[1] == thinking_r1.rstrip() + assert call2.args[1] == thinking_r2.rstrip() - await consumer.stop() - await task + await consumer.stop() + await task - _run(_test()) - - def test_no_response_fallback_when_both_empty(self): + async def test_no_response_fallback_when_both_empty(self): """When both final_content and subagent_text are empty, 'No response' is used.""" events = [ {"type": "done", "content": ""}, ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="test", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "No response" + assert outbound.content == "No response" - await consumer.stop() - await task + await consumer.stop() + await task - _run(_test()) - - def test_done_content_overrides_subagent_text(self): + async def test_done_content_overrides_subagent_text(self): """Done event with content takes priority over sub-agent text buffer.""" events = [ { @@ -463,24 +447,21 @@ class TestConsumerSubagentTextFallback: ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="test", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "Final summary from done event." + assert outbound.content == "Final summary from done event." - await consumer.stop() - await task - - _run(_test()) + await consumer.stop() + await task # ═══════════════════════════════════════════════════════════════════ @@ -560,7 +541,7 @@ class TestJoinSubagentText: class TestConsumerParallelSubagentFallback: """Consumer should group parallel sub-agent text by agent name.""" - def test_parallel_agents_grouped_with_attribution(self): + async def test_parallel_agents_grouped_with_attribution(self): """Multiple sub-agents produce grouped, attributed output.""" events = [ { @@ -585,27 +566,24 @@ class TestConsumerParallelSubagentFallback: ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="test", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert "[research]: Found papers. Key insight." in outbound.content - assert "[analysis]: Metric is high." in outbound.content + assert "[research]: Found papers. Key insight." in outbound.content + assert "[analysis]: Metric is high." in outbound.content - await consumer.stop() - await task + await consumer.stop() + await task - _run(_test()) - - def test_single_agent_no_attribution_prefix(self): + async def test_single_agent_no_attribution_prefix(self): """Single sub-agent fallback has no [name]: prefix.""" events = [ { @@ -618,31 +596,28 @@ class TestConsumerParallelSubagentFallback: ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="test", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - assert outbound.content == "Only agent." - assert "[research]" not in outbound.content + assert outbound.content == "Only agent." + assert "[research]" not in outbound.content - await consumer.stop() - await task - - _run(_test()) + await consumer.stop() + await task class TestConsumerSameNameInterleaved: """Two instances of the same agent type with interleaved chunks.""" - def test_same_name_interleaved_chunks_separated_by_instance_id(self): + async def test_same_name_interleaved_chunks_separated_by_instance_id(self): """Two research-agent instances with different instance_ids are properly separated. With the instance_id fix, chunks are keyed by instance_id so @@ -678,32 +653,29 @@ class TestConsumerSameNameInterleaved: ] consumer, bus = _make_consumer(events) - async def _test(): - msg = BusInbound( - channel="stub", - sender_id="u1", - chat_id="c1", - content="test", - ) - await bus.publish_inbound(msg) + msg = BusInbound( + channel="stub", + sender_id="u1", + chat_id="c1", + content="test", + ) + await bus.publish_inbound(msg) - task = asyncio.create_task(consumer.run()) - outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) + task = asyncio.create_task(consumer.run()) + outbound = await asyncio.wait_for(bus.consume_outbound(), timeout=5.0) - # Fixed: instances are now properly separated with numbered labels - assert ( - "[research-agent #1]: Instance-1 sentence A. Instance-1 sentence B." - in outbound.content - ) - assert ( - "[research-agent #2]: Instance-2 sentence X. Instance-2 sentence Y." - in outbound.content - ) + # Fixed: instances are now properly separated with numbered labels + assert ( + "[research-agent #1]: Instance-1 sentence A. Instance-1 sentence B." + in outbound.content + ) + assert ( + "[research-agent #2]: Instance-2 sentence X. Instance-2 sentence Y." + in outbound.content + ) - await consumer.stop() - await task - - _run(_test()) + await consumer.stop() + await task class TestDelegationPromptSummarize: diff --git a/tests/test_telegram_channel.py b/tests/test_telegram_channel.py index e30bafc..af1a461 100644 --- a/tests/test_telegram_channel.py +++ b/tests/test_telegram_channel.py @@ -4,7 +4,6 @@ import pytest from EvoScientist.channels.base import ChannelError from EvoScientist.channels.telegram.channel import TelegramChannel, TelegramConfig -from tests.conftest import run_async as _run class TestTelegramConfig: @@ -32,18 +31,18 @@ class TestTelegramChannel: assert channel.config is config assert channel._running is False - def test_start_raises_without_token(self): + async def test_start_raises_without_token(self): config = TelegramConfig(bot_token="") channel = TelegramChannel(config) with pytest.raises(ChannelError, match="bot token"): - _run(channel.start()) + await channel.start() - def test_stop_when_not_running(self): + async def test_stop_when_not_running(self): config = TelegramConfig(bot_token="test") channel = TelegramChannel(config) - _run(channel.stop()) + await channel.stop() - def test_send_returns_false_without_app(self): + async def test_send_returns_false_without_app(self): from EvoScientist.channels.base import OutboundMessage config = TelegramConfig(bot_token="test") @@ -54,5 +53,5 @@ class TestTelegramChannel: content="hello", metadata={"chat_id": "123"}, ) - result = _run(channel.send(msg)) + result = await channel.send(msg) assert result is False diff --git a/tests/test_threads_command.py b/tests/test_threads_command.py index d86cbbe..40836ec 100644 --- a/tests/test_threads_command.py +++ b/tests/test_threads_command.py @@ -4,7 +4,6 @@ from unittest.mock import MagicMock from rich.table import Table -from tests.conftest import run_async as _run from tests.fakes import FakeGraphGateway, FakeThreadStore @@ -24,15 +23,15 @@ def _ctx(**overrides): class TestThreadsCommand: - def test_empty_list_prints_message(self): + async def test_empty_list_prints_message(self): from EvoScientist.commands.implementation.session import ThreadsCommand ctx, ui = _ctx() - _run(ThreadsCommand().execute(ctx, [])) + await ThreadsCommand().execute(ctx, []) ui.append_system.assert_called_once() assert "No saved sessions" in ui.append_system.call_args.args[0] - def test_renders_table_with_current_marker(self): + async def test_renders_table_with_current_marker(self): from EvoScientist.commands.implementation.session import ThreadsCommand ctx, ui = _ctx(thread_id="current") @@ -54,7 +53,7 @@ class TestThreadsCommand: ] store = FakeThreadStore(threads=threads) ctx.graph_gateway = FakeGraphGateway(thread_store=store) - _run(ThreadsCommand().execute(ctx, [])) + await ThreadsCommand().execute(ctx, []) ui.mount_renderable.assert_called_once() table = ui.mount_renderable.call_args.args[0] assert isinstance(table, Table) @@ -64,7 +63,7 @@ class TestThreadsCommand: assert "/delete" in footer assert "/new" in footer - def test_footer_hint_suppressed_in_channel_mode(self): + async def test_footer_hint_suppressed_in_channel_mode(self): """Channels don't get the footer — keeps outbound text short.""" from EvoScientist.commands.implementation.session import ThreadsCommand @@ -80,10 +79,10 @@ class TestThreadsCommand: ] store = FakeThreadStore(threads=threads) ctx.graph_gateway = FakeGraphGateway(thread_store=store) - _run(ThreadsCommand().execute(ctx, [])) + await ThreadsCommand().execute(ctx, []) ui.append_system.assert_not_called() - def test_channel_mode_drops_model_column(self): + async def test_channel_mode_drops_model_column(self): """Non-interactive (channel) UIs get a narrower table.""" from EvoScientist.commands.implementation.session import ThreadsCommand @@ -99,7 +98,7 @@ class TestThreadsCommand: ] store = FakeThreadStore(threads=threads) ctx.graph_gateway = FakeGraphGateway(thread_store=store) - _run(ThreadsCommand().execute(ctx, [])) + await ThreadsCommand().execute(ctx, []) # Channel mode: no Model column. 4 columns: ID, Preview, Msgs, Last Used. table = ui.mount_renderable.call_args.args[0] column_headers = [col.header for col in table.columns] diff --git a/tests/test_tool_error_handler.py b/tests/test_tool_error_handler.py index 4d07dd7..7e2e95a 100644 --- a/tests/test_tool_error_handler.py +++ b/tests/test_tool_error_handler.py @@ -149,40 +149,34 @@ class TestWrapToolCallAsync: def setup_method(self): self.mw = ToolErrorHandlerMiddleware() - @staticmethod - def _run(coro): - from tests.conftest import run_async - - return run_async(coro) - - def test_success_passes_through(self): + async def test_success_passes_through(self): expected = ToolMessage(content="ok", tool_call_id="tc_001", name="t") async def handler(req): return expected req = _make_request() - result = self._run(self.mw.awrap_tool_call(req, handler)) + result = await self.mw.awrap_tool_call(req, handler) assert result is expected - def test_command_passes_through(self): + async def test_command_passes_through(self): cmd = Command(update={"messages": []}) async def handler(req): return cmd req = _make_request() - result = self._run(self.mw.awrap_tool_call(req, handler)) + result = await self.mw.awrap_tool_call(req, handler) assert result is cmd - def test_exception_returns_error_tool_message(self): + async def test_exception_returns_error_tool_message(self): async def handler(req): raise RuntimeError("MCP server timed out") req = _make_request("slow_tool", "tc_async") - result = self._run(self.mw.awrap_tool_call(req, handler)) + result = await self.mw.awrap_tool_call(req, handler) assert isinstance(result, ToolMessage) assert result.status == "error" @@ -190,21 +184,21 @@ class TestWrapToolCallAsync: assert result.name == "slow_tool" assert "MCP server timed out" in result.content - def test_exception_does_not_propagate(self): + async def test_exception_does_not_propagate(self): async def handler(req): raise ConnectionError("connection lost") req = _make_request() - result = self._run(self.mw.awrap_tool_call(req, handler)) + result = await self.mw.awrap_tool_call(req, handler) assert isinstance(result, ToolMessage) - def test_keyboard_interrupt_propagates(self): + async def test_keyboard_interrupt_propagates(self): async def handler(req): raise KeyboardInterrupt() req = _make_request() with pytest.raises(KeyboardInterrupt): - self._run(self.mw.awrap_tool_call(req, handler)) + await self.mw.awrap_tool_call(req, handler) # --------------------------------------------------------------------------- diff --git a/tests/test_tui_banner_position.py b/tests/test_tui_banner_position.py index 8c03c36..969476c 100644 --- a/tests/test_tui_banner_position.py +++ b/tests/test_tui_banner_position.py @@ -117,8 +117,8 @@ def _capture_app(monkeypatch) -> object: # --------------------------------------------------------------------------- -def test_clear_chat_resets_scroll_after_long_anchored_conversation( - monkeypatch, run_async +async def test_clear_chat_resets_scroll_after_long_anchored_conversation( + monkeypatch, ): """Repro of issue #301: clear after a long anchored stream → banner on top. @@ -128,104 +128,91 @@ def test_clear_chat_resets_scroll_after_long_anchored_conversation( the now-empty bottom of the previous content. """ - async def scenario(): - from textual.containers import VerticalScroll - from textual.widgets import Static + from textual.containers import VerticalScroll + from textual.widgets import Static - app = _capture_app(monkeypatch) - async with app.run_test(size=(80, 24)) as pilot: - await pilot.pause() - chat = app.query_one("#chat", VerticalScroll) - welcome = app.query_one("#welcome", Static) + app = _capture_app(monkeypatch) + async with app.run_test(size=(80, 24)) as pilot: + await pilot.pause() + chat = app.query_one("#chat", VerticalScroll) + welcome = app.query_one("#welcome", Static) - for i in range(80): - await chat.mount(Static(f"prior message {i}\n" * 2)) - await pilot.pause() - chat.scroll_end(animate=False) - await pilot.pause() - chat.anchor() - await pilot.pause() + for i in range(80): + await chat.mount(Static(f"prior message {i}\n" * 2)) + await pilot.pause() + chat.scroll_end(animate=False) + await pilot.pause() + chat.anchor() + await pilot.pause() - assert chat.scroll_y > 0, "precondition: viewport must be scrolled" + assert chat.scroll_y > 0, "precondition: viewport must be scrolled" - app.clear_chat() - app._append_system("New session: tid", style="green") - await pilot.pause() - await pilot.pause() + app.clear_chat() + app._append_system("New session: tid", style="green") + await pilot.pause() + await pilot.pause() - _assert_banner_at_top( - chat, welcome, label="after /new on long anchored convo" - ) - assert len(chat.children) == 2 - - run_async(scenario()) + _assert_banner_at_top(chat, welcome, label="after /new on long anchored convo") + assert len(chat.children) == 2 -def test_clear_chat_with_anchor_released_also_resets(monkeypatch, run_async): +async def test_clear_chat_with_anchor_released_also_resets(monkeypatch): """User scrolled up (anchor released) before /new → still lands at top.""" - async def scenario(): - from textual.containers import VerticalScroll - from textual.widgets import Static + from textual.containers import VerticalScroll + from textual.widgets import Static - app = _capture_app(monkeypatch) - async with app.run_test(size=(80, 24)) as pilot: - await pilot.pause() - chat = app.query_one("#chat", VerticalScroll) - welcome = app.query_one("#welcome", Static) + app = _capture_app(monkeypatch) + async with app.run_test(size=(80, 24)) as pilot: + await pilot.pause() + chat = app.query_one("#chat", VerticalScroll) + welcome = app.query_one("#welcome", Static) - for i in range(80): - await chat.mount(Static(f"msg {i}\n" * 2)) - await pilot.pause() - chat.anchor() - chat.scroll_to(y=80, animate=False) - await pilot.pause() + for i in range(80): + await chat.mount(Static(f"msg {i}\n" * 2)) + await pilot.pause() + chat.anchor() + chat.scroll_to(y=80, animate=False) + await pilot.pause() - app.clear_chat() - app._append_system("New session: tid", style="green") - await pilot.pause() - await pilot.pause() + app.clear_chat() + app._append_system("New session: tid", style="green") + await pilot.pause() + await pilot.pause() - _assert_banner_at_top( - chat, welcome, label="after /new with released anchor" - ) - assert len(chat.children) == 2 - - run_async(scenario()) + _assert_banner_at_top(chat, welcome, label="after /new with released anchor") + assert len(chat.children) == 2 -def test_clear_chat_short_conversation_anchored(monkeypatch, run_async): +async def test_clear_chat_short_conversation_anchored(monkeypatch): """Even with a short conversation, anchor + clear should not push banner down.""" - async def scenario(): - from textual.containers import VerticalScroll - from textual.widgets import Static + from textual.containers import VerticalScroll + from textual.widgets import Static - app = _capture_app(monkeypatch) - async with app.run_test(size=(80, 24)) as pilot: - await pilot.pause() - chat = app.query_one("#chat", VerticalScroll) - welcome = app.query_one("#welcome", Static) + app = _capture_app(monkeypatch) + async with app.run_test(size=(80, 24)) as pilot: + await pilot.pause() + chat = app.query_one("#chat", VerticalScroll) + welcome = app.query_one("#welcome", Static) - # Just enough content to overflow the viewport. - for i in range(30): - await chat.mount(Static(f"short msg {i}\n" * 2)) - await pilot.pause() - chat.scroll_end(animate=False) - chat.anchor() - await pilot.pause() + # Just enough content to overflow the viewport. + for i in range(30): + await chat.mount(Static(f"short msg {i}\n" * 2)) + await pilot.pause() + chat.scroll_end(animate=False) + chat.anchor() + await pilot.pause() - app.clear_chat() - app._append_system("New session: tid", style="green") - await pilot.pause() - await pilot.pause() + app.clear_chat() + app._append_system("New session: tid", style="green") + await pilot.pause() + await pilot.pause() - _assert_banner_at_top(chat, welcome, label="after /new on short convo") - - run_async(scenario()) + _assert_banner_at_top(chat, welcome, label="after /new on short convo") -def test_clear_chat_then_full_user_turn_keeps_banner_at_top(monkeypatch, run_async): +async def test_clear_chat_then_full_user_turn_keeps_banner_at_top(monkeypatch): """Repro of the user-reported scenario: clear → mount welcome banner → mount new-session → mount user message → mount assistant reply, in a normal-sized terminal where the resulting content fits in the viewport. @@ -237,69 +224,64 @@ def test_clear_chat_then_full_user_turn_keeps_banner_at_top(monkeypatch, run_asy wipe even when more widgets are mounted afterwards. """ - async def scenario(): - from textual.containers import VerticalScroll - from textual.widgets import Static + from textual.containers import VerticalScroll + from textual.widgets import Static - app = _capture_app(monkeypatch) - # Tall-ish terminal: welcome + a few messages must fit in the - # viewport, mirroring the user's manual-test setup. - async with app.run_test(size=(80, 40)) as pilot: - await pilot.pause() - chat = app.query_one("#chat", VerticalScroll) - welcome = app.query_one("#welcome", Static) + app = _capture_app(monkeypatch) + # Tall-ish terminal: welcome + a few messages must fit in the + # viewport, mirroring the user's manual-test setup. + async with app.run_test(size=(80, 40)) as pilot: + await pilot.pause() + chat = app.query_one("#chat", VerticalScroll) + welcome = app.query_one("#welcome", Static) - # Long conversation, then /new. - for i in range(80): - await chat.mount(Static(f"prior message {i}\n" * 2)) - await pilot.pause() - chat.scroll_end(animate=False) - chat.anchor() - await pilot.pause() + # Long conversation, then /new. + for i in range(80): + await chat.mount(Static(f"prior message {i}\n" * 2)) + await pilot.pause() + chat.scroll_end(animate=False) + chat.anchor() + await pilot.pause() - app.clear_chat() - # Render the actual banner (not the empty placeholder) and add - # the /new system message — this is exactly what - # ``start_new_session`` does after clearing. - app._render_welcome() - app._append_system("New session: tid", style="green") - await pilot.pause() - await pilot.pause() + app.clear_chat() + # Render the actual banner (not the empty placeholder) and add + # the /new system message — this is exactly what + # ``start_new_session`` does after clearing. + app._render_welcome() + app._append_system("New session: tid", style="green") + await pilot.pause() + await pilot.pause() - # User types "hello" — _run_turn mounts UserMessage then calls - # ``container.scroll_end(animate=False)`` (line 1305 in the - # real code). In a tall viewport this still lands at scroll_y - # == 0 because content fits. - from EvoScientist.cli.widgets.assistant_message import AssistantMessage - from EvoScientist.cli.widgets.user_message import UserMessage + # User types "hello" — _run_turn mounts UserMessage then calls + # ``container.scroll_end(animate=False)`` (line 1305 in the + # real code). In a tall viewport this still lands at scroll_y + # == 0 because content fits. + from EvoScientist.cli.widgets.assistant_message import AssistantMessage + from EvoScientist.cli.widgets.user_message import UserMessage - await chat.mount(UserMessage("hello")) - chat.scroll_end(animate=False) - await pilot.pause() + await chat.mount(UserMessage("hello")) + chat.scroll_end(animate=False) + await pilot.pause() - await chat.mount( - AssistantMessage( - "Hello. What research problem are we working on today?" - ) - ) - await pilot.pause() - await pilot.pause() + await chat.mount( + AssistantMessage("Hello. What research problem are we working on today?") + ) + await pilot.pause() + await pilot.pause() - _assert_banner_at_top( - chat, - welcome, - label=( - f"after full /new → user msg → reply " - f"(max={chat.max_scroll_y}, " - f"viewport={chat.scrollable_content_region.height}, " - f"content={chat.content_size.height})" - ), - ) - - run_async(scenario()) + _assert_banner_at_top( + chat, + welcome, + label=( + f"after full /new → user msg → reply " + f"(max={chat.max_scroll_y}, " + f"viewport={chat.scrollable_content_region.height}, " + f"content={chat.content_size.height})" + ), + ) -def test_short_turn_keeps_banner_at_top_after_layout_refresh(monkeypatch, run_async): +async def test_short_turn_keeps_banner_at_top_after_layout_refresh(monkeypatch): """Regression for the second symptom of issue #301: after a short user/assistant turn that fits in the viewport, end-of-stream ``_anchor_chat`` must NOT leave the chat anchored. @@ -313,83 +295,73 @@ def test_short_turn_keeps_banner_at_top_after_layout_refresh(monkeypatch, run_as overflows (``max_scroll_y > 0``). """ - async def scenario(): - from textual.containers import VerticalScroll - from textual.widgets import Static + from textual.containers import VerticalScroll + from textual.widgets import Static - from EvoScientist.cli.widgets.assistant_message import AssistantMessage - from EvoScientist.cli.widgets.user_message import UserMessage + from EvoScientist.cli.widgets.assistant_message import AssistantMessage + from EvoScientist.cli.widgets.user_message import UserMessage - app = _capture_app(monkeypatch) - # Tall terminal: welcome + a short exchange fits with room to spare, - # which is exactly the bug condition (content < viewport). - async with app.run_test(size=(80, 40)) as pilot: - await pilot.pause() - chat = app.query_one("#chat", VerticalScroll) - welcome = app.query_one("#welcome", Static) + app = _capture_app(monkeypatch) + # Tall terminal: welcome + a short exchange fits with room to spare, + # which is exactly the bug condition (content < viewport). + async with app.run_test(size=(80, 40)) as pilot: + await pilot.pause() + chat = app.query_one("#chat", VerticalScroll) + welcome = app.query_one("#welcome", Static) - await chat.mount(UserMessage("hi")) - await pilot.pause() - await chat.mount( - AssistantMessage("Hi. What are you looking to work on today?") - ) - await pilot.pause() + await chat.mount(UserMessage("hi")) + await pilot.pause() + await chat.mount(AssistantMessage("Hi. What are you looking to work on today?")) + await pilot.pause() - # End-of-stream re-anchor (matches _stream_with_widgets). - app._anchor_chat(chat) - await pilot.pause() + # End-of-stream re-anchor (matches _stream_with_widgets). + app._anchor_chat(chat) + await pilot.pause() - # Any subsequent mount triggers a layout refresh — this is when - # the compositor would push scroll_y negative without the fix. - # In production this happens via Markdown re-renders, status-bar - # updates, the system "usage" line, etc. - await chat.mount(Static("trailing line\n")) - await pilot.pause() - await pilot.pause() + # Any subsequent mount triggers a layout refresh — this is when + # the compositor would push scroll_y negative without the fix. + # In production this happens via Markdown re-renders, status-bar + # updates, the system "usage" line, etc. + await chat.mount(Static("trailing line\n")) + await pilot.pause() + await pilot.pause() - _assert_banner_at_top( - chat, welcome, label="after short turn + trailing mount" - ) - - run_async(scenario()) + _assert_banner_at_top(chat, welcome, label="after short turn + trailing mount") -def test_long_turn_keeps_viewport_pinned_to_bottom(monkeypatch, run_async): +async def test_long_turn_keeps_viewport_pinned_to_bottom(monkeypatch): """When the conversation overflows, ``_anchor_chat`` must still engage the anchor so streaming output remains visible. The issue #301 fix only suppresses anchoring when content fits — long content must continue to behave as before. """ - async def scenario(): - from textual.containers import VerticalScroll - from textual.widgets import Static + from textual.containers import VerticalScroll + from textual.widgets import Static - app = _capture_app(monkeypatch) - async with app.run_test(size=(80, 24)) as pilot: - await pilot.pause() - chat = app.query_one("#chat", VerticalScroll) + app = _capture_app(monkeypatch) + async with app.run_test(size=(80, 24)) as pilot: + await pilot.pause() + chat = app.query_one("#chat", VerticalScroll) - for i in range(50): - await chat.mount(Static(f"prior message {i}\n" * 2)) - await pilot.pause() + for i in range(50): + await chat.mount(Static(f"prior message {i}\n" * 2)) + await pilot.pause() - app._anchor_chat(chat) - await pilot.pause() - assert chat.scroll_y == chat.max_scroll_y, ( - "long content must anchor to bottom after _anchor_chat" - ) + app._anchor_chat(chat) + await pilot.pause() + assert chat.scroll_y == chat.max_scroll_y, ( + "long content must anchor to bottom after _anchor_chat" + ) - # Trailing mount must keep the viewport pinned to the new bottom. - await chat.mount(Static("trailing line\n")) - await pilot.pause() - await pilot.pause() - assert chat.scroll_y == chat.max_scroll_y, ( - "anchored viewport must follow new bottom after trailing mount" - ) - assert chat.scroll_y > 0, "long content must have positive scroll_y" - - run_async(scenario()) + # Trailing mount must keep the viewport pinned to the new bottom. + await chat.mount(Static("trailing line\n")) + await pilot.pause() + await pilot.pause() + assert chat.scroll_y == chat.max_scroll_y, ( + "anchored viewport must follow new bottom after trailing mount" + ) + assert chat.scroll_y > 0, "long content must have positive scroll_y" # --------------------------------------------------------------------------- diff --git a/tests/test_tui_command_sync.py b/tests/test_tui_command_sync.py index 747ec94..43015cd 100644 --- a/tests/test_tui_command_sync.py +++ b/tests/test_tui_command_sync.py @@ -5,7 +5,6 @@ from types import SimpleNamespace import pytest from EvoScientist.commands.base import ChannelRuntime, CommandContext -from tests.conftest import run_async as _run pytest.importorskip("textual") @@ -43,7 +42,7 @@ class _StubApp: self.refresh_calls.append(reset_streaming_text) -def test_sync_tui_command_completion_adopts_agent_swap(monkeypatch): +async def test_sync_tui_command_completion_adopts_agent_swap(monkeypatch): import EvoScientist.cli.tui_interactive as tui_mod from EvoScientist import EvoScientist as evosci_mod @@ -63,7 +62,7 @@ def test_sync_tui_command_completion_adopts_agent_swap(monkeypatch): monkeypatch.setattr(tui_mod, "_channels_is_running", lambda: True) app._channel_runtime.bind("old-agent", "old-thread") - _run(tui_mod._sync_tui_command_completion(app, ctx, "old-agent", cmd)) + await tui_mod._sync_tui_command_completion(app, ctx, "old-agent", cmd) assert app._agent_loader.adopt_calls == ["new-agent"] assert app.model_updates == [("gpt-5.5", "openai")] @@ -72,7 +71,7 @@ def test_sync_tui_command_completion_adopts_agent_swap(monkeypatch): assert app._channel_runtime.thread_id == "thread-1" -def test_sync_tui_command_completion_refreshes_without_agent_swap(monkeypatch): +async def test_sync_tui_command_completion_refreshes_without_agent_swap(monkeypatch): import EvoScientist.cli.tui_interactive as tui_mod app = _StubApp() @@ -85,14 +84,16 @@ def test_sync_tui_command_completion_refreshes_without_agent_swap(monkeypatch): monkeypatch.setattr(tui_mod, "_channels_is_running", lambda: False) - _run(tui_mod._sync_tui_command_completion(app, ctx, "same-agent", cmd)) + await tui_mod._sync_tui_command_completion(app, ctx, "same-agent", cmd) assert app._agent_loader.adopt_calls == [] assert app.model_updates == [] assert app.refresh_calls == [True] -def test_sync_tui_rebinds_runtime_on_thread_rotation_without_agent_swap(monkeypatch): +async def test_sync_tui_rebinds_runtime_on_thread_rotation_without_agent_swap( + monkeypatch, +): """Regression: ``/new`` and ``/resume`` rotate ``app._conversation_tid`` without swapping the agent. The runtime must still pick up the new thread id so the bus contract stays consistent with serve mode.""" @@ -111,7 +112,7 @@ def test_sync_tui_rebinds_runtime_on_thread_rotation_without_agent_swap(monkeypa monkeypatch.setattr(tui_mod, "_channels_is_running", lambda: True) - _run(tui_mod._sync_tui_command_completion(app, ctx, "same-agent", cmd)) + await tui_mod._sync_tui_command_completion(app, ctx, "same-agent", cmd) assert app._channel_runtime.agent == "same-agent" assert app._channel_runtime.thread_id == "rotated-thread" diff --git a/tests/test_tui_widgets.py b/tests/test_tui_widgets.py index cdae349..9015968 100644 --- a/tests/test_tui_widgets.py +++ b/tests/test_tui_widgets.py @@ -10,14 +10,16 @@ import importlib import unittest from unittest.mock import AsyncMock +import pytest + # --------------------------------------------------------------------------- # Textual might not be installed — skip entire module if missing # --------------------------------------------------------------------------- _has_textual = importlib.util.find_spec("textual") is not None -@unittest.skipUnless(_has_textual, "textual not installed") -class TestLoadingWidget(unittest.TestCase): +@pytest.mark.skipif(not _has_textual, reason="textual not installed") +class TestLoadingWidget: """LoadingWidget construction and attributes.""" def test_construction(self): @@ -43,7 +45,7 @@ class TestLoadingWidget(unittest.TestCase): assert w._frame == 1 assert w._elapsed == 0.1 - def test_cleanup_stops_timer_and_removes(self): + async def test_cleanup_stops_timer_and_removes(self): from EvoScientist.cli.widgets.loading_widget import LoadingWidget class _Timer: @@ -58,9 +60,7 @@ class TestLoadingWidget(unittest.TestCase): w._timer_handle = timer w.remove = AsyncMock() - from tests.conftest import run_async - - run_async(w.cleanup()) + await w.cleanup() assert timer.stopped is True assert w._timer_handle is None diff --git a/tests/test_wechat_channel.py b/tests/test_wechat_channel.py index ce8d7eb..772c38e 100644 --- a/tests/test_wechat_channel.py +++ b/tests/test_wechat_channel.py @@ -20,7 +20,6 @@ from EvoScientist.channels.wechat.crypto import ( _pkcs7_unpad, parse_xml, ) -from tests.conftest import run_async as _run # ── Config tests ────────────────────────────────────────────────── @@ -83,36 +82,36 @@ class TestWeChatChannelInit: channel = WeChatChannel(config, backend="wechatmp") assert channel._backend == "wechatmp" - def test_start_raises_without_corp_id(self): + async def test_start_raises_without_corp_id(self): config = WeComConfig(corp_id="", agent_id="1", secret="s") channel = WeChatChannel(config, backend="wecom") with pytest.raises(ChannelError, match="corp_id"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_secret(self): + async def test_start_raises_without_secret(self): config = WeComConfig(corp_id="corp", agent_id="1", secret="") channel = WeChatChannel(config, backend="wecom") with pytest.raises(ChannelError, match="secret"): - _run(channel.start()) + await channel.start() - def test_start_raises_without_agent_id(self): + async def test_start_raises_without_agent_id(self): config = WeComConfig(corp_id="corp", agent_id="", secret="s") channel = WeChatChannel(config, backend="wecom") with pytest.raises(ChannelError, match="agent_id"): - _run(channel.start()) + await channel.start() - def test_start_raises_mp_without_app_id(self): + async def test_start_raises_mp_without_app_id(self): config = WeChatMPConfig(app_id="", app_secret="s") channel = WeChatChannel(config, backend="wechatmp") with pytest.raises(ChannelError, match="app_id"): - _run(channel.start()) + await channel.start() - def test_stop_when_not_running(self): + async def test_stop_when_not_running(self): config = WeComConfig(corp_id="c", agent_id="1", secret="s") channel = WeChatChannel(config, backend="wecom") - _run(channel.stop()) # Should not raise + await channel.stop() # Should not raise - def test_send_returns_false_without_client(self): + async def test_send_returns_false_without_client(self): from EvoScientist.channels.base import OutboundMessage config = WeComConfig(corp_id="c", agent_id="1", secret="s") @@ -123,7 +122,7 @@ class TestWeChatChannelInit: content="hello", metadata={"chat_id": "user1"}, ) - result = _run(channel.send(msg)) + result = await channel.send(msg) assert result is False @@ -324,144 +323,123 @@ class TestMessageProcessing: ) return WeChatChannel(config, backend="wecom") - def test_text_message_queued(self): + async def test_text_message_queued(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "text", - "Content": "Hello!", - "FromUserName": "user1", - "ToUserName": "bot", - "MsgId": "100", - "CreateTime": str(int(time.time())), - } - ) - # Check message was enqueued - assert not channel._queue.empty() - msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) - assert msg.content == "Hello!" - assert msg.sender_id == "user1" - assert msg.channel == "wechat" + await channel._process_message( + { + "MsgType": "text", + "Content": "Hello!", + "FromUserName": "user1", + "ToUserName": "bot", + "MsgId": "100", + "CreateTime": str(int(time.time())), + } + ) + # Check message was enqueued + assert not channel._queue.empty() + msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) + assert msg.content == "Hello!" + assert msg.sender_id == "user1" + assert msg.channel == "wechat" - _run(_test()) - - def test_location_message(self): + async def test_location_message(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "location", - "Location_X": "39.9", - "Location_Y": "116.4", - "Label": "Beijing", - "FromUserName": "user1", - "ToUserName": "bot", - "MsgId": "101", - "CreateTime": str(int(time.time())), - } - ) - msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) - assert "Beijing" in msg.content - assert "39.9" in msg.content + await channel._process_message( + { + "MsgType": "location", + "Location_X": "39.9", + "Location_Y": "116.4", + "Label": "Beijing", + "FromUserName": "user1", + "ToUserName": "bot", + "MsgId": "101", + "CreateTime": str(int(time.time())), + } + ) + msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) + assert "Beijing" in msg.content + assert "39.9" in msg.content - _run(_test()) - - def test_voice_recognition(self): + async def test_voice_recognition(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "voice", - "Recognition": "你好世界", - "FromUserName": "user1", - "ToUserName": "bot", - "MsgId": "102", - "CreateTime": str(int(time.time())), - } - ) - msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) - assert "你好世界" in msg.content + await channel._process_message( + { + "MsgType": "voice", + "Recognition": "你好世界", + "FromUserName": "user1", + "ToUserName": "bot", + "MsgId": "102", + "CreateTime": str(int(time.time())), + } + ) + msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) + assert "你好世界" in msg.content - _run(_test()) - - def test_link_message(self): + async def test_link_message(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "link", - "Title": "Test Link", - "Description": "A description", - "Url": "https://example.com", - "FromUserName": "user1", - "ToUserName": "bot", - "MsgId": "103", - "CreateTime": str(int(time.time())), - } - ) - msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) - assert "Test Link" in msg.content - assert "https://example.com" in msg.content + await channel._process_message( + { + "MsgType": "link", + "Title": "Test Link", + "Description": "A description", + "Url": "https://example.com", + "FromUserName": "user1", + "ToUserName": "bot", + "MsgId": "103", + "CreateTime": str(int(time.time())), + } + ) + msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) + assert "Test Link" in msg.content + assert "https://example.com" in msg.content - _run(_test()) - - def test_subscribe_event(self): + async def test_subscribe_event(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "event", - "Event": "subscribe", - "FromUserName": "user1", - "ToUserName": "bot", - "MsgId": "", - "CreateTime": str(int(time.time())), - } - ) - msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) - assert "关注" in msg.content + await channel._process_message( + { + "MsgType": "event", + "Event": "subscribe", + "FromUserName": "user1", + "ToUserName": "bot", + "MsgId": "", + "CreateTime": str(int(time.time())), + } + ) + msg = await asyncio.wait_for(channel._queue.get(), timeout=1.0) + assert "关注" in msg.content - _run(_test()) - - def test_unsubscribe_ignored(self): + async def test_unsubscribe_ignored(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "event", - "Event": "unsubscribe", - "FromUserName": "user1", - "ToUserName": "bot", - "MsgId": "", - "CreateTime": str(int(time.time())), - } - ) - assert channel._queue.empty() + await channel._process_message( + { + "MsgType": "event", + "Event": "unsubscribe", + "FromUserName": "user1", + "ToUserName": "bot", + "MsgId": "", + "CreateTime": str(int(time.time())), + } + ) + assert channel._queue.empty() - _run(_test()) - - def test_empty_message_ignored(self): + async def test_empty_message_ignored(self): channel = self._make_channel() - async def _test(): - await channel._process_message( - { - "MsgType": "text", - "Content": "", - "FromUserName": "", - "ToUserName": "bot", - } - ) - assert channel._queue.empty() - - _run(_test()) + await channel._process_message( + { + "MsgType": "text", + "Content": "", + "FromUserName": "", + "ToUserName": "bot", + } + ) + assert channel._queue.empty() # ── Registration test ───────────────────────────────────────────── diff --git a/uv.lock b/uv.lock index 5ad4352..947e9fa 100644 --- a/uv.lock +++ b/uv.lock @@ -992,6 +992,7 @@ dev = [ { name = "build" }, { name = "pre-commit" }, { name = "pytest" }, + { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-timeout" }, { name = "ruff" }, @@ -1032,6 +1033,7 @@ dev = [ { name = "build" }, { name = "pre-commit" }, { name = "pytest" }, + { name = "pytest-asyncio" }, { name = "pytest-cov" }, { name = "pytest-timeout" }, { name = "ruff" }, @@ -1075,6 +1077,7 @@ requires-dist = [ { name = "pycryptodome", marker = "extra == 'all-channels'", specifier = ">=3.20" }, { name = "pycryptodome", marker = "extra == 'wechat'", specifier = ">=3.20" }, { name = "pytest", marker = "extra == 'dev'", specifier = ">=8.0" }, + { name = "pytest-asyncio", marker = "extra == 'dev'", specifier = ">=1.0" }, { name = "pytest-cov", marker = "extra == 'dev'", specifier = ">=5.0" }, { name = "pytest-timeout", marker = "extra == 'dev'", specifier = ">=2.4" }, { name = "python-dotenv", specifier = ">=1.0" }, @@ -1103,6 +1106,7 @@ dev = [ { name = "build", specifier = ">=1.0" }, { name = "pre-commit", specifier = ">=3.5.0" }, { name = "pytest", specifier = ">=8.0" }, + { name = "pytest-asyncio", specifier = ">=1.0" }, { name = "pytest-cov", specifier = ">=5.0" }, { name = "pytest-timeout", specifier = ">=2.4" }, { name = "ruff", specifier = ">=0.5" }, @@ -3503,6 +3507,19 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/d4/24/a372aaf5c9b7208e7112038812994107bc65a84cd00e0354a88c2c77a617/pytest-9.0.3-py3-none-any.whl", hash = "sha256:2c5efc453d45394fdd706ade797c0a81091eccd1d6e4bccfcd476e2b8e0ab5d9", size = 375249, upload-time = "2026-04-07T17:16:16.13Z" }, ] +[[package]] +name = "pytest-asyncio" +version = "1.4.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "pytest" }, + { name = "typing-extensions", marker = "python_full_version < '3.13'" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/43/7c/d36d04db312ecf4298932ef77e6e4a9e8ad017906e24e34f0b0c361a2473/pytest_asyncio-1.4.0.tar.gz", hash = "sha256:c6c0d2259945122819f171a32ecea2c349ead889ee28176caaf492143424be42", size = 58514, upload-time = "2026-05-26T09:56:04.083Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/03/e2/08a497ef684b88559c9cc5f4ad53a37e7b99e727094a86d6ea32536d5d3c/pytest_asyncio-1.4.0-py3-none-any.whl", hash = "sha256:933ca923a23075a87fb7070c0ec272a6848489824d887c85c812670932835aa1", size = 16930, upload-time = "2026-05-26T09:56:02.576Z" }, +] + [[package]] name = "pytest-cov" version = "7.1.0" @@ -4178,11 +4195,11 @@ wheels = [ [[package]] name = "soupsieve" -version = "2.8.3" +version = "2.8.4" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/7b/ae/2d9c981590ed9999a0d91755b47fc74f74de286b0f5cee14c9269041e6c4/soupsieve-2.8.3.tar.gz", hash = "sha256:3267f1eeea4251fb42728b6dfb746edc9acaffc4a45b27e19450b676586e8349", size = 118627, upload-time = "2026-01-20T04:27:02.457Z" } +sdist = { url = "https://files.pythonhosted.org/packages/47/2c/0a5f6f8ee0d5589e48c7640213ed5175d52cf540a06725b628cc1a45d6ce/soupsieve-2.8.4.tar.gz", hash = "sha256:e121fd02e975c695e4e9e8774a5ee35d74714b59307868dcc5319ad2d9e3328e", size = 121110, upload-time = "2026-05-24T13:55:57.154Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/46/2c/1462b1d0a634697ae9e55b3cecdcb64788e8b7d63f54d923fcd0bb140aed/soupsieve-2.8.3-py3-none-any.whl", hash = "sha256:ed64f2ba4eebeab06cc4962affce381647455978ffc1e36bb79a545b91f45a95", size = 37016, upload-time = "2026-01-20T04:27:01.012Z" }, + { url = "https://files.pythonhosted.org/packages/5e/f5/0c41cb68dcae6b7de4fac4188a3a9589e21fb31df21ea3a2e888db95e6c9/soupsieve-2.8.4-py3-none-any.whl", hash = "sha256:e7e6b0769c8f51ed59acab6e994b00621096cfb1c640a7509295987388fbaf65", size = 37304, upload-time = "2026-05-24T13:55:55.406Z" }, ] [[package]]