Stable
This commit is contained in:
@@ -176,24 +176,31 @@ async def download_attachment(
|
||||
local_path = media_path(f"{prefix}{safe_name}")
|
||||
|
||||
async with httpx.AsyncClient(proxy=proxy) as client:
|
||||
resp = await client.get(url, headers=headers or {}, timeout=30)
|
||||
if resp.status_code != 200:
|
||||
return None, f"[attachment: {filename} - download failed]"
|
||||
async with client.stream("GET", url, headers=headers or {}, timeout=30) as resp:
|
||||
if resp.status_code != 200:
|
||||
return None, f"[attachment: {filename} - download failed]"
|
||||
|
||||
# Check Content-Length when file_size was not known beforehand
|
||||
if file_size is None:
|
||||
cl = resp.headers.get("content-length")
|
||||
if cl:
|
||||
try:
|
||||
too_large = check_attachment_size(int(cl), filename)
|
||||
if too_large:
|
||||
return None, too_large
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
if len(resp.content) > MAX_ATTACHMENT_BYTES:
|
||||
return None, check_attachment_size(len(resp.content), filename)
|
||||
# Check Content-Length header before downloading body
|
||||
if file_size is None:
|
||||
cl = resp.headers.get("content-length")
|
||||
if cl:
|
||||
try:
|
||||
too_large = check_attachment_size(int(cl), filename)
|
||||
if too_large:
|
||||
return None, too_large
|
||||
except (ValueError, TypeError):
|
||||
pass
|
||||
|
||||
local_path.write_bytes(resp.content)
|
||||
# Stream body with incremental size check
|
||||
chunks: list[bytes] = []
|
||||
total = 0
|
||||
async for chunk in resp.aiter_bytes():
|
||||
total += len(chunk)
|
||||
if total > MAX_ATTACHMENT_BYTES:
|
||||
return None, check_attachment_size(total, filename)
|
||||
chunks.append(chunk)
|
||||
|
||||
local_path.write_bytes(b"".join(chunks))
|
||||
return str(local_path), f"[attachment: {local_path}]"
|
||||
except Exception as e:
|
||||
_logger.warning(f"Failed to download attachment: {e}")
|
||||
@@ -670,7 +677,12 @@ class Channel(ChannelPlugin, ABC):
|
||||
|
||||
def _should_process(self, raw: RawIncoming) -> bool:
|
||||
"""Decide whether to process a message based on mention gating."""
|
||||
if not raw.is_group or self.require_mention == "off":
|
||||
if self.require_mention == "off":
|
||||
return True
|
||||
if self.require_mention == "always":
|
||||
return raw.was_mentioned
|
||||
# "group" — require mention only in groups
|
||||
if not raw.is_group:
|
||||
return True
|
||||
return raw.was_mentioned
|
||||
|
||||
|
||||
@@ -199,7 +199,7 @@ SIGNAL = ChannelCapabilities(
|
||||
|
||||
EMAIL = ChannelCapabilities(
|
||||
format_type="html",
|
||||
max_text_length=0, # no practical limit
|
||||
max_text_length=999_999, # no practical limit
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
html=True,
|
||||
@@ -208,7 +208,7 @@ EMAIL = ChannelCapabilities(
|
||||
|
||||
IMESSAGE = ChannelCapabilities(
|
||||
format_type="plain",
|
||||
max_text_length=0,
|
||||
max_text_length=999_999,
|
||||
typing=False, # Apple does not expose typing indicator API
|
||||
media_send=True,
|
||||
media_receive=True,
|
||||
|
||||
@@ -346,7 +346,10 @@ def build_inbound_pipeline(
|
||||
))
|
||||
|
||||
if plugin.pairing is not None:
|
||||
middlewares.append(PairingMiddleware(channel_name=plugin.id))
|
||||
middlewares.append(PairingMiddleware(
|
||||
channel_name=plugin.id,
|
||||
dm_policy=dm_policy,
|
||||
))
|
||||
|
||||
if caps.groups:
|
||||
middlewares.append(GroupHistoryMiddleware())
|
||||
|
||||
@@ -158,7 +158,10 @@ class InboundConsumer:
|
||||
# Evict oldest entry
|
||||
oldest = next(iter(self._sessions))
|
||||
del self._sessions[oldest]
|
||||
self._sessions[sender_id] = self.thread_id or str(uuid.uuid4())
|
||||
if self.thread_id:
|
||||
self._sessions[sender_id] = f"{self.thread_id}:{sender_id}"
|
||||
else:
|
||||
self._sessions[sender_id] = str(uuid.uuid4())
|
||||
return self._sessions[sender_id]
|
||||
|
||||
def _get_channel(self, channel_name: str) -> Channel | None:
|
||||
@@ -361,7 +364,7 @@ class InboundConsumer:
|
||||
await self.bus.publish_outbound(OutboundMessage(
|
||||
channel=msg.channel,
|
||||
chat_id=msg.chat_id,
|
||||
content=f"Error: {e}",
|
||||
content="Sorry, something went wrong. Please try again later.",
|
||||
metadata=msg.metadata,
|
||||
))
|
||||
finally:
|
||||
|
||||
@@ -511,6 +511,12 @@ class FormattingMiddleware(OutboundMiddlewareBase):
|
||||
"""Convert text to channel format."""
|
||||
return self._formatter.format(text)
|
||||
|
||||
async def process_outbound(
|
||||
self, message: OutboundMessage, context: dict[str, Any],
|
||||
) -> OutboundMessage | None:
|
||||
formatted = self._formatter.format(message.content)
|
||||
return dataclasses.replace(message, content=formatted)
|
||||
|
||||
|
||||
# ── Retry ────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -662,11 +668,13 @@ class MentionGatingMiddleware(InboundMiddleware):
|
||||
return raw
|
||||
|
||||
def _should_process(self, raw: RawIncoming) -> bool:
|
||||
if not raw.is_group or self.require_mention == "off":
|
||||
if self.require_mention == "off":
|
||||
return True
|
||||
if self.require_mention == "always":
|
||||
return raw.was_mentioned
|
||||
# "group" — require mention in groups
|
||||
# "group" — require mention only in groups
|
||||
if not raw.is_group:
|
||||
return True
|
||||
return raw.was_mentioned
|
||||
|
||||
|
||||
@@ -779,10 +787,12 @@ class PairingMiddleware(InboundMiddleware):
|
||||
self,
|
||||
channel_name: str,
|
||||
send_response_fn: Callable[[str, str], Any] | None = None,
|
||||
dm_policy: str = "allowlist",
|
||||
) -> None:
|
||||
self._manager = PairingManager()
|
||||
self._channel_name = channel_name
|
||||
self._send_response_fn = send_response_fn
|
||||
self._dm_policy = dm_policy
|
||||
|
||||
async def process_inbound(
|
||||
self, raw: RawIncoming, context: dict[str, Any],
|
||||
@@ -790,8 +800,7 @@ class PairingMiddleware(InboundMiddleware):
|
||||
if raw.is_group:
|
||||
return raw # pairing only applies to DMs
|
||||
|
||||
dm_policy = context.get("dm_policy", "allowlist")
|
||||
if dm_policy != "pairing":
|
||||
if self._dm_policy != "pairing":
|
||||
return raw
|
||||
|
||||
if self._manager.is_approved(self._channel_name, raw.sender_id):
|
||||
|
||||
Binary file not shown.
@@ -1108,7 +1108,7 @@ class TestInboundConsumer:
|
||||
assert tid1 == tid2
|
||||
|
||||
def test_shared_thread_id_bug(self):
|
||||
"""[B-20] If thread_id is non-empty, all senders share the same session."""
|
||||
"""[B-20] If thread_id is non-empty, senders get unique thread IDs with shared prefix."""
|
||||
bus = MessageBus()
|
||||
mgr = ChannelManager(bus)
|
||||
mgr.register(StubChannel())
|
||||
@@ -1118,8 +1118,10 @@ class TestInboundConsumer:
|
||||
)
|
||||
tid1 = consumer._get_thread_id("alice")
|
||||
tid2 = consumer._get_thread_id("bob")
|
||||
# BUG: Both get the same thread_id
|
||||
assert tid1 == tid2 == "shared_thread"
|
||||
# Fixed: Each sender gets a unique thread_id using thread_id as prefix
|
||||
assert tid1 != tid2
|
||||
assert tid1 == "shared_thread:alice"
|
||||
assert tid2 == "shared_thread:bob"
|
||||
|
||||
def test_session_eviction_is_fifo_not_lru(self):
|
||||
"""[B-19] Sessions evict oldest by insertion, not by access."""
|
||||
|
||||
Reference in New Issue
Block a user