From 745de4c8327ffe2a6fddfc7a8faa26c5b5d929cb Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:49:12 -0700 Subject: [PATCH] refactor(gateway/platforms): simplify helpers.py (strip rules table, unified mention compiler, chunker cleanups), media_cache.py, _http_client_limits.py, __init__.py --- gateway/platforms/__init__.py | 9 +- gateway/platforms/_http_client_limits.py | 30 +- gateway/platforms/_shared.py | 1 - gateway/platforms/helpers.py | 405 ++++++++--------------- gateway/platforms/media_cache.py | 191 +++-------- 5 files changed, 194 insertions(+), 442 deletions(-) diff --git a/gateway/platforms/__init__.py b/gateway/platforms/__init__.py index 6276791f55..2f69df929a 100644 --- a/gateway/platforms/__init__.py +++ b/gateway/platforms/__init__.py @@ -5,13 +5,7 @@ from .base import BasePlatformAdapter, MessageEvent, SendResult # QQAdapter / YuanbaoAdapter are exposed lazily (PEP 562 ``__getattr__``): eager # imports cost ~48 ms / ~8 MB RSS on every CLI invocation and nothing in-tree # imports them from the package root. -__all__ = [ - "BasePlatformAdapter", - "MessageEvent", - "SendResult", - "QQAdapter", - "YuanbaoAdapter", -] +__all__ = ["BasePlatformAdapter", "MessageEvent", "SendResult", "QQAdapter", "YuanbaoAdapter"] _LAZY_ADAPTERS = {"QQAdapter": ".qqbot", "YuanbaoAdapter": ".yuanbao"} @@ -21,7 +15,6 @@ def __getattr__(name): if module is None: raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from importlib import import_module - return getattr(import_module(module, __name__), name) diff --git a/gateway/platforms/_http_client_limits.py b/gateway/platforms/_http_client_limits.py index 489df6d953..d24f532129 100644 --- a/gateway/platforms/_http_client_limits.py +++ b/gateway/platforms/_http_client_limits.py @@ -1,15 +1,11 @@ """Shared HTTP client factory for long-lived platform adapters. -Persistent ``httpx.AsyncClient`` pools amortise TLS setup, but httpx's default -``keepalive_expiry`` (5s) lets peer-initiated FIN sit in ``CLOSE_WAIT`` behind -transparent proxies (macOS + Cloudflare Warp) — multiplied across 7 adapters -plus LLM/MCP clients that walks into the default 256 fd limit (#18451). -``platform_httpx_limits()`` returns a tighter ``httpx.Limits``: -``max_keepalive_connections=10`` (platform APIs rarely parallelise beyond -this), ``keepalive_expiry=2.0`` (close idle sockets aggressively). - -Override via ``HERMES_GATEWAY_HTTPX_KEEPALIVE_EXPIRY`` / -``HERMES_GATEWAY_HTTPX_MAX_KEEPALIVE`` env vars when tuning under load. +httpx's default ``keepalive_expiry`` (5s) lets peer-initiated FIN sit in +``CLOSE_WAIT`` behind transparent proxies (macOS + Cloudflare Warp); across 7 +adapters plus LLM/MCP clients that walks into the default 256 fd limit. +``platform_httpx_limits()`` returns tighter ``httpx.Limits``: 10 keepalive +connections (platform APIs rarely parallelise beyond this), 2.0s expiry. +Override via ``HERMES_GATEWAY_HTTPX_KEEPALIVE_EXPIRY`` / ``HERMES_GATEWAY_HTTPX_MAX_KEEPALIVE``. """ from __future__ import annotations @@ -39,19 +35,13 @@ def _positive_env(name: str, default, cast): def platform_httpx_limits() -> "httpx.Limits | None": - """``httpx.Limits`` tuned for persistent platform-adapter clients. - - Returns ``None`` when httpx isn't importable so callers can fall back to - httpx's built-in default without a hard dependency on this helper. - """ + """``httpx.Limits`` tuned for persistent platform-adapter clients; ``None`` without httpx.""" if httpx is None: return None + # max_connections stays at the httpx default (100) — plenty of headroom. return httpx.Limits( max_keepalive_connections=_positive_env( - "HERMES_GATEWAY_HTTPX_MAX_KEEPALIVE", _DEFAULT_MAX_KEEPALIVE, int - ), - # max_connections stays at the httpx default (100) — plenty of headroom. + "HERMES_GATEWAY_HTTPX_MAX_KEEPALIVE", _DEFAULT_MAX_KEEPALIVE, int), keepalive_expiry=_positive_env( - "HERMES_GATEWAY_HTTPX_KEEPALIVE_EXPIRY", _DEFAULT_KEEPALIVE_EXPIRY_S, float - ), + "HERMES_GATEWAY_HTTPX_KEEPALIVE_EXPIRY", _DEFAULT_KEEPALIVE_EXPIRY_S, float), ) diff --git a/gateway/platforms/_shared.py b/gateway/platforms/_shared.py index 56cb988b4a..6c92935c64 100644 --- a/gateway/platforms/_shared.py +++ b/gateway/platforms/_shared.py @@ -39,7 +39,6 @@ def profile_scoped() -> bool: """ try: from agent.secret_scope import current_secret_scope, is_multiplex_active - return bool(is_multiplex_active() and current_secret_scope() is not None) except Exception: return False diff --git a/gateway/platforms/helpers.py b/gateway/platforms/helpers.py index 071a6e4362..4d6ed53fe6 100644 --- a/gateway/platforms/helpers.py +++ b/gateway/platforms/helpers.py @@ -7,21 +7,17 @@ import logging import re import time from pathlib import Path -from typing import Dict from utils import atomic_json_write logger = logging.getLogger(__name__) -# ─── Message Deduplication ──────────────────────────────────────────────────── - - class MessageDeduplicator: """TTL-based message deduplication cache (``if dedup.is_duplicate(msg_id): return``).""" def __init__(self, max_size: int = 2000, ttl_seconds: float = 300): - self._seen: Dict[str, float] = {} + self._seen: dict[str, float] = {} self._max_size = max_size self._ttl = ttl_seconds @@ -33,8 +29,7 @@ class MessageDeduplicator: if msg_id in self._seen: if now - self._seen[msg_id] < self._ttl: return True - # Entry has expired — remove it and treat as new - del self._seen[msg_id] + del self._seen[msg_id] # expired: treat as new self._seen[msg_id] = now if len(self._seen) > self._max_size: cutoff = now - self._ttl @@ -46,9 +41,7 @@ class MessageDeduplicator: def contains(self, msg_id: str) -> bool: """Return whether *msg_id* is live in the cache without inserting it.""" - if not msg_id: - return False - seen_at = self._seen.get(msg_id) + seen_at = self._seen.get(msg_id) if msg_id else None if seen_at is None: return False if time.time() - seen_at < self._ttl: @@ -61,46 +54,34 @@ class MessageDeduplicator: self._seen.pop(msg_id, None) def clear(self): - """Clear all tracked messages.""" self._seen.clear() -# ─── Markdown Stripping ────────────────────────────────────────────────────── - -# Pre-compiled regexes for performance -_RE_BOLD = re.compile(r"\*\*(.+?)\*\*", re.DOTALL) -_RE_ITALIC_STAR = re.compile(r"\*(.+?)\*", re.DOTALL) -_RE_BOLD_UNDER = re.compile(r"\b__(?![\s_])(.+?)(? str: """Strip markdown formatting for plain-text platforms (SMS, iMessage, etc.).""" - text = _RE_BOLD.sub(r"\1", text) - text = _RE_ITALIC_STAR.sub(r"\1", text) - text = _RE_BOLD_UNDER.sub(r"\1", text) - text = _RE_ITALIC_UNDER.sub(r"\1", text) - text = _RE_CODE_BLOCK.sub("", text) - text = _RE_INLINE_CODE.sub(r"\1", text) - text = _RE_HEADING.sub("", text) - text = _RE_LINK.sub(r"\1", text) - text = _RE_MULTI_NEWLINE.sub("\n\n", text) + for pattern, repl in _STRIP_RULES: + text = pattern.sub(repl, text) return text.strip() -# ─── Thread Participation Tracking ─────────────────────────────────────────── - - class ThreadParticipationTracker: - """Persistent set of threads the bot has participated in (``_threads.json``). - - ``thread_id in tracker`` checks membership; ``tracker.mark(thread_id)`` persists. - """ + """Persistent set of threads the bot has participated in (``_threads.json``); + ``thread_id in tracker`` checks membership, ``tracker.mark(thread_id)`` persists.""" _MAX_TRACKED = 500 @@ -114,23 +95,18 @@ class ThreadParticipationTracker: return get_hermes_home() / f"{self._platform}_threads.json" def _load(self) -> list[str]: - path = self._state_path() - if path.exists(): - try: - data = json.loads(path.read_text(encoding="utf-8")) - if isinstance(data, list): - return [str(thread_id) for thread_id in data] - except Exception: - pass - return [] + try: + data = json.loads(self._state_path().read_text(encoding="utf-8")) + except Exception: + return [] + return [str(thread_id) for thread_id in data] if isinstance(data, list) else [] def _save(self) -> None: - path = self._state_path() thread_list = list(self._threads) if len(thread_list) > self._max_tracked: thread_list = thread_list[-self._max_tracked:] self._threads = dict.fromkeys(thread_list) - atomic_json_write(path, thread_list, indent=None) + atomic_json_write(self._state_path(), thread_list, indent=None) def mark(self, thread_id: str) -> None: """Mark *thread_id* as participated and persist.""" @@ -145,9 +121,6 @@ class ThreadParticipationTracker: self._threads.clear() -# ─── Phone Number Redaction ────────────────────────────────────────────────── - - def redact_phone(phone: str) -> str: """Redact a phone number for logging, preserving country code and last 4.""" if not phone: @@ -157,19 +130,16 @@ def redact_phone(phone: str) -> str: return phone[:4] + "****" + phone[-4:] -# ─── GFM Markdown Table → Bullet Conversion ───────────────────────────────── -# Discord calls convert_table_to_bullets(); Telegram imports the primitives but -# keeps its own MarkdownV2-aware renderer. - -# GFM delimiter row: optional outer pipes, dash cells (optional alignment colons). -# Requires at least one internal '|' so a lone '---' rule is NOT matched. +# ─── GFM table → bullets. Discord calls convert_table_to_bullets(); Telegram imports the +# primitives but keeps its own MarkdownV2-aware renderer. +# Delimiter row: optional outer pipes, dash cells (optional alignment colons). Requires at +# least one internal '|' so a lone '---' rule is NOT matched. TABLE_SEPARATOR_RE = re.compile(r'^\s*\|?\s*:?-+:?\s*(?:\|\s*:?-+:?\s*){1,}\|?\s*$') def is_table_row(line: str) -> bool: """Return True if *line* could plausibly be a table data row.""" - stripped = line.strip() - return bool(stripped) and '|' in stripped + return '|' in line.strip() def split_markdown_table_row(line: str) -> list[str]: @@ -179,18 +149,12 @@ def split_markdown_table_row(line: str) -> list[str]: def _render_table_block(table_block: list[str]) -> str: - """Render a detected GFM table as bold-heading + bullet groups. - - Same alignment logic as Telegram's renderer: without a row-label column the - full row is the data and the bullet duplicating the heading is skipped. - """ - if len(table_block) < 3: - return "\n".join(table_block) - headers = split_markdown_table_row(table_block[0]) + """Render a GFM table as bold-heading + bullet groups (same alignment logic as Telegram's + renderer: without a row-label column the full row is data and the heading bullet is skipped).""" + headers = split_markdown_table_row(table_block[0]) if len(table_block) >= 3 else [] if len(headers) < 2: return "\n".join(table_block) - first_data_row = split_markdown_table_row(table_block[2]) - has_row_label_col = len(first_data_row) == len(headers) + 1 + has_row_label_col = len(split_markdown_table_row(table_block[2])) == len(headers) + 1 rendered_groups: list[str] = [] for index, row in enumerate(table_block[2:], start=1): cells = split_markdown_table_row(row) @@ -200,24 +164,15 @@ def _render_table_block(table_block: list[str]) -> str: else: heading = next((cell for cell in cells if cell), f"Row {index}") data_cells = cells - if len(data_cells) < len(headers): - data_cells.extend([""] * (len(headers) - len(data_cells))) - elif len(data_cells) > len(headers): - data_cells = data_cells[: len(headers)] - bullets = [ - f"• {header}: {value}" - for header, value in zip(headers, data_cells) - if has_row_label_col or value != heading - ] + data_cells = (data_cells + [""] * len(headers))[: len(headers)] + bullets = [f"• {header}: {value}" for header, value in zip(headers, data_cells) + if has_row_label_col or value != heading] rendered_groups.append("\n".join([f"**{heading}**", *bullets])) return "\n\n".join(rendered_groups) def convert_table_to_bullets(text: str) -> str: - """Rewrite GFM pipe tables into bold-heading + bullet groups. - - Tables inside fenced code blocks are left alone. - """ + """Rewrite GFM pipe tables into bold-heading + bullet groups; fenced code is left alone.""" if '|' not in text or '-' not in text: return text lines = text.split('\n') @@ -226,73 +181,55 @@ def convert_table_to_bullets(text: str) -> str: i = 0 while i < len(lines): line = lines[i] - stripped = line.lstrip() - if stripped.startswith('```'): - in_fence = not in_fence - if in_fence or stripped.startswith('```'): - out.append(line) - i += 1 - continue - if '|' in line and i + 1 < len(lines) and TABLE_SEPARATOR_RE.match(lines[i + 1]): - table_block = [line, lines[i + 1]] + is_fence_line = line.lstrip().startswith('```') + in_fence ^= is_fence_line + if not (in_fence or is_fence_line) and '|' in line and i + 1 < len(lines) \ + and TABLE_SEPARATOR_RE.match(lines[i + 1]): j = i + 2 while j < len(lines) and is_table_row(lines[j]): - table_block.append(lines[j]) j += 1 - out.append(_render_table_block(table_block)) + out.append(_render_table_block(lines[i:j])) i = j - continue - out.append(line) - i += 1 + else: + out.append(line) + i += 1 return '\n'.join(out) -# ─── Mention-pattern compilation ───────────────────────────────────────────── - - -def compile_mention_patterns( - raw, - *, - log_prefix: str, - platform_label: str | None = None, - display_label: str | None = None, - defaults: 'list[str] | None' = None, - logger_: 'logging.Logger | None' = None, -) -> 'list[re.Pattern]': +def compile_mention_patterns(raw, *, log_prefix: str, platform_label: str | None = None, + display_label: str | None = None, defaults: 'list[str] | None' = None, + logger_: 'logging.Logger | None' = None) -> 'list[re.Pattern]': """Compile regex wake-word/mention patterns from config or env values. - * **Config-style** (dingtalk, telegram): pass ``platform_label``. ``raw`` - must be a list or string (anything else warns and yields ``[]``); - non-string entries are skipped; a summary info log is emitted on load. - * **Wakeword-style** (photon, bluebubbles): pass ``defaults``. ``raw`` may be - None (defaults), a string (JSON list or comma/newline separated), a list, - or a scalar; entries are coerced via ``str()``. - - ``log_prefix`` is interpolated into every log line so per-adapter output - stays byte-identical to the historical inline implementations. + * **Config-style** (dingtalk, telegram): pass ``platform_label``. ``raw`` must be a + list or string (else warn + ``[]``); non-string entries skipped; info log on load. + * **Wakeword-style** (photon, bluebubbles): pass ``defaults``. ``raw`` may be None + (defaults), a string (JSON list or comma/newline separated), a list, or a scalar. + ``log_prefix`` is interpolated into every log line so per-adapter output stays + byte-identical to the historical inline implementations. """ log = logger_ or logger - if platform_label is not None: - display = display_label or platform_label - patterns = raw - if patterns is None: - return [] - if isinstance(patterns, str): - patterns = [patterns] - if not isinstance(patterns, list): - log.warning( - "[%s] %s mention_patterns must be a list or string; got %s", - log_prefix, platform_label, type(patterns).__name__, - ) - return [] + + def _compile(patterns, warn_fmt, *warn_args): compiled: list[re.Pattern] = [] for pattern in patterns: - if not isinstance(pattern, str) or not pattern.strip(): - continue try: compiled.append(re.compile(pattern, re.IGNORECASE)) except re.error as exc: - log.warning("[%s] Invalid %s mention pattern %r: %s", log_prefix, display, pattern, exc) + log.warning(warn_fmt, log_prefix, *warn_args, pattern, exc) + return compiled + + if platform_label is not None: + display = display_label or platform_label + if raw is None: + return [] + patterns = [raw] if isinstance(raw, str) else raw + if not isinstance(patterns, list): + log.warning("[%s] %s mention_patterns must be a list or string; got %s", + log_prefix, platform_label, type(patterns).__name__) + return [] + compiled = _compile([p for p in patterns if isinstance(p, str) and p.strip()], + "[%s] Invalid %s mention pattern %r: %s", display) if compiled: log.info("[%s] Loaded %d %s mention pattern(s)", log_prefix, len(compiled), display) return compiled @@ -305,45 +242,26 @@ def compile_mention_patterns( except Exception: loaded = None patterns = loaded if isinstance(loaded, list) else [ - part.strip() for line in text.splitlines() for part in line.split(",") - ] - elif isinstance(raw, list): - patterns = raw + part.strip() for line in text.splitlines() for part in line.split(",")] else: - patterns = [raw] - compiled = [] - for pattern in patterns: - text = str(pattern).strip() - if not text: - continue - try: - compiled.append(re.compile(text, re.IGNORECASE)) - except re.error as exc: - log.warning("[%s] Invalid mention pattern %r: %s", log_prefix, text, exc) - return compiled + patterns = raw if isinstance(raw, list) else [raw] + texts = [t for t in (str(p).strip() for p in patterns) if t] + return _compile(texts, "[%s] Invalid mention pattern %r: %s") # ─── Fence-Aware Markdown Chunking ─────────────────────────────────────────── -# Shared core for the chunkers in gateway/stream_consumer.py (newline-preferred -# + fence balancing: ``prefer_paragraphs=False, balance_fences=True``), yuanbao -# (atomic blocks + paragraph splitting, fences kept as atoms: -# ``prefer_paragraphs=True, balance_fences=False``) and weixin (own block -# splitter, reuses ``greedy_pack_blocks``). +# Shared core for gateway/stream_consumer.py (``prefer_paragraphs=False, balance_fences=True``), +# yuanbao (``prefer_paragraphs=True, balance_fences=False``) and weixin (``greedy_pack_blocks``). def text_has_unclosed_fence(text: str) -> bool: """Return True when *text* ends inside an unclosed ``` code fence.""" - in_fence = False - for line in text.split('\n'): - if line.startswith('```'): - in_fence = not in_fence - return in_fence + return sum(line.startswith('```') for line in text.split('\n')) % 2 == 1 def text_ends_with_table_row(text: str) -> bool: """True when the last non-empty line starts and ends with ``|``.""" - trimmed = text.rstrip() - return bool(trimmed) and _is_pipe_row(trimmed.split('\n')[-1]) + return _is_pipe_row(text.rstrip().split('\n')[-1]) def is_fence_atom(text: str) -> bool: @@ -364,6 +282,15 @@ def is_table_atom(text: str) -> bool: _SENTENCE_END_NEWLINE_RE = re.compile(r'[。!?.!?]\n') +def _cp_budget(text, budget, len_fn): + """Code-point count of the longest prefix of *text* within *budget* ``len_fn`` units + (callers guarantee ``len_fn(text) > budget``, so plain ``len`` needs no search).""" + if len_fn is len: + return budget + from gateway.platforms.base import _custom_unit_to_cp # heavyweight; lazy + return _custom_unit_to_cp(text, budget, len_fn) + + def split_at_paragraph_boundary(text, max_chars, len_fn=None): """Split at the nearest paragraph boundary within *max_chars*; return (head, tail). @@ -374,34 +301,22 @@ def split_at_paragraph_boundary(text, max_chars, len_fn=None): _len = len_fn or len if _len(text) <= max_chars: return text, '' - if _len is len: - window = text[:max_chars] - else: - from gateway.platforms.base import _custom_unit_to_cp # heavyweight; lazy - window = text[:_custom_unit_to_cp(text, max_chars, _len)] + window = text[:_cp_budget(text, max_chars, _len)] pos = window.rfind('\n\n') - if pos > 0: - return text[:pos + 2], text[pos + 2:] - best_pos = -1 - for m in _SENTENCE_END_NEWLINE_RE.finditer(window): - best_pos = m.end() - if best_pos > 0: - return text[:best_pos], text[best_pos:] - pos = window.rfind('\n') - if pos > 0: - return text[:pos + 1], text[pos + 1:] - cut = len(window) + sentence_ends = [m.end() for m in _SENTENCE_END_NEWLINE_RE.finditer(window)] + cut = pos + 2 if pos > 0 else (sentence_ends[-1] if sentence_ends else 0) + if not cut: + pos = window.rfind('\n') + cut = pos + 1 if pos > 0 else len(window) return text[:cut], text[cut:] def split_markdown_atoms(text: str) -> "list[str]": """Split markdown into indivisible atoms: fenced code blocks, tables (consecutive ``|...|`` lines) and paragraphs. Blank lines belong to no atom.""" - lines = text.split('\n') atoms: "list[str]" = [] current_lines: "list[str]" = [] in_fence = False - _is_table_line = _is_pipe_row def _flush_current() -> None: if current_lines: @@ -409,24 +324,22 @@ def split_markdown_atoms(text: str) -> "list[str]": if atom.strip(): atoms.append(atom) current_lines.clear() - for line in lines: + + for line in text.split('\n'): if in_fence: current_lines.append(line) - if line.startswith('```') and len(current_lines) > 1: + if line.startswith('```'): in_fence = False _flush_current() elif line.startswith('```'): _flush_current() in_fence = True current_lines.append(line) - elif _is_table_line(line): - if current_lines and not _is_table_line(current_lines[-1]): - _flush_current() - current_lines.append(line) elif line.strip() == '': _flush_current() else: - if current_lines and _is_table_line(current_lines[-1]): + # A table line and a non-table line never share an atom. + if current_lines and _is_pipe_row(current_lines[-1]) != _is_pipe_row(line): _flush_current() current_lines.append(line) _flush_current() @@ -447,15 +360,12 @@ def infer_block_separator(prev_chunk: str, next_chunk: str) -> str: def merge_streaming_fences(chunks: "list[str]") -> "list[str]": """Rejoin chunks truncated mid-fence: while chunk *i* has an unclosed fence and a successor exists, merge the successor in via :func:`infer_block_separator`.""" - if not chunks: - return [] result: "list[str]" = [] i = 0 while i < len(chunks): current = chunks[i] while text_has_unclosed_fence(current) and i + 1 < len(chunks): - sep = infer_block_separator(current, chunks[i + 1]) - current = current + sep + chunks[i + 1] + current = current + infer_block_separator(current, chunks[i + 1]) + chunks[i + 1] i += 1 result.append(current) i += 1 @@ -470,15 +380,10 @@ def balance_fences_across_chunks(chunks: "list[str]") -> "list[str]": out: "list[str]" = [] carry_lang = None for chunk in chunks: - prefix = f"```{carry_lang}\n" if carry_lang is not None else "" + body = f"```{carry_lang}\n{chunk}" if carry_lang is not None else chunk in_code, lang = fence_state_after(chunk, carry_lang is not None, carry_lang or "") - body = prefix + chunk - if in_code: - body += "\n```" - carry_lang = lang - else: - carry_lang = None - out.append(body) + carry_lang = lang if in_code else None + out.append(body + "\n```" if in_code else body) return out @@ -487,20 +392,14 @@ def fence_state_after(text: str, in_code: bool = False, lang: str = "") -> "tupl for line in text.split("\n"): stripped = line.strip() if stripped.startswith("```"): - if in_code: - in_code, lang = False, "" - else: - tag = stripped[3:].strip() - in_code, lang = True, (tag.split()[0] if tag else "") + tag = stripped[3:].split() + in_code, lang = (False, "") if in_code else (True, tag[0] if tag else "") return in_code, lang def greedy_pack_blocks(blocks, max_length, len_fn=None, sep="\n\n", overflow=None): - """Greedily pack *blocks* (joined with *sep*) into chunks of at most *max_length*. - - A block that alone exceeds the limit goes through *overflow(block)* (must - return a list of chunks) when provided, else is emitted as-is. - """ + """Greedily pack *blocks* (joined with *sep*) into chunks of at most *max_length*; an + oversized block goes through *overflow(block)* (-> list of chunks) if given, else as-is.""" _len = len_fn or len packed: "list[str]" = [] current = "" @@ -514,8 +413,7 @@ def greedy_pack_blocks(blocks, max_length, len_fn=None, sep="\n\n", overflow=Non current = "" if _len(block) <= max_length: current = block - continue - if overflow is not None: + elif overflow is not None: packed.extend(overflow(block)) else: packed.append(block) @@ -524,35 +422,24 @@ def greedy_pack_blocks(blocks, max_length, len_fn=None, sep="\n\n", overflow=Non return packed -def split_text_fence_aware( - text, - limit, - len_fn=None, - *, - prefer_paragraphs=True, - balance_fences=False, -): +def split_text_fence_aware(text, limit, len_fn=None, *, prefer_paragraphs=True, + balance_fences=False): """Split markdown into chunks of at most *limit*, respecting fences. - ``prefer_paragraphs=True`` (yuanbao-derived): atoms (fences, tables, - paragraphs) are greedily merged, oversized non-atomic chunks split at - paragraph boundaries, small neighbours re-merged; a single atom larger than - *limit* is emitted oversize rather than broken. - ``prefer_paragraphs=False`` (stream_consumer-derived): newline-preferred - hard splitting with headroom reserved for fence markers. - ``balance_fences=True`` closes/reopens fences at chunk boundaries so each - chunk renders standalone. + ``prefer_paragraphs=True`` (yuanbao-derived): atoms (fences, tables, paragraphs) are + greedily merged, oversized non-atomic chunks split at paragraph boundaries, small + neighbours re-merged; a single atom larger than *limit* is emitted oversize, not broken. + ``prefer_paragraphs=False`` (stream_consumer-derived): newline-preferred hard splitting + with headroom reserved for fence markers. ``balance_fences=True`` closes/reopens fences + at chunk boundaries so each chunk renders standalone. """ - _len = len_fn or len if not text: return [] if prefer_paragraphs: chunks = _chunk_markdown_paragraphs(text, limit, len_fn) else: - chunks = _chunk_newline_preferred(text, limit, _len) - if balance_fences: - chunks = balance_fences_across_chunks(chunks) - return chunks + chunks = _chunk_newline_preferred(text, limit, len_fn or len) + return balance_fences_across_chunks(chunks) if balance_fences else chunks def _chunk_markdown_paragraphs(text, max_chars, len_fn=None): @@ -560,34 +447,26 @@ def _chunk_markdown_paragraphs(text, max_chars, len_fn=None): _len = len_fn or len if _len(text) <= max_chars: return [text] - atoms = split_markdown_atoms(text) # Phase 2: greedy merge; oversized fence/table atoms stay indivisible. chunks: "list[str]" = [] indivisible_set: "set[int]" = set() current_parts: "list[str]" = [] current_len = 0 - - def _flush_parts() -> None: - if current_parts: - chunks.append('\n\n'.join(current_parts)) - for atom in atoms: + for atom in split_markdown_atoms(text): atom_len = _len(atom) sep_len = 2 if current_parts else 0 - projected_len = current_len + sep_len + atom_len - if projected_len > max_chars and current_parts: - _flush_parts() - current_parts = [] - current_len = 0 - sep_len = 0 - if (not current_parts - and atom_len > max_chars + if current_len + sep_len + atom_len > max_chars and current_parts: + chunks.append('\n\n'.join(current_parts)) + current_parts, current_len, sep_len = [], 0, 0 + if (not current_parts and atom_len > max_chars and (is_fence_atom(atom) or is_table_atom(atom))): indivisible_set.add(len(chunks)) chunks.append(atom) continue current_parts.append(atom) current_len += sep_len + atom_len - _flush_parts() + if current_parts: + chunks.append('\n\n'.join(current_parts)) # Phase 3: split still-oversized divisible chunks at paragraph boundaries. result: "list[str]" = [] for idx, chunk in enumerate(chunks): @@ -604,17 +483,14 @@ def _chunk_markdown_paragraphs(text, max_chars, len_fn=None): if remaining: result.append(remaining) # Phase 4: merge small chunks with neighbours. - if len(result) > 1: - merged: "list[str]" = [result[0]] - for chunk in result[1:]: - prev = merged[-1] - combined = prev + '\n\n' + chunk - if _len(combined) <= max_chars: - merged[-1] = combined - else: - merged.append(chunk) - result = merged - return [c for c in result if c] + merged: "list[str]" = result[:1] + for chunk in result[1:]: + combined = merged[-1] + '\n\n' + chunk + if _len(combined) <= max_chars: + merged[-1] = combined + else: + merged.append(chunk) + return [c for c in merged if c] def _chunk_newline_preferred(text, limit, len_fn): @@ -622,17 +498,14 @@ def _chunk_newline_preferred(text, limit, len_fn): if len_fn(text) <= limit: return [text] # Reserve headroom for fence markers a balancing pass may add. - split_limit = limit - if "```" in text: - split_limit = max(limit - 16, limit // 2, 1) - from gateway.platforms.base import _custom_unit_to_cp # heavyweight; lazy + split_limit = max(limit - 16, limit // 2, 1) if "```" in text else limit chunks: "list[str]" = [] remaining = text while len_fn(remaining) > split_limit: - _cp_budget = _custom_unit_to_cp(remaining, split_limit, len_fn) - split_at = remaining.rfind("\n", 0, _cp_budget) - if split_at < _cp_budget // 2: - split_at = _cp_budget + budget = _cp_budget(remaining, split_limit, len_fn) + split_at = remaining.rfind("\n", 0, budget) + if split_at < budget // 2: + split_at = budget chunks.append(remaining[:split_at]) remaining = remaining[split_at:].lstrip("\n") if remaining: diff --git a/gateway/platforms/media_cache.py b/gateway/platforms/media_cache.py index 4d565a61c0..8ec4d76c80 100644 --- a/gateway/platforms/media_cache.py +++ b/gateway/platforms/media_cache.py @@ -1,37 +1,11 @@ """Shared mime↔extension dispatch for inbound (downloaded) platform media. -Historically every gateway adapter hand-rolled its own mime→extension map -before handing downloaded bytes to the cache primitives in -``gateway.platforms.base`` (``cache_image_from_bytes``, -``cache_audio_from_bytes``, ``cache_document_from_bytes``). Those maps -*disagree* with each other on purpose — e.g. BlueBubbles coerces -``image/heic`` to ``.jpg`` because downstream vision tools can't read HEIC, -while WhatsApp Cloud pins ``audio/ogg`` to ``.ogg`` (not the RFC-correct -``.oga`` Python's ``mimetypes`` returns) because the STT pipeline whitelists -extensions. - -This module owns: - -* ``DEFAULT_MIME_TO_EXT`` — the union table of entries the adapters already - agree on (plus a few uncontroversial document types). -* ``DEFAULT_EXT_TO_MIME`` — the canonical inverse (used by Signal to map a - sniffed extension back to a content type). -* ``ext_for_mime`` / ``mime_for_ext`` — lookup helpers that accept - per-adapter ``overrides`` so each adapter's historical (divergent) - behavior is preserved byte-for-byte. -* ``cache_media_bytes`` — one-call dispatch: classify the mime, resolve the - extension, and write to the right cache (image / audio / document). - -Behavior-preservation contract: adapters that had divergent maps pass them -as ``overrides`` (and, where their historical code never consulted -``mimetypes`` or a shared table, disable those fallbacks via -``use_defaults`` / ``use_mimetypes``). The parity tests in -``tests/gateway/test_media_cache.py`` hardcode the historical outputs as -the contract. - -NOTE: ``gateway/platforms/weixin.py`` also has a private mime map -(``_mime_from_filename``) but is intentionally NOT migrated here — another -in-flight branch edits that file. Follow-up: fold it in once that lands. +Adapters historically hand-rolled divergent mime→extension maps on purpose (BlueBubbles +coerces ``image/heic`` to ``.jpg`` for vision tools; WhatsApp Cloud pins ``audio/ogg`` to +``.ogg`` for the STT extension whitelist). This module owns the agreed-upon union tables +plus lookup helpers taking per-adapter ``overrides`` (and ``use_defaults``/``use_mimetypes`` +toggles) so each adapter's historical output stays byte-identical — +``tests/gateway/test_media_cache.py`` pins those outputs as the contract. """ from __future__ import annotations @@ -40,47 +14,24 @@ import mimetypes import uuid from typing import Mapping, Optional -# --------------------------------------------------------------------------- -# Shared tables -# --------------------------------------------------------------------------- - -# Union of the per-adapter maps where the adapters already agree (or where -# only one adapter pinned the type and no other adapter contradicts it). -# Entries deliberately favor the common-in-the-wild extension over the -# RFC-correct one (``audio/ogg`` → ``.ogg``, not ``.oga``) because the -# downstream STT/vision pipelines whitelist real-world extensions. +# Union of the per-adapter maps where they agree. Favors the common-in-the-wild extension +# over the RFC-correct one (``audio/ogg`` → ``.ogg``, not ``.oga``): downstream STT/vision +# pipelines whitelist real-world extensions. (weixin.py's private ``_mime_from_filename`` +# is not folded in yet — another in-flight branch edits that file.) DEFAULT_MIME_TO_EXT: dict[str, str] = { - # --- images (bluebubbles + whatsapp_cloud agree; matches mimetypes) --- - "image/jpeg": ".jpg", - "image/png": ".png", - "image/gif": ".gif", - "image/webp": ".webp", - # --- audio --- - "audio/ogg": ".ogg", # bluebubbles + whatsapp_cloud agree - "audio/x-opus+ogg": ".ogg", # whatsapp voice notes (opus-in-ogg) - "audio/opus": ".ogg", # whatsapp voice notes (opus-in-ogg) - "audio/mpeg": ".mp3", - "audio/mp3": ".mp3", # non-standard but seen in the wild - "audio/wav": ".wav", - "audio/mp4": ".m4a", # bluebubbles + whatsapp_cloud agree - "audio/x-m4a": ".m4a", - "audio/aac": ".aac", - # --- video / documents (from signal's inverse table) --- - "video/mp4": ".mp4", - "application/pdf": ".pdf", - "application/zip": ".zip", + "image/jpeg": ".jpg", "image/png": ".png", "image/gif": ".gif", "image/webp": ".webp", + "audio/ogg": ".ogg", "audio/x-opus+ogg": ".ogg", "audio/opus": ".ogg", # whatsapp voice notes + "audio/mpeg": ".mp3", "audio/mp3": ".mp3", "audio/wav": ".wav", + "audio/mp4": ".m4a", "audio/x-m4a": ".m4a", "audio/aac": ".aac", + "video/mp4": ".mp4", "application/pdf": ".pdf", "application/zip": ".zip", } -# Canonical inverse. Kept explicit (rather than mechanically inverted) -# because the forward table is many-to-one — e.g. both ``audio/mpeg`` and -# ``audio/mp3`` map to ``.mp3`` and the inverse must pick the canonical -# mime. This is byte-identical to Signal's historical ``_EXT_TO_MIME``. +# Explicit inverse (the forward table is many-to-one, so the inverse must pick +# the canonical mime). Byte-identical to Signal's historical ``_EXT_TO_MIME``. DEFAULT_EXT_TO_MIME: dict[str, str] = { - ".jpg": "image/jpeg", ".jpeg": "image/jpeg", ".png": "image/png", - ".gif": "image/gif", ".webp": "image/webp", - ".ogg": "audio/ogg", ".mp3": "audio/mpeg", ".wav": "audio/wav", - ".m4a": "audio/mp4", ".aac": "audio/aac", - ".mp4": "video/mp4", ".pdf": "application/pdf", + ".jpg": "image/jpeg", ".jpeg": "image/jpeg", ".png": "image/png", ".gif": "image/gif", + ".webp": "image/webp", ".ogg": "audio/ogg", ".mp3": "audio/mpeg", ".wav": "audio/wav", + ".m4a": "audio/mp4", ".aac": "audio/aac", ".mp4": "video/mp4", ".pdf": "application/pdf", ".zip": "application/zip", } @@ -90,102 +41,48 @@ def _normalize_mime(mime: str) -> str: return (mime or "").split(";")[0].strip().lower() -# --------------------------------------------------------------------------- -# Lookups -# --------------------------------------------------------------------------- - -def ext_for_mime( - mime: str, - *, - overrides: Optional[Mapping[str, str]] = None, - use_defaults: bool = True, - use_mimetypes: bool = True, - fallback: Optional[str] = None, -) -> Optional[str]: - """Resolve a mime type to a file extension (including the dot). - - Resolution order: ``overrides`` → ``DEFAULT_MIME_TO_EXT`` (if - ``use_defaults``) → ``mimetypes.guess_extension`` (if - ``use_mimetypes``) → ``fallback``. - - Adapters with historical divergent maps pass them via ``overrides`` - and disable the stages their old code never consulted, keeping their - outputs byte-identical to the pre-refactor behavior. - """ +def ext_for_mime(mime: str, *, overrides: Optional[Mapping[str, str]] = None, + use_defaults: bool = True, use_mimetypes: bool = True, + fallback: Optional[str] = None) -> Optional[str]: + """Resolve a mime type to a dotted extension: ``overrides`` → ``DEFAULT_MIME_TO_EXT`` (if + ``use_defaults``) → ``mimetypes.guess_extension`` (if ``use_mimetypes``) → ``fallback``.""" primary = _normalize_mime(mime) if not primary: return fallback - if overrides: - ext = overrides.get(primary) - if ext: - return ext - if use_defaults: - ext = DEFAULT_MIME_TO_EXT.get(primary) - if ext: - return ext - if use_mimetypes: - ext = mimetypes.guess_extension(primary) + stages = [overrides.get if overrides else None, + DEFAULT_MIME_TO_EXT.get if use_defaults else None, + mimetypes.guess_extension if use_mimetypes else None] + for lookup in stages: + ext = lookup(primary) if lookup else None if ext: return ext return fallback -def mime_for_ext( - ext: str, - *, - overrides: Optional[Mapping[str, str]] = None, - fallback: str = "application/octet-stream", -) -> str: - """Inverse lookup: file extension → canonical mime type. - - Resolution order: ``overrides`` → ``DEFAULT_EXT_TO_MIME`` → ``fallback``. - """ +def mime_for_ext(ext: str, *, overrides: Optional[Mapping[str, str]] = None, + fallback: str = "application/octet-stream") -> str: + """Inverse lookup: ``overrides`` → ``DEFAULT_EXT_TO_MIME`` → ``fallback``.""" key = (ext or "").strip().lower() - if overrides: - mime = overrides.get(key) - if mime: - return mime - return DEFAULT_EXT_TO_MIME.get(key, fallback) + return (overrides or {}).get(key) or DEFAULT_EXT_TO_MIME.get(key, fallback) -# --------------------------------------------------------------------------- -# One-call cache dispatch -# --------------------------------------------------------------------------- - -def cache_media_bytes( - data: bytes, - mime: str, - *, - filename_hint: str = "", - kind_hint: Optional[str] = None, - ext_overrides: Optional[Mapping[str, str]] = None, -) -> str: +def cache_media_bytes(data: bytes, mime: str, *, filename_hint: str = "", + kind_hint: Optional[str] = None, + ext_overrides: Optional[Mapping[str, str]] = None) -> str: """Cache downloaded media bytes and return the local file path. - Picks the image / audio / document cache primitive from - ``gateway.platforms.base`` based on the mime class (or an explicit - ``kind_hint`` of ``"image"``, ``"audio"`` or ``"document"``). - ``filename_hint`` is used for document caching (falls back to a - generated name with the resolved extension). ``ext_overrides`` is - threaded through to :func:`ext_for_mime` for adapters that need their - historical mappings. + Picks the image / audio / document cache primitive by mime class (or explicit ``kind_hint`` + ``"image"``/``"audio"``/``"document"``). ``filename_hint`` names document files (else a + generated name with the resolved extension); ``ext_overrides`` feeds :func:`ext_for_mime`. """ - # Local import: base is a large module and some adapters import this - # module very early; keep import-time coupling minimal. + # Local import: base is heavyweight and some adapters import this module very early. from gateway.platforms.base import ( - cache_audio_from_bytes, - cache_document_from_bytes, - cache_image_from_bytes, - ) + cache_audio_from_bytes, cache_document_from_bytes, cache_image_from_bytes) primary = _normalize_mime(mime) kind = kind_hint if kind is None: - if primary.startswith("image/"): - kind = "image" - elif primary.startswith("audio/"): - kind = "audio" - else: - kind = "document" + kind = ("image" if primary.startswith("image/") + else "audio" if primary.startswith("audio/") else "document") if kind == "image": ext = ext_for_mime(primary, overrides=ext_overrides, fallback=".jpg") or ".jpg" return cache_image_from_bytes(data, ext)