diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 9e04633..4615ccc 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -28,6 +28,6 @@ jobs: python-version: ${{ matrix.python-version }} cache-dependency-glob: "**/pyproject.toml" - name: Install dependencies - run: uv sync --dev + run: uv sync --dev --extra all-channels - name: Run pytest run: uv run pytest -v --timeout=30 diff --git a/EvoScientist/channels/base.py b/EvoScientist/channels/base.py index 1b9502c..0003882 100644 --- a/EvoScientist/channels/base.py +++ b/EvoScientist/channels/base.py @@ -13,7 +13,8 @@ from collections import OrderedDict from collections.abc import AsyncIterator, Awaitable, Callable from collections.abc import Callable as CallableABC from dataclasses import dataclass, field -from datetime import datetime +from datetime import UTC, datetime +from email.utils import parsedate_to_datetime from pathlib import Path from typing import Any @@ -746,68 +747,144 @@ class Channel(TraceMixin, ChannelPlugin, ABC): # ── Send retry abstraction ────────────────────────────────────── - _non_retryable_patterns: tuple[str, ...] = () + # HTTP status codes that should never be retried. Listed explicitly + # rather than as a 4xx range: 408 and 425 are retryable by definition and + # 429 is handled by the rate-limit path. + _non_retryable_status_codes: tuple[int, ...] = (400, 401, 403, 404) + + # Structured SDK error codes that should never be retried (e.g. Slack invalid_auth) + # Channel-specific message patterns (e.g. Feishu 10003, DingTalk 40014) are handled + # via _non_retryable_patterns in respective channel subclasses. + _non_retryable_error_codes: tuple[str, ...] = ( + "invalid_auth", + "invalid_token", + "expired_token", + "token_expired", + "token_revoked", + "account_inactive", + "not_authed", + "no_permission", + "missing_scope", + ) + + _non_retryable_patterns: tuple[str, ...] = ( + "unauthorized", + "forbidden", + "permission denied", + "invalid token", + "invalid api key", + "authentication failed", + ) _rate_limit_patterns: tuple[str, ...] = ("429", "ratelimit") _rate_limit_delay: float = 1.0 def _extract_retry_after(self, exc: Exception) -> float | None: """Extract retry-wait seconds from an exception. - Returns ``None`` to signal that the error is **not retryable**. + Returns a retry delay in seconds, or ``None`` when the error is + explicitly non-retryable. Pipeline: - 1. SDK-provided ``retry_after`` attribute (Telegram / Slack SDKs). - 2. HTTP ``Retry-After`` header via :meth:`_parse_retry_after_header`. - 3. Non-retryable pattern match → ``None``. - 4. Rate-limit pattern match → ``_rate_limit_delay``. - 5. Default ``1.0`` s (generic transient-error retry). + 1. Non-retryable detection → ``None``. Evaluates HTTP status codes + (e.g. 401, 403), structured SDK error codes (e.g. Slack + ``"invalid_auth"``), and message pattern matching + (e.g. ``"unauthorized"``, ``"forbidden"``). + 2. Server-supplied delay via :meth:`_extract_retry_delay` + (httpx ``Retry-After``; channels override for their SDK). + 3. Rate-limit pattern match → ``_rate_limit_delay``. + 4. Default ``1.0`` s for generic transient errors. Channels can customize behavior declaratively via class attributes - ``_non_retryable_patterns``, ``_rate_limit_patterns``, and - ``_rate_limit_delay``, or override this method entirely. + ``_non_retryable_patterns``, ``_rate_limit_patterns``, + ``_non_retryable_status_codes``, ``_non_retryable_error_codes``, + and ``_rate_limit_delay``, or override this method entirely. """ - # 1. SDK retry_after attribute - retry = getattr(exc, "retry_after", None) - if retry is not None: - return float(retry) + # 1. Non-retryable detection: evaluate status codes, structured SDK + # error codes, and message patterns independently. + status_code = self._extract_status_code(exc) + if status_code is not None and status_code in self._non_retryable_status_codes: + return None - # 2. HTTP Retry-After header - header_val = self._parse_retry_after_header(exc) - if header_val is not None: - return header_val + sdk_error = self._extract_sdk_error_code(exc) + if sdk_error is not None and sdk_error in self._non_retryable_error_codes: + return None msg = str(exc).lower() - - # 3. Non-retryable patterns if self._non_retryable_patterns and any( p in msg for p in self._non_retryable_patterns ): return None - # 4. Rate-limit patterns + # 2. Server-supplied delay + delay = self._extract_retry_delay(exc) + if delay is not None: + return delay + + # 3. Rate-limit patterns if self._rate_limit_patterns and any( p in msg for p in self._rate_limit_patterns ): return self._rate_limit_delay - # 5. Default + # 4. Default: transient error, retry with the standard delay return 1.0 - def _parse_retry_after_header(self, exc: Exception) -> float | None: - """Try to extract a ``Retry-After`` value from an HTTP response.""" - resp = getattr(exc, "response", None) - if resp is None: - return None - headers = getattr(resp, "headers", None) - if not headers: - return None - raw = headers.get("Retry-After") or headers.get("retry-after") - if raw is None: - return None + def _extract_status_code(self, exc: Exception) -> int | None: + """Extract HTTP status from an httpx error. + + Channels with other SDKs (e.g. ``SlackChannel``, ``DiscordChannel``) + override this method. + """ + import httpx + + if isinstance(exc, httpx.HTTPStatusError): + return exc.response.status_code + + return None + + def _extract_sdk_error_code(self, exc: Exception) -> str | None: + """Extract structured SDK error code string from an exception. + + Plain HTTP carries no structured error code by default (returns ``None``). + Subclasses with specialized SDKs (e.g. ``SlackChannel``, ``DiscordChannel``) + override this method. + """ + return None + + def _extract_retry_delay(self, exc: Exception) -> float | None: + """Retry delay the server asked for, in seconds, or ``None``. + + Base implementation reads the ``Retry-After`` header of an httpx + error. Channels whose SDK reports the delay differently + (``SlackChannel``, ``TelegramChannel``, ``DiscordChannel``) override. + """ + import httpx + + if isinstance(exc, httpx.HTTPStatusError): + raw = exc.response.headers.get("retry-after") + if raw is not None: + return self._parse_retry_after(raw) + return None + + @staticmethod + def _parse_retry_after(raw: str) -> float | None: + """Convert a ``Retry-After`` header value to seconds. + + RFC 9110 allows either delay-seconds or an HTTP-date; a date is + returned as the non-negative number of seconds until it. Unparseable + values yield ``None`` so the caller can fall back to its own delay. + """ try: return float(raw) + except ValueError: + pass + try: + when = parsedate_to_datetime(raw) except (ValueError, TypeError): return None + if when.tzinfo is None: + when = when.replace(tzinfo=UTC) + return max(0.0, (when - datetime.now(UTC)).total_seconds()) async def _send_with_retry( self, diff --git a/EvoScientist/channels/dingtalk/channel.py b/EvoScientist/channels/dingtalk/channel.py index e875741..252f8c9 100644 --- a/EvoScientist/channels/dingtalk/channel.py +++ b/EvoScientist/channels/dingtalk/channel.py @@ -35,7 +35,11 @@ class DingTalkChannel(Channel, WebSocketMixin, TokenMixin): capabilities = DINGTALK_CAPS name = "dingtalk" _ready_attrs = ("_http_client", "_access_token") - _non_retryable_patterns = ("invalidauthentication", "forbidden", "40014") + _non_retryable_patterns = ( + *Channel._non_retryable_patterns, + "invalidauthentication", + "40014", + ) _mention_pattern = r"@\S+\s*" _mention_strip_count = 1 diff --git a/EvoScientist/channels/discord/channel.py b/EvoScientist/channels/discord/channel.py index d86611d..620979b 100644 --- a/EvoScientist/channels/discord/channel.py +++ b/EvoScientist/channels/discord/channel.py @@ -208,6 +208,25 @@ class DiscordChannel(Channel): return str(self._client.user.id) return None + # ── Retry error code extraction (override base) ───────────────── + + def _extract_status_code(self, exc: Exception) -> int | None: + """Extract HTTP status code from discord.HTTPException or fallback to base.""" + import discord + + if isinstance(exc, discord.HTTPException): + return exc.status + return super()._extract_status_code(exc) + + def _extract_retry_delay(self, exc: Exception) -> float | None: + """Honor ``discord.RateLimited``, raised when a 429 exceeds + ``max_ratelimit_timeout`` and discord.py stops retrying internally.""" + import discord + + if isinstance(exc, discord.RateLimited): + return exc.retry_after + return super()._extract_retry_delay(exc) + # ── Inbound ───────────────────────────────────────────────────── async def _on_message(self, message) -> None: diff --git a/EvoScientist/channels/email/channel.py b/EvoScientist/channels/email/channel.py index 6d5e426..783359f 100644 --- a/EvoScientist/channels/email/channel.py +++ b/EvoScientist/channels/email/channel.py @@ -73,7 +73,12 @@ class EmailChannel(Channel, PollingMixin): name = "email" capabilities = EMAIL_CAPS - _non_retryable_patterns = ("auth", "login", "credential") + _non_retryable_patterns = ( + *Channel._non_retryable_patterns, + "auth", + "login", + "credential", + ) def __init__(self, config: EmailConfig): super().__init__(config) diff --git a/EvoScientist/channels/feishu/channel.py b/EvoScientist/channels/feishu/channel.py index 6901488..998423d 100644 --- a/EvoScientist/channels/feishu/channel.py +++ b/EvoScientist/channels/feishu/channel.py @@ -257,6 +257,7 @@ class FeishuChannel(Channel, WebhookMixin, TokenMixin): name = "feishu" _ready_attrs = ("_http_client", "_access_token") _non_retryable_patterns = ( + *Channel._non_retryable_patterns, "app_access_token is empty", # invalid credentials "10003", # invalid app_id "10014", # invalid app_secret diff --git a/EvoScientist/channels/qq/channel.py b/EvoScientist/channels/qq/channel.py index 1f13eba..7b5cc15 100644 --- a/EvoScientist/channels/qq/channel.py +++ b/EvoScientist/channels/qq/channel.py @@ -116,7 +116,7 @@ class QQChannel(Channel): capabilities = QQ_CAPS _ready_attrs = ("_client", "_running") - _non_retryable_patterns = () + _non_retryable_patterns = Channel._non_retryable_patterns _mention_pattern = r"@\S+\s*" _mention_strip_count = 1 _markdown_fallback_exc_types: ClassVar[tuple[type[Exception], ...]] = ( diff --git a/EvoScientist/channels/signal/channel.py b/EvoScientist/channels/signal/channel.py index a7aa31e..75163c3 100644 --- a/EvoScientist/channels/signal/channel.py +++ b/EvoScientist/channels/signal/channel.py @@ -32,7 +32,11 @@ class SignalChannel(Channel): name = "signal" capabilities = SIGNAL_CAPS - _non_retryable_patterns = ("unregistered", "auth") + _non_retryable_patterns = ( + *Channel._non_retryable_patterns, + "unregistered", + "auth", + ) def __init__(self, config: SignalConfig): super().__init__(config) diff --git a/EvoScientist/channels/slack/channel.py b/EvoScientist/channels/slack/channel.py index ff2a031..490b251 100644 --- a/EvoScientist/channels/slack/channel.py +++ b/EvoScientist/channels/slack/channel.py @@ -19,6 +19,13 @@ class SlackConfig(BaseChannelConfig): text_chunk_limit: int = 4096 +def _slack_response_types() -> tuple[type, ...]: + from slack_sdk.web.async_slack_response import AsyncSlackResponse + from slack_sdk.web.slack_response import SlackResponse + + return (SlackResponse, AsyncSlackResponse) + + class SlackChannel(Channel): """Slack channel using slack-sdk Socket Mode.""" @@ -195,6 +202,45 @@ class SlackChannel(Channel): def _get_bot_identifier(self) -> str | None: return getattr(self, "_bot_user_id", None) + # ── Retry error code extraction (override base) ───────────────── + + def _extract_status_code(self, exc: Exception) -> int | None: + """Extract HTTP status code from SlackApiError or fallback to base.""" + from slack_sdk.errors import SlackApiError + + if isinstance(exc, SlackApiError) and isinstance( + exc.response, _slack_response_types() + ): + return exc.response.status_code + return super()._extract_status_code(exc) + + def _extract_sdk_error_code(self, exc: Exception) -> str | None: + """Extract structured error code string from SlackApiError.""" + from slack_sdk.errors import SlackApiError + + if isinstance(exc, SlackApiError) and isinstance( + exc.response, _slack_response_types() + ): + error = exc.response.get("error") + return error.lower() if isinstance(error, str) else None + return super()._extract_sdk_error_code(exc) + + def _extract_retry_delay(self, exc: Exception) -> float | None: + """Read Slack's ``Retry-After`` header from a SlackApiError. + + ``SlackResponse.headers`` is a plain ``dict`` whose key casing depends + on the HTTP client, so match the key case-insensitively (the same + approach slack_sdk's own ``RateLimitErrorRetryHandler`` takes). + """ + from slack_sdk.errors import SlackApiError + + if isinstance(exc, SlackApiError): + for key, raw in exc.response.headers.items(): + if key.lower() == "retry-after": + return self._parse_retry_after(raw) + return None + return super()._extract_retry_delay(exc) + # ── ACK Reactions ─────────────────────────────────────────────── async def _send_ack_reaction( diff --git a/EvoScientist/channels/telegram/channel.py b/EvoScientist/channels/telegram/channel.py index c2ef180..793c0a6 100644 --- a/EvoScientist/channels/telegram/channel.py +++ b/EvoScientist/channels/telegram/channel.py @@ -2,7 +2,7 @@ import logging from dataclasses import dataclass -from datetime import datetime +from datetime import datetime, timedelta from pathlib import Path from typing import ClassVar @@ -34,7 +34,11 @@ class TelegramChannel(Channel): capabilities = TELEGRAM_CAPS _typing_interval: float = 4.0 _ready_attrs = ("_app",) - _non_retryable_patterns = ("parse", "can't parse") + _non_retryable_patterns = ( + *Channel._non_retryable_patterns, + "parse", + "can't parse", + ) _mention_pattern = r"(?i)@{bot_id}\s*" def __init__(self, config: TelegramConfig): @@ -114,6 +118,21 @@ class TelegramChannel(Channel): action="typing", ) + # ── Retry delay extraction (override base) ───────────────────── + + def _extract_retry_delay(self, exc: Exception) -> float | None: + """Honor Telegram flood control (``telegram.error.RetryAfter``). + + ``retry_after`` is an ``int`` by default and a ``timedelta`` when the + ``PTB_TIMEDELTA`` opt-in is enabled. + """ + from telegram.error import RetryAfter + + if isinstance(exc, RetryAfter): + ra = exc.retry_after + return ra.total_seconds() if isinstance(ra, timedelta) else float(ra) + return super()._extract_retry_delay(exc) + # ── Send (template method overrides) ────────────────────────── async def _send_chunk(self, chat_id, formatted_text, raw_text, reply_to, metadata): diff --git a/tests/test_channel_comprehensive.py b/tests/test_channel_comprehensive.py index f4e7bfc..6d3016b 100644 --- a/tests/test_channel_comprehensive.py +++ b/tests/test_channel_comprehensive.py @@ -16,9 +16,10 @@ from __future__ import annotations import asyncio import threading -from datetime import datetime +from datetime import UTC, datetime from unittest.mock import AsyncMock, MagicMock +import httpx import pytest from EvoScientist.channels.base import ( @@ -1221,29 +1222,245 @@ class TestChannelReconnect: class TestExtractRetryAfter: - def test_never_returns_none(self): - """[B-01] Base _extract_retry_after always returns float, never None.""" + def test_generic_errors_use_default_retry_delay(self): + """Generic transient errors should use the base retry delay.""" ch = StubChannel() - # Even for a generic exception, it returns 1.0 instead of None result = ch._extract_retry_after(ValueError("bad")) - # BUG: This should return None for non-retryable errors - # Current behavior: always returns 1.0 - assert result is not None # Documents the bug + assert result == 1.0 - def test_extracts_retry_after_attribute(self): + def test_explicit_auth_errors_do_not_retry(self): ch = StubChannel() + result = ch._extract_retry_after(Exception("HTTP 401 Unauthorized")) + assert result is None - class RateLimitError(Exception): - retry_after = 5.0 - - result = ch._extract_retry_after(RateLimitError("rate limited")) - assert result == 5.0 + def test_5xx_style_errors_still_retry(self): + ch = StubChannel() + result = ch._extract_retry_after(Exception("HTTP 500 Internal Server Error")) + assert result == 1.0 def test_detects_429_in_message(self): ch = StubChannel() result = ch._extract_retry_after(RuntimeError("HTTP 429 Too Many Requests")) assert result == 1.0 + # ── Real HTTP SDK status code tests (httpx / aiohttp) ──────────── + + def test_httpx_401_not_retryable(self): + """httpx.HTTPStatusError with status 401 should return None (no retry).""" + exc = httpx.HTTPStatusError( + "unauthorized", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(401), + ) + assert StubChannel()._extract_retry_after(exc) is None + + def test_httpx_403_not_retryable(self): + """httpx.HTTPStatusError with status 403 should return None (no retry).""" + exc = httpx.HTTPStatusError( + "forbidden", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(403), + ) + assert StubChannel()._extract_retry_after(exc) is None + + @pytest.mark.parametrize("status", [400, 404]) + def test_httpx_permanent_4xx_not_retryable(self, status): + """400 and 404 are permanent for a given request and must not retry.""" + exc = httpx.HTTPStatusError( + "client error", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(status), + ) + assert StubChannel()._extract_retry_after(exc) is None + + def test_httpx_408_still_retries(self): + """Not every 4xx is permanent: 408 Request Timeout keeps the default delay.""" + exc = httpx.HTTPStatusError( + "request timeout", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(408), + ) + assert StubChannel()._extract_retry_after(exc) == 1.0 + + def test_httpx_500_is_retryable(self): + """httpx.HTTPStatusError with status 500 should retry (default 1.0s).""" + exc = httpx.HTTPStatusError( + "server error", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(500), + ) + assert StubChannel()._extract_retry_after(exc) == 1.0 + + def test_aiohttp_401_not_retryable(self): + """aiohttp.ClientResponseError with status 401 should return None (no retry).""" + import aiohttp + from yarl import URL + + exc = aiohttp.ClientResponseError( + request_info=aiohttp.RequestInfo( + url=URL("https://example.invalid"), + method="POST", + headers={}, + real_url=URL("https://example.invalid"), + ), + history=(), + status=401, + message="Unauthorized", + ) + assert StubChannel()._extract_retry_after(exc) is None + + def test_aiohttp_403_not_retryable(self): + """aiohttp.ClientResponseError with status 403 should return None (no retry).""" + import aiohttp + from yarl import URL + + exc = aiohttp.ClientResponseError( + request_info=aiohttp.RequestInfo( + url=URL("https://example.invalid"), + method="POST", + headers={}, + real_url=URL("https://example.invalid"), + ), + history=(), + status=403, + message="Forbidden", + ) + assert StubChannel()._extract_retry_after(exc) is None + + def test_aiohttp_500_is_retryable(self): + """aiohttp.ClientResponseError with status 500 should retry (default 1.0s).""" + import aiohttp + from yarl import URL + + exc = aiohttp.ClientResponseError( + request_info=aiohttp.RequestInfo( + url=URL("https://example.invalid"), + method="POST", + headers={}, + real_url=URL("https://example.invalid"), + ), + history=(), + status=500, + message="Server Error", + ) + assert StubChannel()._extract_retry_after(exc) == 1.0 + + def test_httpx_401_with_retry_after_header_is_still_not_retryable(self): + """Non-retryable 401 takes precedence over Retry-After header.""" + ch = StubChannel() + exc = httpx.HTTPStatusError( + "unauthorized", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(401, headers={"Retry-After": "10"}), + ) + assert ch._extract_retry_after(exc) is None + + # ── _extract_status_code tests ─────────────────────────────────── + + def test_extract_status_code_from_httpx(self): + """_extract_status_code extracts status_code from httpx.HTTPStatusError.""" + ch = StubChannel() + exc = httpx.HTTPStatusError( + "unauthorized", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(401), + ) + assert ch._extract_status_code(exc) == 401 + + def test_extract_status_code_non_http_returns_none(self): + """_extract_status_code returns None for non-HTTP exceptions.""" + ch = StubChannel() + assert ch._extract_status_code(RuntimeError("plain error")) is None + + # ── _extract_sdk_error_code tests ───────────────────────────────── + + def test_base_extract_sdk_error_code_returns_none(self): + """Base Channel._extract_sdk_error_code returns None by default.""" + ch = StubChannel() + assert ch._extract_sdk_error_code(RuntimeError("plain error")) is None + exc = httpx.HTTPStatusError( + "error", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(401), + ) + assert ch._extract_sdk_error_code(exc) is None + + # ── _extract_retry_delay tests ─────────────────────────────── + + def test_extract_retry_delay_integer(self): + """_extract_retry_delay parses integer string from headers.""" + ch = StubChannel() + resp = httpx.Response(429, headers={"Retry-After": "10"}) + exc = httpx.HTTPStatusError( + "rate limited", + request=httpx.Request("POST", "https://example.invalid"), + response=resp, + ) + assert ch._extract_retry_delay(exc) == 10.0 + + def test_extract_retry_delay_float(self): + """_extract_retry_delay parses float string from lowercase headers.""" + ch = StubChannel() + resp = httpx.Response(429, headers={"retry-after": "2.5"}) + exc = httpx.HTTPStatusError( + "rate limited", + request=httpx.Request("POST", "https://example.invalid"), + response=resp, + ) + assert ch._extract_retry_delay(exc) == 2.5 + + def test_extract_retry_delay_invalid_value(self): + """_extract_retry_delay returns None for non-numeric header.""" + ch = StubChannel() + resp = httpx.Response(429, headers={"Retry-After": "invalid-date"}) + exc = httpx.HTTPStatusError( + "rate limited", + request=httpx.Request("POST", "https://example.invalid"), + response=resp, + ) + assert ch._extract_retry_delay(exc) is None + + def test_extract_retry_delay_missing(self): + """_extract_retry_delay returns None when no Retry-After header exists.""" + ch = StubChannel() + resp = httpx.Response(429, headers={"Content-Type": "application/json"}) + exc = httpx.HTTPStatusError( + "rate limited", + request=httpx.Request("POST", "https://example.invalid"), + response=resp, + ) + assert ch._extract_retry_delay(exc) is None + + def test_extract_retry_delay_http_date(self): + """An HTTP-date Retry-After is honored as seconds until that time.""" + from datetime import datetime, timedelta + from email.utils import format_datetime + + when = datetime.now(UTC) + timedelta(seconds=60) + resp = httpx.Response( + 503, headers={"Retry-After": format_datetime(when, usegmt=True)} + ) + exc = httpx.HTTPStatusError( + "unavailable", + request=httpx.Request("POST", "https://example.invalid"), + response=resp, + ) + delay = StubChannel()._extract_retry_after(exc) + assert delay is not None + assert 55.0 <= delay <= 60.0 + + def test_extract_retry_delay_http_date_in_past_is_zero(self): + """A past HTTP-date yields 0.0 rather than a negative delay.""" + resp = httpx.Response( + 503, headers={"Retry-After": "Wed, 21 Oct 2015 07:28:00 GMT"} + ) + exc = httpx.HTTPStatusError( + "unavailable", + request=httpx.Request("POST", "https://example.invalid"), + response=resp, + ) + assert StubChannel()._extract_retry_delay(exc) == 0.0 + class TestChannelAttachments: def test_check_attachment_size_within_limit(self): diff --git a/tests/test_discord_channel.py b/tests/test_discord_channel.py index 18eec2e..b98b8a2 100644 --- a/tests/test_discord_channel.py +++ b/tests/test_discord_channel.py @@ -1,5 +1,7 @@ """Tests for Discord channel implementation.""" +import importlib.util + import pytest from EvoScientist.channels.base import ChannelError @@ -37,3 +39,73 @@ class TestDiscordChannel: ) result = await channel.send(msg) assert result is False + + +@pytest.mark.skipif( + importlib.util.find_spec("discord") is None, + reason="discord.py not installed", +) +class TestDiscordRetryErrorExtraction: + """Test Discord-specific status code extraction.""" + + def test_discord_status_code_401_not_retryable(self): + from unittest.mock import MagicMock + + import discord + + ch = DiscordChannel(DiscordConfig(bot_token="test")) + resp = MagicMock() + resp.status = 401 + resp.reason = "Unauthorized" + resp.headers = {} + exc = discord.HTTPException(resp, "401 Unauthorized") + assert ch._extract_status_code(exc) == 401 + assert ch._extract_retry_after(exc) is None + + def test_discord_status_code_403_not_retryable(self): + from unittest.mock import MagicMock + + import discord + + ch = DiscordChannel(DiscordConfig(bot_token="test")) + resp = MagicMock() + resp.status = 403 + resp.reason = "Forbidden" + resp.headers = {} + exc = discord.HTTPException(resp, "50001 Missing Access") + assert ch._extract_status_code(exc) == 403 + assert ch._extract_retry_after(exc) is None + + def test_discord_status_code_500_is_retryable(self): + from unittest.mock import MagicMock + + import discord + + ch = DiscordChannel(DiscordConfig(bot_token="test")) + resp = MagicMock() + resp.status = 500 + resp.reason = "Internal Server Error" + resp.headers = {} + exc = discord.HTTPException(resp, "500 Internal Server Error") + assert ch._extract_status_code(exc) == 500 + assert ch._extract_retry_after(exc) == 1.0 + + def test_discord_fallback_to_httpx(self): + import httpx + + ch = DiscordChannel(DiscordConfig(bot_token="test")) + exc = httpx.HTTPStatusError( + "unauthorized", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(401), + ) + assert ch._extract_status_code(exc) == 401 + assert ch._extract_retry_after(exc) is None + + def test_discord_rate_limited_uses_retry_after(self): + import discord + + ch = DiscordChannel(DiscordConfig(bot_token="test")) + exc = discord.RateLimited(12.5) + assert ch._extract_retry_delay(exc) == 12.5 + assert ch._extract_retry_after(exc) == 12.5 diff --git a/tests/test_slack_channel.py b/tests/test_slack_channel.py index 2c291b1..9a7de06 100644 --- a/tests/test_slack_channel.py +++ b/tests/test_slack_channel.py @@ -1,5 +1,7 @@ """Tests for Slack channel implementation.""" +import importlib.util + import pytest from EvoScientist.channels.base import ChannelError @@ -75,3 +77,188 @@ class TestSlackChannelRegistration: channels = available_channels() assert "slack" in channels + + +@pytest.mark.skipif( + importlib.util.find_spec("slack_sdk") is None, + reason="slack-sdk not installed", +) +class TestSlackRetryErrorExtraction: + """Test Slack-specific status code and SDK error code extraction.""" + + def test_extract_slack_auth_error_not_retryable(self): + from slack_sdk.errors import SlackApiError + from slack_sdk.web.slack_response import SlackResponse + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + resp = SlackResponse( + client=None, + http_verb="POST", + api_url="https://slack.com/api/chat.postMessage", + req_args={}, + data={"ok": False, "error": "invalid_auth"}, + headers={}, + status_code=200, + ) + exc = SlackApiError("The request to the Slack API failed.", response=resp) + assert ch._extract_sdk_error_code(exc) == "invalid_auth" + assert ch._extract_retry_after(exc) is None + + def test_extract_slack_token_expired_not_retryable(self): + from slack_sdk.errors import SlackApiError + from slack_sdk.web.slack_response import SlackResponse + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + resp = SlackResponse( + client=None, + http_verb="POST", + api_url="https://slack.com/api/chat.postMessage", + req_args={}, + data={"ok": False, "error": "token_expired"}, + headers={}, + status_code=200, + ) + exc = SlackApiError("The token has expired.", response=resp) + assert ch._extract_sdk_error_code(exc) == "token_expired" + assert ch._extract_retry_after(exc) is None + + def test_extract_slack_status_code_401_not_retryable(self): + from slack_sdk.errors import SlackApiError + from slack_sdk.web.slack_response import SlackResponse + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + resp = SlackResponse( + client=None, + http_verb="POST", + api_url="https://slack.com/api/chat.postMessage", + req_args={}, + data={"ok": False, "error": "unknown_custom"}, + headers={}, + status_code=401, + ) + exc = SlackApiError("Unauthorized", response=resp) + assert ch._extract_status_code(exc) == 401 + assert ch._extract_retry_after(exc) is None + + def test_extract_slack_status_code_500_is_retryable(self): + from slack_sdk.errors import SlackApiError + from slack_sdk.web.slack_response import SlackResponse + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + resp = SlackResponse( + client=None, + http_verb="POST", + api_url="https://slack.com/api/chat.postMessage", + req_args={}, + data={"ok": False, "error": "internal_error"}, + headers={}, + status_code=500, + ) + exc = SlackApiError("Internal Server Error", response=resp) + assert ch._extract_status_code(exc) == 500 + assert ch._extract_retry_after(exc) == 1.0 + + def test_slack_channel_fallback_to_httpx(self): + import httpx + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + exc = httpx.HTTPStatusError( + "unauthorized", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(401), + ) + assert ch._extract_status_code(exc) == 401 + assert ch._extract_retry_after(exc) is None + + def test_slack_ratelimited_uses_retry_after_header(self): + from slack_sdk.errors import SlackApiError + from slack_sdk.web.slack_response import SlackResponse + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + resp = SlackResponse( + client=None, + http_verb="POST", + api_url="https://slack.com/api/chat.postMessage", + req_args={}, + data={"ok": False, "error": "ratelimited"}, + headers={"Retry-After": "30"}, + status_code=429, + ) + exc = SlackApiError("ratelimited", response=resp) + assert ch._extract_retry_delay(exc) == 30.0 + assert ch._extract_retry_after(exc) == 30.0 + + def test_slack_malformed_retry_after_falls_through(self): + from slack_sdk.errors import SlackApiError + from slack_sdk.web.slack_response import SlackResponse + + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + resp = SlackResponse( + client=None, + http_verb="POST", + api_url="https://slack.com/api/chat.postMessage", + req_args={}, + data={"ok": False, "error": "ratelimited"}, + headers={"Retry-After": "soon"}, + status_code=429, + ) + exc = SlackApiError("ratelimited", response=resp) + assert ch._extract_retry_delay(exc) is None + assert ch._extract_retry_after(exc) == ch._rate_limit_delay + + +@pytest.mark.skipif( + importlib.util.find_spec("slack_sdk") is None + or importlib.util.find_spec("aiohttp") is None, + reason="slack_sdk or aiohttp not installed", +) +class TestSlackRetryWithRawClientResponse: + """slack_sdk wraps the raw aiohttp response in SlackApiError when a + JSON-declared body fails to parse; the retry path must survive that.""" + + async def test_malformed_json_body_is_retried_and_surfaces_sdk_error( + self, monkeypatch + ): + import aiohttp + from aiohttp import web + from slack_sdk.errors import SlackApiError + from slack_sdk.web.async_client import AsyncWebClient + + from EvoScientist.channels.retry import RetryConfig + + for var in ("HTTP_PROXY", "HTTPS_PROXY", "http_proxy", "https_proxy"): + monkeypatch.delenv(var, raising=False) + + calls = 0 + + async def handler(request): + nonlocal calls + calls += 1 + return web.Response( + status=200, text="<>", content_type="application/json" + ) + + app = web.Application() + app.router.add_post("/api/chat.postMessage", handler) + runner = web.AppRunner(app) + await runner.setup() + try: + await web.TCPSite(runner, "127.0.0.1", 0).start() + port = runner.addresses[0][1] + client = AsyncWebClient( + token="xoxb-test", + base_url=f"http://127.0.0.1:{port}/api/", + retry_handlers=[], + ) + ch = SlackChannel(SlackConfig(bot_token="xoxb-test", app_token="xapp-test")) + ch._retry_config = RetryConfig( + attempts=3, min_delay_s=0.01, max_delay_s=0.02, jitter=0.0 + ) + with pytest.raises(SlackApiError) as excinfo: + await ch._send_with_retry( + lambda: client.chat_postMessage(channel="C1", text="hi") + ) + assert isinstance(excinfo.value.response, aiohttp.ClientResponse) + assert calls == 3 + finally: + await runner.cleanup() diff --git a/tests/test_telegram_channel.py b/tests/test_telegram_channel.py index e0e2b2d..a5a54ec 100644 --- a/tests/test_telegram_channel.py +++ b/tests/test_telegram_channel.py @@ -1,5 +1,6 @@ """Tests for Telegram channel implementation.""" +import importlib.util import sys from datetime import datetime from types import ModuleType, SimpleNamespace @@ -275,3 +276,27 @@ class TestTelegramChannel: message_id=789, ) return SimpleNamespace(message=message) + + +@pytest.mark.skipif( + importlib.util.find_spec("telegram") is None, + reason="python-telegram-bot not installed", +) +class TestTelegramRetryDelay: + def test_retry_after_honored(self): + from telegram.error import RetryAfter + + ch = TelegramChannel(TelegramConfig(bot_token="t")) + assert ch._extract_retry_delay(RetryAfter(7)) == 7.0 + assert ch._extract_retry_after(RetryAfter(7)) == 7.0 + + def test_other_errors_fall_through(self): + import httpx + + ch = TelegramChannel(TelegramConfig(bot_token="t")) + exc = httpx.HTTPStatusError( + "429", + request=httpx.Request("POST", "https://example.invalid"), + response=httpx.Response(429, headers={"Retry-After": "3"}), + ) + assert ch._extract_retry_delay(exc) == 3.0