This commit is contained in:
MuXinCG
2026-02-15 17:21:29 +08:00
parent 7a94951af9
commit e672f71c12
7 changed files with 58 additions and 29 deletions
+29 -17
View File
@@ -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
+2 -2
View File
@@ -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,
+4 -1
View File
@@ -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())
+5 -2
View File
@@ -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:
+13 -4
View File
@@ -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.
+5 -3
View File
@@ -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."""