Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 39 additions & 6 deletions gateway/platforms/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -2104,6 +2104,37 @@ def text(self) -> str:
return str.__str__(self)


def _refresh_merged_event_trigger_metadata(
existing: MessageEvent,
event: MessageEvent,
) -> None:
"""Align a merged event's trigger and reply context with its newest input."""
# Synthetic/internal follow-ups may not have a platform message id; keep
# the last concrete anchor in that case. Reply context is still replaced
# so an unquoted newer message cannot inherit an older quoted-message
# prompt prefix.
if event.message_id is not None:
existing.message_id = event.message_id
if event.source is not None:
existing.source = event.source
existing.raw_message = event.raw_message
existing.platform_update_id = event.platform_update_id
existing.reply_to_message_id = event.reply_to_message_id
existing.reply_to_text = event.reply_to_text
existing.reply_to_author_id = event.reply_to_author_id
existing.reply_to_author_name = event.reply_to_author_name
existing.reply_to_is_own_message = event.reply_to_is_own_message
existing.auto_skill = event.auto_skill
existing.channel_prompt = event.channel_prompt
existing.channel_context = event.channel_context
# A merged event may bypass authorization only when every constituent was
# already trusted as internal. Never let an internal notification elevate
# a user-originated follow-up (or vice versa).
existing.internal = existing.internal and event.internal
existing.metadata = {**(existing.metadata or {}), **(event.metadata or {})}
existing.timestamp = event.timestamp


def merge_pending_message_event(
pending_messages: Dict[str, MessageEvent],
session_key: str,
Expand Down Expand Up @@ -2134,6 +2165,7 @@ def merge_pending_message_event(
existing.media_types.extend(event.media_types)
if event.text:
existing.text = BasePlatformAdapter._merge_caption(existing.text, event.text)
_refresh_merged_event_trigger_metadata(existing, event)
return

if existing_has_media or incoming_has_media:
Expand All @@ -2152,6 +2184,7 @@ def merge_pending_message_event(
and event.message_type != MessageType.TEXT
):
existing.message_type = event.message_type
_refresh_merged_event_trigger_metadata(existing, event)
return

if (
Expand All @@ -2161,6 +2194,11 @@ def merge_pending_message_event(
):
if event.text:
existing.text = f"{existing.text}\n{event.text}" if existing.text else event.text
# The merged event becomes the next turn's trigger. Keep its reply
# target and injected reply context aligned with the newest input;
# otherwise Telegram can visibly quote one message while the model
# receives stale quoted-message context from another.
_refresh_merged_event_trigger_metadata(existing, event)
return

pending_messages[session_key] = event
Expand Down Expand Up @@ -4347,12 +4385,7 @@ async def _queue_text_debounce(self, session_key: str, event: MessageEvent) -> N
if state.event.text
else event.text
)
latest_message_id = getattr(event, "message_id", None)
latest_anchor = latest_message_id or getattr(event, "reply_to_message_id", None)
if latest_message_id is not None:
state.event.message_id = str(latest_message_id)
if latest_anchor is not None and hasattr(state.event, "reply_to_message_id"):
state.event.reply_to_message_id = str(latest_anchor)
_refresh_merged_event_trigger_metadata(state.event, event)
state.last_ts = now

if state.task is not None and not state.task.done():
Expand Down
4 changes: 4 additions & 0 deletions plugins/platforms/telegram/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -266,6 +266,7 @@ class _MockContextTypes:
SUPPORTED_DOCUMENT_TYPES,
SUPPORTED_IMAGE_DOCUMENT_TYPES,
_TEXT_INJECT_EXTENSIONS,
_refresh_merged_event_trigger_metadata,
utf16_len,
)
from plugins.platforms.telegram.telegram_ids import (
Expand Down Expand Up @@ -8191,6 +8192,7 @@ def _enqueue_text_event(self, event: MessageEvent) -> None:
if event.media_urls:
existing.media_urls.extend(event.media_urls)
existing.media_types.extend(event.media_types)
_refresh_merged_event_trigger_metadata(existing, event)

# Cancel any pending flush and restart the timer
prior_task = self._pending_text_batch_tasks.get(key)
Expand Down Expand Up @@ -8295,6 +8297,7 @@ def _enqueue_photo_event(self, batch_key: str, event: MessageEvent) -> None:
existing.media_types.extend(event.media_types)
if event.text:
existing.text = self._merge_caption(existing.text, event.text)
_refresh_merged_event_trigger_metadata(existing, event)

prior_task = self._pending_photo_batch_tasks.get(batch_key)
if prior_task and not prior_task.done():
Expand Down Expand Up @@ -8603,6 +8606,7 @@ async def _queue_media_group_event(self, media_group_id: str, event: MessageEven
existing.media_types.extend(event.media_types)
if event.text:
existing.text = self._merge_caption(existing.text, event.text)
_refresh_merged_event_trigger_metadata(existing, event)

prior_task = self._media_group_tasks.get(media_group_id)
if prior_task:
Expand Down
53 changes: 52 additions & 1 deletion tests/gateway/test_active_session_text_merge.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,12 @@ def _make_event(
user_id: str = "u1",
user_name: str | None = None,
thread_id: str | None = None,
message_id: str | None = None,
reply_to_message_id: str | None = None,
reply_to_text: str | None = None,
reply_to_author_id: str | None = None,
reply_to_author_name: str | None = None,
reply_to_is_own_message: bool = False,
) -> MessageEvent:
source = SessionSource(
platform=Platform.TELEGRAM,
Expand All @@ -60,7 +66,12 @@ def _make_event(
text=text,
message_type=MessageType.TEXT,
source=source,
message_id=f"msg-{text[:8]}",
message_id=message_id or f"msg-{text[:8]}",
reply_to_message_id=reply_to_message_id,
reply_to_text=reply_to_text,
reply_to_author_id=reply_to_author_id,
reply_to_author_name=reply_to_author_name,
reply_to_is_own_message=reply_to_is_own_message,
)


Expand Down Expand Up @@ -155,6 +166,46 @@ async def test_debounce_buffers_rapid_text_then_flushes_to_pending():
assert adapter._pending_messages[session_key].text == "part two\npart three"


@pytest.mark.asyncio
async def test_debounce_flush_keeps_latest_trigger_and_reply_context_aligned():
adapter = _make_adapter()
adapter._busy_text_debounce_seconds = 1.0

first = _make_event(
"first follow-up",
message_id="100",
reply_to_message_id="90",
reply_to_text="old quote",
reply_to_author_id="old-author",
reply_to_author_name="Old",
reply_to_is_own_message=True,
)
latest = _make_event(
"latest follow-up",
message_id="101",
reply_to_message_id="99",
reply_to_text="latest quote",
reply_to_author_id="new-author",
reply_to_author_name="New",
reply_to_is_own_message=False,
)
session_key = build_session_key(first.source)
adapter._active_sessions[session_key] = asyncio.Event()

await adapter.handle_message(first)
await adapter.handle_message(latest)
assert await adapter._flush_text_debounce_now(session_key) is True

pending = adapter._pending_messages[session_key]
assert pending.text == "first follow-up\nlatest follow-up"
assert pending.message_id == "101"
assert pending.reply_to_message_id == "99"
assert pending.reply_to_text == "latest quote"
assert pending.reply_to_author_id == "new-author"
assert pending.reply_to_author_name == "New"
assert pending.reply_to_is_own_message is False


@pytest.mark.asyncio
async def test_debounce_resets_timer_on_new_arrival():
adapter = _make_adapter()
Expand Down
144 changes: 143 additions & 1 deletion tests/gateway/test_session_race_guard.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,12 @@
import pytest

from gateway.config import GatewayConfig, Platform, PlatformConfig
from gateway.platforms.base import MessageEvent, MessageType, merge_pending_message_event
from gateway.platforms.base import (
MessageEvent,
MessageType,
_reply_anchor_for_event,
merge_pending_message_event,
)
from gateway.run import GatewayRunner, _AGENT_PENDING_SENTINEL
from gateway.session import SessionSource, build_session_key

Expand Down Expand Up @@ -197,6 +202,134 @@ async def slow_inner(self_inner, ev, src, qk, generation):
await task1


def test_merge_pending_text_followups_reply_to_latest_message():
"""A merged busy-turn reply must quote the newest user message."""
pending = {}
source = SessionSource(
platform=Platform.TELEGRAM,
chat_id="12345",
chat_type="dm",
user_id="u1",
thread_id="34457",
)
session_key = build_session_key(source)
first = MessageEvent(
text="first follow-up",
message_type=MessageType.TEXT,
source=source,
message_id="100",
reply_to_message_id="50",
reply_to_text="older quoted message",
reply_to_author_id="old-author",
reply_to_author_name="Old Author",
reply_to_is_own_message=True,
)
latest = MessageEvent(
text="latest follow-up",
message_type=MessageType.TEXT,
source=source,
message_id="101",
reply_to_message_id="60",
reply_to_text="latest quoted message",
reply_to_author_id="new-author",
reply_to_author_name="New Author",
reply_to_is_own_message=False,
)

merge_pending_message_event(pending, session_key, first, merge_text=True)
merge_pending_message_event(pending, session_key, latest, merge_text=True)

merged = pending[session_key]
assert merged.text == "first follow-up\nlatest follow-up"
assert merged.message_id == "101"
assert _reply_anchor_for_event(merged) == "101"
assert merged.reply_to_message_id == "60"
assert merged.reply_to_text == "latest quoted message"
assert merged.reply_to_author_id == "new-author"
assert merged.reply_to_author_name == "New Author"
assert merged.reply_to_is_own_message is False


def test_merge_pending_text_followup_clears_stale_reply_context():
pending = {}
source = SessionSource(platform=Platform.TELEGRAM, chat_id="12345", chat_type="dm")
first = MessageEvent(
text="first",
source=source,
message_id="100",
reply_to_message_id="50",
reply_to_text="stale quote",
reply_to_author_id="old-author",
reply_to_author_name="Old Author",
reply_to_is_own_message=True,
)
latest = MessageEvent(text="latest", source=source, message_id="101")

merge_pending_message_event(pending, "session", first, merge_text=True)
merge_pending_message_event(pending, "session", latest, merge_text=True)

merged = pending["session"]
assert merged.message_id == "101"
assert merged.reply_to_message_id is None
assert merged.reply_to_text is None
assert merged.reply_to_author_id is None
assert merged.reply_to_author_name is None
assert merged.reply_to_is_own_message is False


def test_merge_pending_text_followup_without_id_preserves_existing_anchor():
pending = {}
source = SessionSource(platform=Platform.TELEGRAM, chat_id="12345", chat_type="dm")
first = MessageEvent(text="first", source=source, message_id="100")
synthetic = MessageEvent(text="synthetic", source=source, message_id=None)

merge_pending_message_event(pending, "session", first, merge_text=True)
merge_pending_message_event(pending, "session", synthetic, merge_text=True)

assert pending["session"].message_id == "100"


def test_merge_pending_refreshes_latest_event_provenance_without_privilege_escalation():
pending = {}
source = SessionSource(platform=Platform.TELEGRAM, chat_id="12345", chat_type="dm")
first_raw = object()
latest_raw = object()
first = MessageEvent(
text="internal completion",
source=source,
raw_message=first_raw,
message_id="100",
platform_update_id=10,
channel_prompt="old prompt",
channel_context="old context",
internal=True,
metadata={"old": True, "shared": "old"},
)
latest = MessageEvent(
text="user follow-up",
source=source,
raw_message=latest_raw,
message_id="101",
platform_update_id=11,
channel_prompt="new prompt",
channel_context="new context",
internal=False,
metadata={"new": True, "shared": "new"},
)

merge_pending_message_event(pending, "session", first, merge_text=True)
merge_pending_message_event(pending, "session", latest, merge_text=True)

merged = pending["session"]
assert merged.raw_message is latest_raw
assert merged.platform_update_id == 11
assert merged.timestamp == latest.timestamp
assert merged.channel_prompt == "new prompt"
assert merged.channel_context == "new context"
assert merged.metadata == {"old": True, "new": True, "shared": "new"}
assert merged.internal is False


def test_merge_pending_message_event_merges_text_and_photo_followups():
pending = {}
source = SessionSource(
Expand All @@ -211,11 +344,17 @@ def test_merge_pending_message_event_merges_text_and_photo_followups():
text="first follow-up",
message_type=MessageType.TEXT,
source=source,
message_id="100",
reply_to_message_id="90",
reply_to_text="old quote",
)
photo_event = MessageEvent(
text="see screenshot",
message_type=MessageType.PHOTO,
source=source,
message_id="101",
reply_to_message_id="99",
reply_to_text="latest quote",
media_urls=["/tmp/test.png"],
media_types=["image/png"],
)
Expand All @@ -228,6 +367,9 @@ def test_merge_pending_message_event_merges_text_and_photo_followups():
assert merged.text == "first follow-up\n\nsee screenshot"
assert merged.media_urls == ["/tmp/test.png"]
assert merged.media_types == ["image/png"]
assert merged.message_id == "101"
assert merged.reply_to_message_id == "99"
assert merged.reply_to_text == "latest quote"


def test_merge_pending_message_event_promotes_document_followups_over_text():
Expand Down
Loading