01845f4311
* chore: add pytest-asyncio in auto mode * test: migrate channel and stream tests to native async Convert run_async() wrapper tests to plain 'async def test_*' under pytest-asyncio auto mode. collect_events() in stream_v3_fakes becomes a coroutine awaited at every call site. * test: migrate command and model/middleware tests to native async Convert run_async() wrappers (import, alias, and fixture forms) to plain 'async def test_*'. Multi-call tests merge onto one loop as sequential awaits; none asserted on loop identity. * test: migrate TUI, notifier, gateway, and session tests to native async TUI/notifier/gateway files convert run_async wrappers to plain async tests. test_sessions.py's unittest.TestCase classes move to unittest.IsolatedAsyncioTestCase (pytest-asyncio does not await async methods on plain TestCase; converting blindly would have made ~70 tests silently vacuous). Its setUpClass keeps a one-shot asyncio.run() since IsolatedAsyncioTestCase has no async class-level hook. TestLoadingWidget in test_tui_widgets.py drops its TestCase base for the same reason. * test: replace direct asyncio.run() calls with native async tests Convert tests that called asyncio.run() (directly or via a local _run helper) to plain 'async def test_*'; delete the local helpers. * test: drop undeclared anyio markers and delete run_async helper The @pytest.mark.anyio tests relied on anyio being a transitive dep of httpx; auto-mode pytest-asyncio collects them natively. run_async() and its fixture are unreferenced after the migration, so remove them — pytest-asyncio's per-test loop teardown covers the pending-task cancellation the helper existed for (verified: full suite runs with no 'Event loop is closed' errors or destroyed-task warnings). * test: add autouse fixture for watcher cleanup * refactor: remove redundant hasattr calls * refactor: add typed middleware event sink and thread through assembly Add MiddlewareEventSink protocol + NoOpSink in middleware/events.py with a documented any-thread non-blocking contract (contract test uses a deliberately-slow fake sink). Thread an optional `events` parameter through create_cli_agent -> _get_default_middleware -> tool selector / model fallback constructors; subagent stacks are always forced to NoOpSink. * refactor: inject a notifier port into async-watcher and background middleware Add public pre_cancel_watcher() and enqueue_task_notification() to cli/async_notifier.py and a small NotifierPort protocol (middleware/notifier.py) that the module satisfies structurally. AsyncWatcherMiddleware and BackgroundExecutionMiddleware now receive the port by constructor injection at the composition root, deleting the lazy 'from ..cli import async_notifier' imports and the private _watcher_by_thread / _enqueue pokes. * refactor: invert tool-selection ownership onto a frontend event sink The adaptive tool selector now reports on_tool_selection_started / on_tool_selection / on_tool_selection_ended to the injected sink instead of writing four process-global module variables. The frontend sink (stream/sink.py FrontendEventSink) owns the selected/total/active state with consume-once + dedup-vs-last-emitted semantics; stream/tool_selection.py reads that sink object (a ToolSelectionView) rather than reaching into tool_selector's globals. Deleted: the 4 module globals, the cross-module mutations in tool_selection.py, the track_stream_selection flag, the now-vestigial _ToolSelectionTrackerMiddleware, reset_tool_selection_state_for_tests, and the autouse conftest fixture. The sink is threaded from the two interactive frontends through create_runtime_gateways -> LocalGraphGateway (read side) and _load_agent -> create_cli_agent (write side); subagent / headless stacks get NoOpSink. * refactor: route model-fallback narration through the injected event sink Delete the _ui_emit_fn / set_ui_emit module global and the ..stream.console import from model_fallback.py. The fallback middleware now reports through its injected sink: the fallback transition via the structured on_model_fallback (the frontend formats the '-> Falling back to ...' line), and the surrounding narration (primary-failure header, per-attempt outcome, exhaustion, non-fallbackable rejection) via emit_fallback_notice, preserving the exact user-facing text. The TUI binds its _append_system as the sink's fallback display where it used to call set_ui_emit (cleared on exit); the Rich CLI's sink prints to the console. _try_fallbacks / _guard_and_fallback take the sink. * refactor: declare events on the GraphGateway protocol Both gateway implementations now carry an explicit events attribute (LangGraphServerGateway holds None — no frontend renders middleware events across the HTTP boundary), so the four call sites use plain attribute access instead of getattr probing an implicit contract. * refactor: bind fallback display via the closure-scoped concrete sink The App methods used gateway.events (typed as the read-side view) and hasattr-probed for the concrete FrontendEventSink API. The enclosing factory creates that sink two hundred lines up — close over it directly: no probing, fully typed, and it becomes a constructor parameter naturally when the App class is hoisted out of the factory. * fix: end tool selection before fallback handler * fix: keep fallback display errors non-fatal * fix: preserve selector suppression for default streams * fix: restore fallback notice console display * refactor: consolidate fallback narration events * refactor: clean middleware event sink plumbing * fix: type gateway session events * refactor: make all event protocols runtime-checkable MiddlewareEventSink already carried @runtime_checkable (the stream binding guard isinstance-checks it); ToolSelectionView and SessionEvents now match, so mirroring that pattern against any of the three protocols works instead of raising TypeError. * fix(cli): close QuickJS workers after one-shot failures * fix(cli): honor no-thinking in final output * fix(channels): report failed startup accurately * fix(channels): make Telegram cleanup idempotent * fix(tui): skip command sync during exit * fix(channels): preserve startup state during retries * refactor(channels): share pending startup status * refactor(cli): expose channel startup snapshot * fix(tui): move channel startup off event loop * test(channels): release retry gate on assertion failure --------- Co-authored-by: Xi Zhang <106144707+X-iZhang@users.noreply.github.com>
1758 lines
60 KiB
Python
1758 lines
60 KiB
Python
"""Comprehensive channel test suite — covers all major functionalities and known bug scenarios.
|
|
|
|
Bug IDs prefixed with [B-xx] map to the internal bug report.
|
|
Test groups:
|
|
1. DedupCache — dedup correctness, TTL, LRU, boundary
|
|
2. RetryConfig / retry — exponential backoff, jitter, should_retry
|
|
3. chunk_text — text splitting, code fences, edge cases
|
|
4. markdown_utils — placeholder integrity, escape_fn, inline/block
|
|
5. Channel base — send, debounce, typing, allow-list, reconnect
|
|
6. ChannelManager — register, dispatch, health, add/remove, drain
|
|
7. InboundConsumer — worker pool, session, timeout, error handling
|
|
8. MessageBus — pub/sub, backpressure, subscriber dispatch
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import asyncio
|
|
from datetime import datetime
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
|
|
from EvoScientist.channels.base import (
|
|
ChannelError,
|
|
InboundMessage,
|
|
OutboundMessage,
|
|
RawIncoming,
|
|
chunk_text,
|
|
)
|
|
from EvoScientist.channels.bus.events import (
|
|
InboundMessage as BusInbound,
|
|
)
|
|
from EvoScientist.channels.bus.events import (
|
|
OutboundMessage as BusOutbound,
|
|
)
|
|
from EvoScientist.channels.bus.message_bus import MessageBus
|
|
from EvoScientist.channels.channel_manager import ChannelManager
|
|
from EvoScientist.channels.consumer import InboundConsumer
|
|
from EvoScientist.channels.formatter import convert_markdown
|
|
from EvoScientist.channels.middleware import DedupCache
|
|
from EvoScientist.channels.retry import RetryConfig, RetryInfo, retry_async
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# Helpers
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
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
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestDedupCache:
|
|
def test_first_message_is_not_duplicate(self):
|
|
dc = DedupCache()
|
|
assert dc.is_duplicate("msg_001") is False
|
|
|
|
def test_same_id_is_duplicate(self):
|
|
dc = DedupCache()
|
|
dc.is_duplicate("msg_001")
|
|
assert dc.is_duplicate("msg_001") is True
|
|
|
|
def test_empty_id_never_duplicate(self):
|
|
dc = DedupCache()
|
|
assert dc.is_duplicate("") is False
|
|
assert dc.is_duplicate("") is False
|
|
|
|
def test_ttl_expiry(self):
|
|
clock = ManualClock()
|
|
dc = DedupCache(ttl_seconds=0.05, clock=clock)
|
|
dc.is_duplicate("msg_001")
|
|
clock.advance(0.051)
|
|
# After TTL, the entry should be pruned
|
|
assert dc.is_duplicate("msg_001") is False
|
|
|
|
def test_max_size_trim(self):
|
|
dc = DedupCache(max_size=5, trim_to=2)
|
|
for i in range(6):
|
|
dc.is_duplicate(f"m{i}")
|
|
# After exceeding max_size, trimmed to trim_to
|
|
assert dc.size <= 3 # 2 kept + the just-inserted one
|
|
|
|
def test_lru_refresh(self):
|
|
"""Accessing an entry refreshes its position (LRU)."""
|
|
dc = DedupCache(max_size=3, trim_to=1, ttl_seconds=60)
|
|
dc.is_duplicate("a")
|
|
dc.is_duplicate("b")
|
|
# Re-access "a" to move it to end
|
|
dc.is_duplicate("a")
|
|
dc.is_duplicate("c")
|
|
# Now exceed — oldest insertion-order should be "b"
|
|
dc.is_duplicate("d")
|
|
# "a" was refreshed, so "b" should have been evicted
|
|
assert dc.is_duplicate("b") is False # "b" was evicted
|
|
|
|
def test_clear(self):
|
|
dc = DedupCache()
|
|
dc.is_duplicate("x")
|
|
dc.clear()
|
|
assert dc.size == 0
|
|
assert dc.is_duplicate("x") is False
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 2. Retry
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestRetryAsync:
|
|
async def test_success_on_first_attempt(self):
|
|
call_count = 0
|
|
|
|
async def _fn():
|
|
nonlocal call_count
|
|
call_count += 1
|
|
return "ok"
|
|
|
|
result = await retry_async(_fn)
|
|
assert result == "ok"
|
|
assert call_count == 1
|
|
|
|
async def test_retries_on_failure_then_succeeds(self):
|
|
attempts = []
|
|
|
|
async def _fn():
|
|
attempts.append(1)
|
|
if len(attempts) < 3:
|
|
raise RuntimeError("transient")
|
|
return "recovered"
|
|
|
|
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
|
|
|
|
async def test_exhausts_retries_raises(self):
|
|
async def _fn():
|
|
raise ValueError("permanent")
|
|
|
|
with pytest.raises(ValueError, match="permanent"):
|
|
await retry_async(
|
|
_fn,
|
|
config=RetryConfig(attempts=2, min_delay_s=0.01),
|
|
)
|
|
|
|
async def test_should_retry_false_aborts(self):
|
|
"""[B-01] should_retry returning False should abort immediately."""
|
|
call_count = 0
|
|
|
|
async def _fn():
|
|
nonlocal call_count
|
|
call_count += 1
|
|
raise PermissionError("forbidden")
|
|
|
|
with pytest.raises(PermissionError):
|
|
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
|
|
|
|
async def test_server_retry_after_respected(self):
|
|
"""retry_after_s callback provides server-supplied delay."""
|
|
delays = []
|
|
|
|
async def _fn():
|
|
if len(delays) < 1:
|
|
raise RuntimeError("429")
|
|
return "ok"
|
|
|
|
def _on_retry(info: RetryInfo):
|
|
delays.append(info.delay_s)
|
|
|
|
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
|
|
|
|
async def test_jitter_applied(self):
|
|
"""With jitter > 0, delays should vary."""
|
|
delays = []
|
|
|
|
async def _fn():
|
|
if len(delays) < 5:
|
|
raise RuntimeError("fail")
|
|
return "ok"
|
|
|
|
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:
|
|
assert len({f"{d:.4f}" for d in delays}) > 1
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 3. chunk_text
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestChunkText:
|
|
def test_short_text_single_chunk(self):
|
|
assert chunk_text("hello", 100) == ["hello"]
|
|
|
|
def test_empty_text(self):
|
|
assert chunk_text("", 100) == []
|
|
|
|
def test_exact_limit(self):
|
|
text = "a" * 100
|
|
assert chunk_text(text, 100) == [text]
|
|
|
|
def test_splits_at_paragraph_break(self):
|
|
text = "first paragraph\n\nsecond paragraph"
|
|
chunks = chunk_text(text, 25)
|
|
assert len(chunks) == 2
|
|
assert "first" in chunks[0]
|
|
assert "second" in chunks[1]
|
|
|
|
def test_splits_at_newline(self):
|
|
text = "line one\nline two\nline three"
|
|
chunks = chunk_text(text, 15)
|
|
assert all(len(c) <= 15 for c in chunks)
|
|
assert len(chunks) >= 2
|
|
|
|
def test_splits_at_space(self):
|
|
text = "word " * 30
|
|
chunks = chunk_text(text, 20)
|
|
assert all(len(c) <= 20 for c in chunks)
|
|
|
|
def test_hard_cut_no_separators(self):
|
|
text = "a" * 200
|
|
chunks = chunk_text(text, 50)
|
|
assert all(len(c) <= 50 for c in chunks)
|
|
|
|
def test_code_block_fence_split(self):
|
|
"""[B-08] Code block fence splitting should not break mid-block without refencing."""
|
|
code = "```python\nprint('hello')\nprint('world')\n```"
|
|
text = "Before.\n\n" + code + "\n\nAfter some text here."
|
|
# Use a limit that forces a split inside the code block
|
|
chunks = chunk_text(text, 25)
|
|
|
|
# Verify we get multiple chunks and none are empty
|
|
assert len(chunks) >= 2
|
|
assert all(c.strip() for c in chunks)
|
|
|
|
# Check that the code block was properly re-fenced
|
|
# The first chunk should open the block but not close it (if it splits mid-block)
|
|
# Actually, the new implementation adds closing fences to the split part and opens the next.
|
|
# Let's just check that all chunks are valid markdown and the code is preserved.
|
|
reconstructed = (
|
|
"".join(c for c in chunks)
|
|
.replace("```python\n", "")
|
|
.replace("\n```", "")
|
|
.replace("```\n", "")
|
|
)
|
|
assert "print('hello')" in reconstructed
|
|
assert "print('world')" in reconstructed
|
|
|
|
# At least one chunk should have a re-fenced code block if it split
|
|
has_refence = any("```python\n" in c and c.count("```") == 2 for c in chunks)
|
|
assert has_refence, "No chunks were properly re-fenced"
|
|
|
|
# It's possible it split exactly on the fence, so we can't assert has_refence strictly without knowing the exact cut,
|
|
# but we can assert that every chunk has balanced or correctly formatted fences.
|
|
for c in chunks:
|
|
if "```" in c:
|
|
assert c.count("```") % 2 == 0, f"Unbalanced fences in chunk: {c}"
|
|
|
|
def test_long_code_block_refencing(self):
|
|
"""Verify that a very long code block is split and each chunk gets fences."""
|
|
code = "```js\n" + "line of code\n" * 10 + "```"
|
|
chunks = chunk_text(code, 50)
|
|
assert len(chunks) > 1
|
|
for chunk in chunks:
|
|
assert chunk.startswith("```js\n") or chunk.startswith("```")
|
|
assert chunk.rstrip().endswith("```")
|
|
assert chunk.count("```") >= 2
|
|
|
|
def test_code_block_preserved_when_fits(self):
|
|
code = "```\ncode\n```"
|
|
text = f"intro\n\n{code}\n\noutro"
|
|
chunks = chunk_text(text, 200)
|
|
assert len(chunks) == 1
|
|
assert "```" in chunks[0]
|
|
|
|
def test_whitespace_only_input(self):
|
|
"""[B-09] Whitespace-heavy input should not produce empty chunks."""
|
|
text = " \n\n \n\n content \n\n "
|
|
chunks = chunk_text(text, 20)
|
|
assert all(c.strip() for c in chunks)
|
|
|
|
def test_very_small_limit(self):
|
|
"""Limit below typical message sizes."""
|
|
text = "Hello, this is a test message."
|
|
chunks = chunk_text(text, 5)
|
|
assert all(len(c) <= 5 for c in chunks)
|
|
assert "".join(c.replace(" ", "") for c in chunks).replace(" ", "") != ""
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 4. markdown_utils — convert_markdown
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestMarkdownUtils:
|
|
@staticmethod
|
|
def _html_converter(text: str) -> str:
|
|
return convert_markdown(
|
|
text,
|
|
code_block_formatter=lambda lang, code: f"<pre>{code}</pre>",
|
|
inline_code_formatter=lambda code: f"<code>{code}</code>",
|
|
inline_rules=[
|
|
(r"\*\*(.+?)\*\*", r"<b>\1</b>"),
|
|
(r"\*(.+?)\*", r"<i>\1</i>"),
|
|
],
|
|
escape_fn=lambda t: (
|
|
t.replace("&", "&").replace("<", "<").replace(">", ">")
|
|
),
|
|
)
|
|
|
|
def test_basic_bold_italic(self):
|
|
result = self._html_converter("**bold** and *italic*")
|
|
assert "<b>bold</b>" in result
|
|
assert "<i>italic</i>" in result
|
|
|
|
def test_code_block_protection(self):
|
|
"""Code inside blocks should NOT have inline rules applied."""
|
|
text = "```\n**not bold**\n```"
|
|
result = self._html_converter(text)
|
|
assert "<b>" not in result
|
|
assert "**not bold**" in result
|
|
|
|
def test_inline_code_protection(self):
|
|
text = "Use `**literal**` please"
|
|
result = self._html_converter(text)
|
|
assert "<code>" in result
|
|
# The **literal** inside backticks should be literal
|
|
assert "**literal**" in result
|
|
|
|
def test_escape_fn_does_not_corrupt_placeholders(self):
|
|
"""[B-28] escape_fn must not corrupt NUL-byte placeholders."""
|
|
text = "```\ncode\n```\nNormal <text>"
|
|
|
|
def bad_escape(t):
|
|
# Strips NUL bytes — would break placeholders
|
|
return t.replace("\x00", "")
|
|
|
|
result = convert_markdown(
|
|
text,
|
|
code_block_formatter=lambda lang, c: f"[CODE]{c}[/CODE]",
|
|
inline_code_formatter=lambda c: f"[IC]{c}[/IC]",
|
|
inline_rules=[],
|
|
escape_fn=bad_escape,
|
|
)
|
|
# If placeholders were corrupted, the code block won't be restored
|
|
# This test DOCUMENTS the bug — it should fail until the bug is fixed
|
|
# After fix: assert "[CODE]" in result
|
|
# Current behavior: placeholder is corrupted
|
|
if "\x00" in text:
|
|
pass # Can't easily test without modifying source
|
|
# At minimum, verify the function doesn't crash
|
|
assert isinstance(result, str)
|
|
|
|
def test_placeholder_collision_with_user_input(self):
|
|
"""[B-28 variant] User input containing placeholder pattern."""
|
|
text = "Normal text with \x00BLOCK0\x00 in it"
|
|
result = convert_markdown(
|
|
text,
|
|
code_block_formatter=lambda lang, c: f"<pre>{c}</pre>",
|
|
inline_code_formatter=lambda c: f"<code>{c}</code>",
|
|
inline_rules=[],
|
|
)
|
|
assert isinstance(result, str)
|
|
|
|
def test_empty_inline_code(self):
|
|
"""[B-29] Empty backtick pairs should not crash."""
|
|
text = "before `` after"
|
|
result = convert_markdown(
|
|
text,
|
|
code_block_formatter=lambda lang, c: c,
|
|
inline_code_formatter=lambda c: f"[{c}]",
|
|
inline_rules=[],
|
|
)
|
|
assert isinstance(result, str)
|
|
|
|
def test_nested_code_fence_on_same_line(self):
|
|
"""[B-30] Opening fence with code on same line."""
|
|
text = "```pythonprint('hi')```"
|
|
result = convert_markdown(
|
|
text,
|
|
code_block_formatter=lambda lang, code: f"LANG={lang}|CODE={code}",
|
|
inline_code_formatter=lambda c: c,
|
|
inline_rules=[],
|
|
)
|
|
assert isinstance(result, str)
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 5. Channel base class
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestChannelSend:
|
|
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
|
|
|
|
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
|
|
|
|
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
|
|
|
|
async def test_send_per_chat_lock_serializes(self):
|
|
"""[B-03] Per-chat locks prevent message reordering."""
|
|
|
|
ch = StubChannel()
|
|
order = []
|
|
|
|
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)
|
|
|
|
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"},
|
|
)
|
|
|
|
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
|
|
|
|
async def test_reply_to_only_on_first_chunk(self):
|
|
"""reply_to should only be passed to the first chunk."""
|
|
|
|
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:
|
|
def test_open_access_when_no_list(self):
|
|
ch = StubChannel()
|
|
assert ch.is_allowed("anyone") is True
|
|
|
|
def test_allowed_sender_passes(self):
|
|
cfg = _FakeConfig(allowed_senders=["alice", "bob"])
|
|
ch = StubChannel(cfg)
|
|
assert ch.is_allowed("alice") is True
|
|
assert ch.is_allowed("bob") is True
|
|
|
|
def test_disallowed_sender_blocked(self):
|
|
cfg = _FakeConfig(allowed_senders=["alice"])
|
|
ch = StubChannel(cfg)
|
|
assert ch.is_allowed("eve") is False
|
|
|
|
def test_composite_sender_id(self):
|
|
"""Pipe-separated composite IDs should match any component."""
|
|
cfg = _FakeConfig(allowed_senders=["12345"])
|
|
ch = StubChannel(cfg)
|
|
assert ch.is_allowed("12345|alice") is True
|
|
|
|
def test_channel_allow_list(self):
|
|
cfg = _FakeConfig(allowed_channels=["chan_1", "chan_2"])
|
|
ch = StubChannel(cfg)
|
|
assert ch.is_channel_allowed("chan_1") is True
|
|
assert ch.is_channel_allowed("chan_3") is False
|
|
|
|
def test_channel_allow_list_empty_allows_all(self):
|
|
cfg = _FakeConfig(allowed_channels=None)
|
|
ch = StubChannel(cfg)
|
|
assert ch.is_channel_allowed("any_channel") is True
|
|
|
|
|
|
class TestChannelMentionGating:
|
|
def test_dm_always_passes(self):
|
|
ch = StubChannel()
|
|
raw = RawIncoming(
|
|
sender_id="u1", chat_id="c1", text="hi", is_group=False, was_mentioned=False
|
|
)
|
|
assert ch._should_process(raw) is True
|
|
|
|
def test_group_mentioned_passes(self):
|
|
ch = StubChannel()
|
|
raw = RawIncoming(
|
|
sender_id="u1", chat_id="c1", text="hi", is_group=True, was_mentioned=True
|
|
)
|
|
assert ch._should_process(raw) is True
|
|
|
|
def test_group_not_mentioned_blocked(self):
|
|
ch = StubChannel()
|
|
ch.require_mention = "group"
|
|
raw = RawIncoming(
|
|
sender_id="u1", chat_id="c1", text="hi", is_group=True, was_mentioned=False
|
|
)
|
|
assert ch._should_process(raw) is False
|
|
|
|
def test_mention_off_passes_all(self):
|
|
ch = StubChannel()
|
|
ch.require_mention = "off"
|
|
raw = RawIncoming(
|
|
sender_id="u1", chat_id="c1", text="hi", is_group=True, was_mentioned=False
|
|
)
|
|
assert ch._should_process(raw) is True
|
|
|
|
|
|
class TestChannelBuildInbound:
|
|
def test_builds_valid_inbound(self):
|
|
ch = StubChannel()
|
|
raw = RawIncoming(
|
|
sender_id="u1",
|
|
chat_id="c1",
|
|
text="hello",
|
|
message_id="m1",
|
|
media_files=["/path/img.jpg"],
|
|
)
|
|
msg = ch._raw_to_inbound(raw)
|
|
assert msg is not None
|
|
assert msg.channel == "stub"
|
|
assert msg.sender_id == "u1"
|
|
assert msg.content == "hello"
|
|
assert msg.media == ["/path/img.jpg"]
|
|
|
|
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
|
|
|
|
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()
|
|
raw = RawIncoming(sender_id="u1", chat_id="c1", text="")
|
|
assert ch._raw_to_inbound(raw) is None
|
|
|
|
def test_media_only_message_passes(self):
|
|
ch = StubChannel()
|
|
raw = RawIncoming(
|
|
sender_id="u1",
|
|
chat_id="c1",
|
|
text="",
|
|
media_files=["/path/file.pdf"],
|
|
)
|
|
msg = ch._raw_to_inbound(raw)
|
|
assert msg is not None
|
|
assert msg.content == "[media only]"
|
|
|
|
def test_annotations_merged(self):
|
|
ch = StubChannel()
|
|
raw = RawIncoming(
|
|
sender_id="u1",
|
|
chat_id="c1",
|
|
text="main text",
|
|
content_annotations=["[attachment: photo.jpg]"],
|
|
)
|
|
msg = ch._raw_to_inbound(raw)
|
|
assert "[attachment: photo.jpg]" in msg.content
|
|
|
|
def test_metadata_preserves_chat_id(self):
|
|
ch = StubChannel()
|
|
raw = RawIncoming(
|
|
sender_id="u1", chat_id="c1", text="hi", metadata={"extra": "data"}
|
|
)
|
|
msg = ch._raw_to_inbound(raw)
|
|
assert msg.metadata["chat_id"] == "c1"
|
|
assert msg.metadata["extra"] == "data"
|
|
|
|
|
|
class TestInboundPipeline:
|
|
"""Tests for the new middleware-based inbound pipeline in _enqueue_raw()."""
|
|
|
|
async def test_pipeline_dedup(self):
|
|
"""Duplicate messages are dropped by the pipeline."""
|
|
|
|
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
|
|
|
|
async def test_pipeline_allowlist_blocks(self):
|
|
"""Non-allowed senders are blocked by the pipeline."""
|
|
|
|
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_pipeline_allowlist_passes(self):
|
|
"""Allowed senders pass through the pipeline."""
|
|
|
|
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
|
|
|
|
async def test_pipeline_channel_allowlist_blocks(self):
|
|
"""Non-allowed channels are blocked by the pipeline."""
|
|
|
|
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
|
|
|
|
async def test_pipeline_inbound_has_is_group(self):
|
|
"""InboundMessage carries is_group and was_mentioned from RawIncoming."""
|
|
|
|
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:
|
|
async def test_single_message_processed(self):
|
|
"""A single message should be published after debounce delay."""
|
|
|
|
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=f"part{i}",
|
|
message_id=f"m{i}",
|
|
metadata={"chat_id": "c1"},
|
|
)
|
|
await ch.queue_message(msg)
|
|
|
|
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
|
|
|
|
async def test_dedup_skips_duplicate(self):
|
|
"""Dedup is now handled in _enqueue_raw pipeline, not queue_message."""
|
|
|
|
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
|
|
|
|
# Only one should be enqueued (dedup catches second)
|
|
assert ch._queue.qsize() == 1
|
|
|
|
async def test_debounce_metadata_from_first_message(self):
|
|
"""[B-05] Metadata from the first message in a debounce window is kept."""
|
|
|
|
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 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"
|
|
|
|
|
|
class TestChannelTyping:
|
|
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
|
|
|
|
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")
|
|
|
|
async def test_stop_typing_idempotent(self):
|
|
ch = StubChannel()
|
|
# Should not raise even if never started
|
|
await ch.stop_typing("nonexistent")
|
|
|
|
|
|
class TestChannelReconnect:
|
|
async def test_run_reconnects_on_error(self, monkeypatch):
|
|
"""Channel.run() should reconnect with backoff on transient errors."""
|
|
|
|
ch = StubChannel()
|
|
mgr = ChannelManager(MessageBus())
|
|
mgr.register(ch)
|
|
start_count = 0
|
|
sleep_count = 0
|
|
first_retry_waiting = asyncio.Event()
|
|
allow_retries = asyncio.Event()
|
|
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 controlled_sleep(_delay):
|
|
nonlocal sleep_count
|
|
sleep_count += 1
|
|
if sleep_count == 1:
|
|
first_retry_waiting.set()
|
|
await allow_retries.wait()
|
|
|
|
ch.start = flaky_start
|
|
monkeypatch.setattr(asyncio, "sleep", controlled_sleep)
|
|
|
|
run_task = asyncio.create_task(ch.run())
|
|
await first_retry_waiting.wait()
|
|
|
|
try:
|
|
assert ch._startup_event.is_set() is False
|
|
assert ch._startup_error is None
|
|
assert mgr.startup_results() == [("stub", False, "starting (bus)")]
|
|
finally:
|
|
allow_retries.set()
|
|
|
|
await run_task
|
|
assert start_count == 3
|
|
assert ch._startup_event.is_set()
|
|
assert ch._startup_error is None
|
|
|
|
async def test_runtime_error_does_not_overwrite_successful_startup(
|
|
self, monkeypatch
|
|
):
|
|
"""A receive failure should preserve the completed startup result."""
|
|
|
|
ch = StubChannel()
|
|
start_count = 0
|
|
receive_count = 0
|
|
state_before_reconnect: list[tuple[bool, str | None]] = []
|
|
original_start = ch.start
|
|
|
|
async def tracking_start():
|
|
nonlocal start_count
|
|
start_count += 1
|
|
if start_count == 2:
|
|
state_before_reconnect.append(
|
|
(ch._startup_event.is_set(), ch._startup_error)
|
|
)
|
|
await original_start()
|
|
if start_count == 2:
|
|
ch._running = False
|
|
|
|
async def flaky_receive():
|
|
nonlocal receive_count
|
|
receive_count += 1
|
|
if receive_count == 1:
|
|
raise ConnectionError("receive transient")
|
|
if False: # pragma: no cover - marks this as an async generator
|
|
yield None
|
|
|
|
async def no_sleep(_delay):
|
|
return None
|
|
|
|
ch.start = tracking_start
|
|
ch.receive = flaky_receive
|
|
monkeypatch.setattr(asyncio, "sleep", no_sleep)
|
|
|
|
await ch.run()
|
|
|
|
assert state_before_reconnect == [(True, None)]
|
|
assert ch._startup_event.is_set()
|
|
assert ch._startup_error is None
|
|
|
|
async def test_run_stops_on_channel_error(self):
|
|
"""ChannelError should stop the channel permanently."""
|
|
|
|
ch = StubChannel()
|
|
|
|
async def fatal_start():
|
|
raise ChannelError("fatal")
|
|
|
|
ch.start = fatal_start
|
|
await ch.run()
|
|
assert ch._running is False
|
|
assert ch._startup_event.is_set()
|
|
assert ch._startup_error == "fatal"
|
|
|
|
|
|
class TestExtractRetryAfter:
|
|
def test_never_returns_none(self):
|
|
"""[B-01] Base _extract_retry_after always returns float, never None."""
|
|
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
|
|
|
|
def test_extracts_retry_after_attribute(self):
|
|
ch = StubChannel()
|
|
|
|
class RateLimitError(Exception):
|
|
retry_after = 5.0
|
|
|
|
result = ch._extract_retry_after(RateLimitError("rate limited"))
|
|
assert result == 5.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
|
|
|
|
|
|
class TestChannelAttachments:
|
|
def test_check_attachment_size_within_limit(self):
|
|
ch = StubChannel()
|
|
result = ch._check_attachment_size(1024, "small.txt")
|
|
assert result is None
|
|
|
|
def test_check_attachment_size_too_large(self):
|
|
ch = StubChannel()
|
|
result = ch._check_attachment_size(30 * 1024 * 1024, "huge.bin")
|
|
assert result is not None
|
|
assert "too large" in result
|
|
|
|
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
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 6. ChannelManager
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestChannelManagerRegister:
|
|
def test_register_and_lookup(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
mgr.register(ch)
|
|
assert mgr.get_channel("stub") is ch
|
|
assert "stub" in mgr.enabled_channels
|
|
|
|
def test_duplicate_raises(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
with pytest.raises(ValueError, match="already registered"):
|
|
mgr.register(StubChannel())
|
|
|
|
def test_register_injects_bus(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
mgr.register(ch)
|
|
assert ch._bus is bus
|
|
|
|
def test_register_applies_kwargs(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
mgr.register(ch, send_thinking=True, initial_debounce=5.0)
|
|
assert ch.send_thinking is True
|
|
assert ch.initial_debounce == 5.0
|
|
|
|
def test_health_entry_created(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
assert "stub" in mgr._health
|
|
|
|
|
|
class TestChannelManagerDispatch:
|
|
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()
|
|
|
|
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.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"
|
|
|
|
async def test_dispatch_unknown_channel_logged(self):
|
|
"""Messages to unknown channels should be logged, not crash."""
|
|
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
|
|
task = asyncio.create_task(mgr._dispatch_outbound())
|
|
await bus.publish_outbound(
|
|
OutboundMessage(
|
|
channel="nonexistent",
|
|
chat_id="c1",
|
|
content="hello",
|
|
)
|
|
)
|
|
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
|
|
|
|
async def test_dispatch_send_return_false_counts_failure(self):
|
|
"""send() returning False should mark the delivery as failed."""
|
|
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
failed_event = asyncio.Event()
|
|
|
|
async def failing_send(msg):
|
|
failed_event.set()
|
|
return False # Indicates failure
|
|
|
|
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",
|
|
)
|
|
)
|
|
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
|
|
|
|
async def test_dispatch_send_media_return_false_counts_failure(self):
|
|
"""send_media() returning False should mark the delivery as failed."""
|
|
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
failed_event = asyncio.Event()
|
|
|
|
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.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
|
|
|
|
|
|
class TestChannelManagerHealth:
|
|
def test_health_tracks_success(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
health = mgr._health["stub"]
|
|
health.total_successes = 5
|
|
health.consecutive_failures = 0
|
|
assert health.total_successes == 5
|
|
|
|
def test_health_tracks_failure(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
health = mgr._health["stub"]
|
|
health.consecutive_failures = 3
|
|
health.total_failures = 10
|
|
health.last_failure_error = "timeout"
|
|
assert health.consecutive_failures == 3
|
|
assert health.last_failure_error == "timeout"
|
|
|
|
|
|
class TestChannelManagerDynamicOps:
|
|
def test_add_channel_runtime(self):
|
|
"""[B-15] add_channel uses channel_type as key for start_times
|
|
but register() uses channel.name — potential mismatch."""
|
|
|
|
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
|
|
|
|
async def test_remove_channel(self):
|
|
"""[B-14] remove_channel removes from dict but doesn't cancel task."""
|
|
|
|
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
|
|
|
|
async def test_remove_nonexistent_channel(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
await mgr.remove_channel("ghost") # should not raise
|
|
|
|
|
|
class TestChannelManagerDrain:
|
|
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",
|
|
)
|
|
)
|
|
|
|
await mgr.stop_all()
|
|
# The drain loop should have sent it
|
|
assert len(sent) == 1
|
|
assert sent[0].content == "drain me"
|
|
|
|
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",
|
|
chat_id="c1",
|
|
content="drain me",
|
|
media=["/tmp/file.png"],
|
|
)
|
|
)
|
|
|
|
await mgr.stop_all()
|
|
|
|
assert len(sent) == 1
|
|
assert len(media_sent) == 1
|
|
|
|
assert "Outbound drain:" not in caplog.text
|
|
|
|
|
|
class TestChannelManagerTracking:
|
|
def test_record_message(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
|
|
mgr.record_message("stub", "received")
|
|
mgr.record_message("stub", "received")
|
|
mgr.record_message("stub", "sent")
|
|
|
|
assert mgr._message_counts["stub"]["received"] == 2
|
|
assert mgr._message_counts["stub"]["sent"] == 1
|
|
|
|
def test_record_message_unknown_channel(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
|
|
# Should not raise, auto-creates entry
|
|
mgr.record_message("unknown", "received")
|
|
assert mgr._message_counts["unknown"]["received"] == 1
|
|
|
|
def test_get_detailed_status(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
|
|
# Simulate start_all setting start_times
|
|
mgr._start_times["stub"] = datetime.now()
|
|
mgr._message_counts["stub"] = {"received": 5, "sent": 3}
|
|
|
|
status = mgr.get_detailed_status()
|
|
assert "stub" in status
|
|
assert status["stub"]["registered"] is True
|
|
assert status["stub"]["received"] == 5
|
|
assert status["stub"]["sent"] == 3
|
|
assert status["stub"]["uptime_seconds"] >= 0
|
|
assert status["stub"]["start_time"] is not None
|
|
|
|
def test_get_detailed_status_no_start_time(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
|
|
status = mgr.get_detailed_status()
|
|
assert status["stub"]["uptime_seconds"] == 0
|
|
assert status["stub"]["start_time"] is None
|
|
|
|
|
|
class TestChannelManagerStatus:
|
|
def test_get_status(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
status = mgr.get_status()
|
|
assert "stub" in status
|
|
assert status["stub"]["registered"] is True
|
|
|
|
def test_running_channels(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
mgr.register(ch)
|
|
assert mgr.running_channels() == []
|
|
ch._running = True
|
|
assert mgr.running_channels() == ["stub"]
|
|
|
|
def test_startup_results_report_fatal_error(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
mgr.register(ch)
|
|
ch._startup_error = "dependency missing"
|
|
ch._startup_event.set()
|
|
|
|
assert mgr.startup_results() == [("stub", False, "failed: dependency missing")]
|
|
|
|
def test_startup_results_do_not_assume_pending_channel_is_connected(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
|
|
assert mgr.startup_results() == [("stub", False, "starting (bus)")]
|
|
|
|
def test_get_stats(self):
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
stats = mgr.get_stats()
|
|
assert "channels" in stats
|
|
assert "running" in stats
|
|
assert "message_counts" in stats
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 7. InboundConsumer
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestInboundConsumer:
|
|
@staticmethod
|
|
def _make_consumer(bus=None, mgr=None, agent=None, **kw):
|
|
bus = bus or MessageBus()
|
|
if mgr is None:
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
if agent is None:
|
|
agent = MagicMock()
|
|
kw.setdefault("graph_gateway", FakeGraphGateway())
|
|
return InboundConsumer(
|
|
bus=bus,
|
|
manager=mgr,
|
|
agent=agent,
|
|
thread_id="",
|
|
max_concurrent=2,
|
|
max_pending=10,
|
|
inference_timeout=2.0,
|
|
drain_timeout=1.0,
|
|
**kw,
|
|
)
|
|
|
|
def test_session_key_format(self):
|
|
msg = BusInbound(channel="tg", sender_id="u1", chat_id="c1", content="hi")
|
|
assert msg.session_key == "tg:c1"
|
|
|
|
async def test_get_thread_id_creates_unique(self):
|
|
consumer = self._make_consumer(
|
|
graph_gateway=FakeGraphGateway(
|
|
generated_thread_ids=["thread-a", "thread-b"]
|
|
)
|
|
)
|
|
tid1 = await consumer._get_thread_id("user_a")
|
|
tid2 = await consumer._get_thread_id("user_b")
|
|
assert tid1 != tid2
|
|
|
|
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 = await consumer._get_thread_id("user_a")
|
|
tid2 = await consumer._get_thread_id("user_a")
|
|
assert tid1 == tid2
|
|
|
|
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)
|
|
mgr.register(StubChannel())
|
|
consumer = InboundConsumer(
|
|
bus=bus,
|
|
manager=mgr,
|
|
agent=MagicMock(),
|
|
thread_id="shared_thread", # Non-empty!
|
|
graph_gateway=FakeGraphGateway(),
|
|
)
|
|
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"
|
|
|
|
async def test_session_eviction_is_lru(self):
|
|
"""Sessions use LRU eviction: recently accessed senders are kept."""
|
|
consumer = self._make_consumer()
|
|
consumer._sessions.clear()
|
|
|
|
# Fill up to limit
|
|
for i in range(10):
|
|
consumer._sessions[f"user_{i}"] = f"thread_{i}"
|
|
|
|
# Access "user_0" via _get_thread_id (triggers LRU move_to_end)
|
|
await consumer._get_thread_id("user_0")
|
|
|
|
# "user_0" should now be at the end (most recently used)
|
|
oldest = next(iter(consumer._sessions))
|
|
assert oldest == "user_1" # user_1 is now the least recently used
|
|
|
|
def test_metrics_initial(self):
|
|
consumer = self._make_consumer()
|
|
m = consumer.metrics
|
|
assert m["total_processed"] == 0
|
|
assert m["total_successes"] == 0
|
|
assert m["total_failures"] == 0
|
|
assert m["total_timeouts"] == 0
|
|
|
|
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:
|
|
def test_error_message_leaks_info(self):
|
|
"""[B-22] Exception messages are sent directly to users."""
|
|
|
|
# This test documents that internal error details are exposed
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
ch = StubChannel()
|
|
mgr.register(ch)
|
|
|
|
_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
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 8. MessageBus
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestMessageBus:
|
|
"""Covers the bus as a pure pub/sub queue.
|
|
|
|
Outbound routing (subscriber dispatch, error handling, stop semantics)
|
|
is owned by ``ChannelManager._dispatch_outbound`` — see
|
|
``TestChannelManagerDispatch`` for that coverage.
|
|
"""
|
|
|
|
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",
|
|
)
|
|
)
|
|
assert bus.inbound_size == 1
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 9. Event dataclasses
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestEvents:
|
|
def test_inbound_defaults(self):
|
|
msg = BusInbound(channel="tg", sender_id="u1", chat_id="c1", content="hi")
|
|
assert msg.media == []
|
|
assert msg.metadata == {}
|
|
assert msg.session_key == "tg:c1"
|
|
assert isinstance(msg.timestamp, datetime)
|
|
|
|
def test_outbound_defaults(self):
|
|
msg = BusOutbound(channel="tg", chat_id="c1", content="reply")
|
|
assert msg.reply_to is None
|
|
assert msg.media == []
|
|
assert msg.metadata == {}
|
|
|
|
def test_inbound_sender_alias(self):
|
|
msg = InboundMessage(channel="x", sender_id="u1", chat_id="c1", content="hi")
|
|
assert msg.sender == "u1"
|
|
|
|
def test_outbound_recipient_alias(self):
|
|
msg = OutboundMessage(channel="x", chat_id="c1", content="hi")
|
|
assert msg.recipient == "c1"
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 10. Integration scenarios
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestIntegration:
|
|
async def test_full_inbound_pipeline(self):
|
|
"""Raw message → build_inbound → queue_message → bus."""
|
|
|
|
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)
|
|
|
|
# _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 _flush_debounce(ch, "user1")
|
|
assert bus.inbound_size == 1
|
|
|
|
async def test_outbound_dispatch_with_media(self):
|
|
"""Dispatch routes media alongside text content."""
|
|
|
|
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"],
|
|
)
|
|
)
|
|
await asyncio.wait_for(media_event.wait(), timeout=1.0)
|
|
task.cancel()
|
|
try:
|
|
await task
|
|
except asyncio.CancelledError:
|
|
pass
|
|
|
|
assert len(media_sent) == 1
|
|
|
|
async def test_debounce_lost_on_stop(self):
|
|
"""Buffered messages should be flushed when stop() is called."""
|
|
|
|
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
|
|
|
|
assert len(ch._message_buffers) == 1
|
|
|
|
# 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"
|
|
|
|
async def test_send_locks_bounded_growth(self):
|
|
"""_send_locks stays bounded via LRU eviction of unlocked entries."""
|
|
|
|
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
|
|
|
|
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
# 11. Edge cases and boundary conditions
|
|
# ═══════════════════════════════════════════════════════════════════
|
|
|
|
|
|
class TestEdgeCases:
|
|
def test_chunk_text_single_char_limit(self):
|
|
chunks = chunk_text("abc", 1)
|
|
assert all(len(c) <= 1 for c in chunks)
|
|
assert len(chunks) == 3
|
|
|
|
def test_chunk_text_unicode(self):
|
|
text = "你好世界" * 100
|
|
chunks = chunk_text(text, 50)
|
|
assert all(len(c) <= 50 for c in chunks)
|
|
|
|
def test_dedup_cache_rapid_same_id(self):
|
|
dc = DedupCache()
|
|
assert dc.is_duplicate("x") is False
|
|
for _ in range(100):
|
|
assert dc.is_duplicate("x") is True
|
|
|
|
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")
|
|
assert raw.text == ""
|
|
assert raw.media_files == []
|
|
assert raw.content_annotations == []
|
|
assert raw.is_group is False
|
|
assert raw.was_mentioned is True
|
|
assert raw.message_id == ""
|
|
|
|
def test_outbound_message_no_metadata_chat_id_resolution(self):
|
|
"""resolve_chat_id falls back to recipient when metadata has no chat_id."""
|
|
ch = StubChannel()
|
|
msg = OutboundMessage(
|
|
channel="stub",
|
|
chat_id="fallback_id",
|
|
content="hi",
|
|
metadata={},
|
|
)
|
|
resolved = ch._resolve_chat_id(msg)
|
|
assert resolved == "fallback_id"
|
|
|
|
def test_health_server_response_structure(self):
|
|
"""HealthServer builds response with expected keys."""
|
|
from EvoScientist.channels.channel_manager import _HealthServer
|
|
|
|
bus = MessageBus()
|
|
mgr = ChannelManager(bus)
|
|
mgr.register(StubChannel())
|
|
|
|
hs = _HealthServer(mgr, 0)
|
|
resp = hs._build_response()
|
|
assert resp["status"] == "healthy"
|
|
assert "uptime_seconds" in resp
|
|
assert "channels" in resp
|
|
assert "queues" in resp
|
|
assert "health" in resp
|