refactor(gateway/platforms): simplify helpers.py (strip rules table, unified mention compiler, chunker cleanups), media_cache.py, _http_client_limits.py, __init__.py
This commit is contained in:
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
+139
-266
@@ -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_])(.+?)(?<![\s_])__\b", re.DOTALL)
|
||||
_RE_ITALIC_UNDER = re.compile(r"\b_(?![\s_])(.+?)(?<![\s_])_\b", re.DOTALL)
|
||||
_RE_CODE_BLOCK = re.compile(r"```[a-zA-Z0-9_+-]*\n?")
|
||||
_RE_INLINE_CODE = re.compile(r"`(.+?)`")
|
||||
_RE_HEADING = re.compile(r"^#{1,6}\s+", re.MULTILINE)
|
||||
_RE_LINK = re.compile(r"\[([^\]]+)\]\([^\)]+\)")
|
||||
_RE_MULTI_NEWLINE = re.compile(r"\n{3,}")
|
||||
# Markdown-stripping rules, applied in order: bold, italic, bold/italic underscore,
|
||||
# code fence markers, inline code, headings, links, then newline squeeze.
|
||||
_STRIP_RULES = (
|
||||
(re.compile(r"\*\*(.+?)\*\*", re.DOTALL), r"\1"),
|
||||
(re.compile(r"\*(.+?)\*", re.DOTALL), r"\1"),
|
||||
(re.compile(r"\b__(?![\s_])(.+?)(?<![\s_])__\b", re.DOTALL), r"\1"),
|
||||
(re.compile(r"\b_(?![\s_])(.+?)(?<![\s_])_\b", re.DOTALL), r"\1"),
|
||||
(re.compile(r"```[a-zA-Z0-9_+-]*\n?"), ""),
|
||||
(re.compile(r"`(.+?)`"), r"\1"),
|
||||
(re.compile(r"^#{1,6}\s+", re.MULTILINE), ""),
|
||||
(re.compile(r"\[([^\]]+)\]\([^\)]+\)"), r"\1"),
|
||||
(re.compile(r"\n{3,}"), "\n\n"),
|
||||
)
|
||||
|
||||
|
||||
def strip_markdown(text: str) -> 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 (``<platform>_threads.json``).
|
||||
|
||||
``thread_id in tracker`` checks membership; ``tracker.mark(thread_id)`` persists.
|
||||
"""
|
||||
"""Persistent set of threads the bot has participated in (``<platform>_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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user