mirror of
https://github.com/MODSetter/SurfSense.git
synced 2026-07-20 23:21:06 +02:00
chore: linting
This commit is contained in:
parent
0a012dbc79
commit
ce952d2ad1
127 changed files with 821 additions and 517 deletions
|
|
@ -31,7 +31,9 @@ def slack_account_credentials(account: ExternalChatAccount) -> dict:
|
|||
"""Decrypt Slack gateway credentials stored as encrypted JSON."""
|
||||
if not account.encrypted_credentials:
|
||||
return {}
|
||||
raw = TokenEncryption(config.SECRET_KEY or "").decrypt_token(account.encrypted_credentials)
|
||||
raw = TokenEncryption(config.SECRET_KEY or "").decrypt_token(
|
||||
account.encrypted_credentials
|
||||
)
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
|
|
@ -44,7 +46,9 @@ def discord_account_credentials(account: ExternalChatAccount) -> dict:
|
|||
"""Decrypt Discord gateway credentials stored as encrypted JSON."""
|
||||
if not account.encrypted_credentials:
|
||||
return {}
|
||||
raw = TokenEncryption(config.SECRET_KEY or "").decrypt_token(account.encrypted_credentials)
|
||||
raw = TokenEncryption(config.SECRET_KEY or "").decrypt_token(
|
||||
account.encrypted_credentials
|
||||
)
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
|
|
@ -135,4 +139,3 @@ async def get_discord_account_by_guild(
|
|||
)
|
||||
)
|
||||
return result.scalars().first()
|
||||
|
||||
|
|
|
|||
|
|
@ -21,7 +21,9 @@ from app.tasks.chat.streaming.flows import stream_new_chat
|
|||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def _events_from_sse(chunks: AsyncIterator[str]) -> AsyncIterator[GatewayStreamEvent]:
|
||||
async def _events_from_sse(
|
||||
chunks: AsyncIterator[str],
|
||||
) -> AsyncIterator[GatewayStreamEvent]:
|
||||
saw_text = False
|
||||
async for chunk in chunks:
|
||||
for raw_line in chunk.splitlines():
|
||||
|
|
@ -98,4 +100,3 @@ async def call_agent_for_gateway(
|
|||
record_gateway_turn_latency(0, platform=platform_label)
|
||||
finally:
|
||||
release_thread_lock(thread.id)
|
||||
|
||||
|
|
|
|||
|
|
@ -52,4 +52,3 @@ async def assert_authorization_invariant(
|
|||
await _fail(session, binding, f"rbac_{exc.status_code}")
|
||||
|
||||
return user
|
||||
|
||||
|
|
|
|||
|
|
@ -1,2 +1 @@
|
|||
"""Base gateway interfaces."""
|
||||
|
||||
|
|
|
|||
|
|
@ -62,9 +62,10 @@ class BasePlatformAdapter(ABC):
|
|||
async def validate_credentials(self) -> dict[str, Any]:
|
||||
"""Validate configured credentials and return account metadata."""
|
||||
|
||||
async def fetch_updates(self, *, offset: int | None) -> AsyncIterator[dict[str, Any]]:
|
||||
async def fetch_updates(
|
||||
self, *, offset: int | None
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
"""Yield provider updates for long-polling adapters."""
|
||||
if False:
|
||||
yield {} # pragma: no cover
|
||||
raise NotImplementedError("This adapter does not support long-polling")
|
||||
|
||||
|
|
|
|||
|
|
@ -16,4 +16,3 @@ def hash_external_id(value: str | int | None) -> str | None:
|
|||
if not normalized:
|
||||
return None
|
||||
return hashlib.sha256(normalized.encode("utf-8")).hexdigest()
|
||||
|
||||
|
|
|
|||
|
|
@ -25,4 +25,3 @@ class BaseStreamTranslator(ABC):
|
|||
@abstractmethod
|
||||
async def translate(self, events: AsyncIterator[GatewayStreamEvent]) -> None:
|
||||
"""Consume agent stream events and emit platform messages."""
|
||||
|
||||
|
|
|
|||
|
|
@ -64,4 +64,3 @@ def resume_binding(binding: ExternalChatBinding) -> None:
|
|||
binding.state = ExternalChatBindingState.BOUND
|
||||
binding.suspended_at = None
|
||||
binding.suspended_reason = None
|
||||
|
||||
|
|
|
|||
|
|
@ -58,8 +58,10 @@ async def _whatsapp_baileys_supervisor() -> None:
|
|||
async with async_session_maker() as session:
|
||||
result = await session.execute(
|
||||
select(ExternalChatAccount).where(
|
||||
ExternalChatAccount.platform == ExternalChatPlatform.WHATSAPP,
|
||||
ExternalChatAccount.mode == ExternalChatAccountMode.SELF_HOST_BYO,
|
||||
ExternalChatAccount.platform
|
||||
== ExternalChatPlatform.WHATSAPP,
|
||||
ExternalChatAccount.mode
|
||||
== ExternalChatAccountMode.SELF_HOST_BYO,
|
||||
ExternalChatAccount.is_system_account.is_(False),
|
||||
ExternalChatAccount.suspended_at.is_(None),
|
||||
)
|
||||
|
|
@ -128,7 +130,9 @@ async def start_byo_long_poll_supervisors() -> None:
|
|||
)
|
||||
_tasks.add(task)
|
||||
task.add_done_callback(_tasks.discard)
|
||||
logger.info("Started BYO Telegram long-poll supervisor account_id=%s", account.id)
|
||||
logger.info(
|
||||
"Started BYO Telegram long-poll supervisor account_id=%s", account.id
|
||||
)
|
||||
|
||||
if config.GATEWAY_WHATSAPP_INTAKE_MODE == "baileys":
|
||||
task = asyncio.create_task(
|
||||
|
|
@ -151,9 +155,12 @@ async def stop_byo_long_poll_supervisors() -> None:
|
|||
task.cancel()
|
||||
if tasks:
|
||||
try:
|
||||
await asyncio.wait_for(asyncio.gather(*tasks, return_exceptions=True), timeout=10)
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(*tasks, return_exceptions=True), timeout=10
|
||||
)
|
||||
except TimeoutError:
|
||||
logger.warning("Timed out waiting for BYO Telegram long-poll supervisors to stop")
|
||||
logger.warning(
|
||||
"Timed out waiting for BYO Telegram long-poll supervisors to stop"
|
||||
)
|
||||
_tasks.clear()
|
||||
_shutdown_event = None
|
||||
|
||||
|
|
|
|||
|
|
@ -39,7 +39,9 @@ def _message_reference_payload(message: discord.Message) -> dict[str, Any] | Non
|
|||
}
|
||||
|
||||
|
||||
def _serialize_message(message: discord.Message, *, bot_user_id: str | None) -> dict[str, Any]:
|
||||
def _serialize_message(
|
||||
message: discord.Message, *, bot_user_id: str | None
|
||||
) -> dict[str, Any]:
|
||||
guild = message.guild
|
||||
channel = message.channel
|
||||
thread_id = str(channel.id) if isinstance(channel, discord.Thread) else None
|
||||
|
|
@ -62,8 +64,7 @@ def _serialize_message(message: discord.Message, *, bot_user_id: str | None) ->
|
|||
"bot": message.author.bot,
|
||||
},
|
||||
"mentions": [
|
||||
{"id": str(user.id), "username": user.name}
|
||||
for user in message.mentions
|
||||
{"id": str(user.id), "username": user.name} for user in message.mentions
|
||||
],
|
||||
"message_reference": _message_reference_payload(message),
|
||||
"created_at": message.created_at.isoformat()
|
||||
|
|
@ -73,7 +74,9 @@ def _serialize_message(message: discord.Message, *, bot_user_id: str | None) ->
|
|||
}
|
||||
|
||||
|
||||
async def _persist_message(message: discord.Message, *, bot_user_id: str | None) -> None:
|
||||
async def _persist_message(
|
||||
message: discord.Message, *, bot_user_id: str | None
|
||||
) -> None:
|
||||
if message.guild is None:
|
||||
return
|
||||
guild_id = str(message.guild.id)
|
||||
|
|
@ -82,7 +85,9 @@ async def _persist_message(message: discord.Message, *, bot_user_id: str | None)
|
|||
async with async_session_maker() as session:
|
||||
account = await get_discord_account_by_guild(session, guild_id=guild_id)
|
||||
if account is None:
|
||||
logger.info("Ignoring Discord message for uninstalled guild_id=%s", guild_id)
|
||||
logger.info(
|
||||
"Ignoring Discord message for uninstalled guild_id=%s", guild_id
|
||||
)
|
||||
return
|
||||
|
||||
inbox_id = await persist_inbound_event(
|
||||
|
|
@ -144,7 +149,9 @@ def _build_client() -> discord.Client:
|
|||
try:
|
||||
await _persist_message(message, bot_user_id=bot_user_id)
|
||||
except Exception:
|
||||
logger.exception("Discord gateway failed to persist message_id=%s", message.id)
|
||||
logger.exception(
|
||||
"Discord gateway failed to persist message_id=%s", message.id
|
||||
)
|
||||
|
||||
return client
|
||||
|
||||
|
|
|
|||
|
|
@ -41,7 +41,9 @@ class DiscordStreamTranslator(BaseStreamTranslator):
|
|||
async def translate(self, events: AsyncIterator[GatewayStreamEvent]) -> None:
|
||||
async for event in events:
|
||||
if event.type in {"text-delta", "text_delta", "text"}:
|
||||
self._buffer += str(event.data.get("text") or event.data.get("delta") or "")
|
||||
self._buffer += str(
|
||||
event.data.get("text") or event.data.get("delta") or ""
|
||||
)
|
||||
elif event.type in {"data-interrupt-request", "interrupt"}:
|
||||
await self._handle_hitl_interrupt()
|
||||
return
|
||||
|
|
@ -53,7 +55,9 @@ class DiscordStreamTranslator(BaseStreamTranslator):
|
|||
async def _flush_final(self) -> None:
|
||||
if not self._buffer:
|
||||
return
|
||||
for chunk in split_text_message(self._buffer, max_chars=DISCORD_MAX_MESSAGE_CHARS):
|
||||
for chunk in split_text_message(
|
||||
self._buffer, max_chars=DISCORD_MAX_MESSAGE_CHARS
|
||||
):
|
||||
await self._send_text(chunk)
|
||||
|
||||
async def _send_text(self, text: str) -> PlatformSendResult:
|
||||
|
|
|
|||
|
|
@ -32,4 +32,3 @@ def filter_hitl_tools(
|
|||
return None
|
||||
blocked = blocked_names or DEFAULT_HITL_TOOL_NAMES
|
||||
return [tool for tool in toolkit if (_tool_name(tool) or "") not in blocked]
|
||||
|
||||
|
|
|
|||
|
|
@ -51,4 +51,3 @@ async def persist_inbound_event(
|
|||
)
|
||||
result = await session.execute(stmt)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
|
|
|||
|
|
@ -128,7 +128,9 @@ async def process_inbound_event(
|
|||
event.status = ExternalChatEventStatus.PROCESSED
|
||||
event.processed_at = datetime.now(UTC)
|
||||
await session.commit()
|
||||
record_gateway_inbox_processed(platform=event.platform.value, status="processed")
|
||||
record_gateway_inbox_processed(
|
||||
platform=event.platform.value, status="processed"
|
||||
)
|
||||
|
||||
|
||||
async def _mark_failed(
|
||||
|
|
@ -173,7 +175,9 @@ async def _resolve_slack_thread_binding(
|
|||
parsed,
|
||||
) -> ExternalChatBinding | None:
|
||||
user_peer_id = parsed.metadata.get("slack_user_peer_id")
|
||||
thread_peer_id = parsed.metadata.get("slack_thread_peer_id") or parsed.external_peer_id
|
||||
thread_peer_id = (
|
||||
parsed.metadata.get("slack_thread_peer_id") or parsed.external_peer_id
|
||||
)
|
||||
if not user_peer_id or not thread_peer_id:
|
||||
return None
|
||||
|
||||
|
|
@ -233,7 +237,9 @@ async def _resolve_discord_thread_binding(
|
|||
parsed,
|
||||
) -> ExternalChatBinding | None:
|
||||
user_peer_id = parsed.metadata.get("discord_user_peer_id")
|
||||
thread_peer_id = parsed.metadata.get("discord_thread_peer_id") or parsed.external_peer_id
|
||||
thread_peer_id = (
|
||||
parsed.metadata.get("discord_thread_peer_id") or parsed.external_peer_id
|
||||
)
|
||||
if not user_peer_id or not thread_peer_id:
|
||||
return None
|
||||
|
||||
|
|
@ -357,7 +363,11 @@ async def _dispatch_inbound_event(
|
|||
return
|
||||
|
||||
if binding is None:
|
||||
if bundle.auto_bind_owner and account.owner_user_id and account.owner_search_space_id:
|
||||
if (
|
||||
bundle.auto_bind_owner
|
||||
and account.owner_user_id
|
||||
and account.owner_search_space_id
|
||||
):
|
||||
binding = ExternalChatBinding(
|
||||
account_id=account.id,
|
||||
user_id=account.owner_user_id,
|
||||
|
|
@ -385,7 +395,9 @@ async def _dispatch_inbound_event(
|
|||
event.external_chat_binding_id = binding.id
|
||||
|
||||
if cmd == "/help":
|
||||
handled = await bundle.commands.handle_help_command(adapter=adapter, event=parsed)
|
||||
handled = await bundle.commands.handle_help_command(
|
||||
adapter=adapter, event=parsed
|
||||
)
|
||||
if handled:
|
||||
event.status = ExternalChatEventStatus.PROCESSED
|
||||
await session.commit()
|
||||
|
|
|
|||
|
|
@ -55,4 +55,3 @@ async def stop_gateway_inbox_worker() -> None:
|
|||
with suppress(TimeoutError, asyncio.CancelledError):
|
||||
await asyncio.wait_for(_task, timeout=10)
|
||||
_task = None
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ from datetime import UTC, datetime, timedelta
|
|||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.db import ExternalChatBindingState, ExternalChatBinding
|
||||
from app.db import ExternalChatBinding, ExternalChatBindingState
|
||||
|
||||
PAIRING_CODE_TTL = timedelta(minutes=10)
|
||||
|
||||
|
|
@ -51,4 +51,3 @@ async def redeem_pairing_code(
|
|||
binding.external_username = external_username
|
||||
binding.external_metadata = external_metadata or {}
|
||||
return binding
|
||||
|
||||
|
|
|
|||
|
|
@ -133,4 +133,3 @@ async def wait_for_token(
|
|||
if wait_ms > 0:
|
||||
await asyncio.sleep(wait_ms / 1000)
|
||||
return wait_ms
|
||||
|
||||
|
|
|
|||
|
|
@ -186,4 +186,6 @@ def resolve_platform_bundle(account: ExternalChatAccount) -> PlatformBundle:
|
|||
auto_bind_owner=False,
|
||||
)
|
||||
|
||||
raise RuntimeError(f"unsupported_gateway_platform:{account.platform.value}:{account.mode.value}")
|
||||
raise RuntimeError(
|
||||
f"unsupported_gateway_platform:{account.platform.value}:{account.mode.value}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,7 +8,12 @@ import uuid
|
|||
|
||||
from sqlalchemy import text
|
||||
|
||||
from app.db import ExternalChatPlatform, ExternalChatAccount, async_session_maker, engine
|
||||
from app.db import (
|
||||
ExternalChatAccount,
|
||||
ExternalChatPlatform,
|
||||
async_session_maker,
|
||||
engine,
|
||||
)
|
||||
from app.gateway.inbox import persist_inbound_event, telegram_event_dedupe_key
|
||||
from app.gateway.telegram.adapter import TelegramAdapter
|
||||
from app.observability.metrics import record_gateway_byo_longpoll_running_delta
|
||||
|
|
@ -39,7 +44,9 @@ async def _run_telegram_account(account_id: int, token: str) -> None:
|
|||
account = await session.get(ExternalChatAccount, account_id)
|
||||
offset = None
|
||||
if account is not None:
|
||||
offset = int((account.cursor_state or {}).get("last_update_id", 0)) + 1
|
||||
offset = (
|
||||
int((account.cursor_state or {}).get("last_update_id", 0)) + 1
|
||||
)
|
||||
|
||||
async for update in adapter.fetch_updates(offset=offset):
|
||||
request_id = f"gateway_{uuid.uuid4().hex[:16]}"
|
||||
|
|
@ -58,8 +65,11 @@ async def _run_telegram_account(account_id: int, token: str) -> None:
|
|||
)
|
||||
await session.commit()
|
||||
if inbox_id is not None:
|
||||
logger.debug("Persisted Telegram polling update inbox_id=%s", inbox_id)
|
||||
logger.debug(
|
||||
"Persisted Telegram polling update inbox_id=%s", inbox_id
|
||||
)
|
||||
finally:
|
||||
record_gateway_byo_longpoll_running_delta(-1, account_id=account_id)
|
||||
await conn.execute(text("SELECT pg_advisory_unlock(:key)"), {"key": lock_key})
|
||||
|
||||
await conn.execute(
|
||||
text("SELECT pg_advisory_unlock(:key)"), {"key": lock_key}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -38,7 +38,9 @@ class SlackAdapter(BasePlatformAdapter):
|
|||
slack_user_id = str(event.get("user") or "")
|
||||
message_ts = str(event.get("ts") or "")
|
||||
thread_ts = str(event.get("thread_ts") or message_ts)
|
||||
bot_user_id = self.bot_user_id or str(raw_payload.get("authorizations", [{}])[0].get("user_id") or "")
|
||||
bot_user_id = self.bot_user_id or str(
|
||||
raw_payload.get("authorizations", [{}])[0].get("user_id") or ""
|
||||
)
|
||||
|
||||
if not channel_id or not slack_user_id or not message_ts:
|
||||
return ParsedInboundEvent(
|
||||
|
|
|
|||
|
|
@ -15,7 +15,9 @@ class SlackGatewayClient:
|
|||
def __init__(self, bot_token: str) -> None:
|
||||
self.bot_token = bot_token
|
||||
|
||||
async def api_call(self, method: str, payload: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
async def api_call(
|
||||
self, method: str, payload: dict[str, Any] | None = None
|
||||
) -> dict[str, Any]:
|
||||
async with httpx.AsyncClient(timeout=20.0) as client:
|
||||
response = await client.post(
|
||||
f"{SLACK_API}/{method}",
|
||||
|
|
@ -55,7 +57,9 @@ class SlackGatewayClient:
|
|||
ts: str,
|
||||
text: str,
|
||||
) -> PlatformSendResult:
|
||||
data = await self.api_call("chat.update", {"channel": channel, "ts": ts, "text": text})
|
||||
data = await self.api_call(
|
||||
"chat.update", {"channel": channel, "ts": ts, "text": text}
|
||||
)
|
||||
return PlatformSendResult(
|
||||
external_message_id=str(data.get("ts") or ts),
|
||||
raw_response=data,
|
||||
|
|
|
|||
|
|
@ -41,7 +41,9 @@ class SlackStreamTranslator(BaseStreamTranslator):
|
|||
async def translate(self, events: AsyncIterator[GatewayStreamEvent]) -> None:
|
||||
async for event in events:
|
||||
if event.type in {"text-delta", "text_delta", "text"}:
|
||||
self._buffer += str(event.data.get("text") or event.data.get("delta") or "")
|
||||
self._buffer += str(
|
||||
event.data.get("text") or event.data.get("delta") or ""
|
||||
)
|
||||
elif event.type in {"data-interrupt-request", "interrupt"}:
|
||||
await self._handle_hitl_interrupt()
|
||||
return
|
||||
|
|
@ -53,7 +55,9 @@ class SlackStreamTranslator(BaseStreamTranslator):
|
|||
async def _flush_final(self) -> None:
|
||||
if not self._buffer:
|
||||
return
|
||||
for chunk in split_text_message(self._buffer, max_chars=SLACK_MAX_MESSAGE_CHARS):
|
||||
for chunk in split_text_message(
|
||||
self._buffer, max_chars=SLACK_MAX_MESSAGE_CHARS
|
||||
):
|
||||
await self._send_text(chunk)
|
||||
|
||||
async def _send_text(self, text: str) -> PlatformSendResult:
|
||||
|
|
|
|||
|
|
@ -1,2 +1 @@
|
|||
"""Telegram gateway adapter."""
|
||||
|
||||
|
|
|
|||
|
|
@ -51,9 +51,7 @@ class TelegramAdapter(BasePlatformAdapter):
|
|||
"channel": "channel",
|
||||
}.get(chat_type, "unknown")
|
||||
display_name = chat.get("title") or " ".join(
|
||||
part
|
||||
for part in (sender.get("first_name"), sender.get("last_name"))
|
||||
if part
|
||||
part for part in (sender.get("first_name"), sender.get("last_name")) if part
|
||||
)
|
||||
|
||||
return ParsedInboundEvent(
|
||||
|
|
@ -62,14 +60,21 @@ class TelegramAdapter(BasePlatformAdapter):
|
|||
external_peer_id=str(chat["id"]) if chat.get("id") is not None else None,
|
||||
external_peer_kind=peer_kind,
|
||||
external_message_id=(
|
||||
str(message["message_id"]) if message.get("message_id") is not None else None
|
||||
str(message["message_id"])
|
||||
if message.get("message_id") is not None
|
||||
else None
|
||||
),
|
||||
external_user_id=str(sender["id"]) if sender.get("id") is not None else None,
|
||||
external_user_id=str(sender["id"])
|
||||
if sender.get("id") is not None
|
||||
else None,
|
||||
text=message.get("text") or message.get("caption"),
|
||||
raw_payload=raw_payload,
|
||||
display_name=display_name or None,
|
||||
username=sender.get("username") or chat.get("username"),
|
||||
metadata={"chat_type": chat_type, "update_id": raw_payload.get("update_id")},
|
||||
metadata={
|
||||
"chat_type": chat_type,
|
||||
"update_id": raw_payload.get("update_id"),
|
||||
},
|
||||
)
|
||||
|
||||
async def send_message(
|
||||
|
|
@ -108,7 +113,8 @@ class TelegramAdapter(BasePlatformAdapter):
|
|||
async def leave_chat(self, *, external_peer_id: str) -> None:
|
||||
await self.client.leave_chat(chat_id=external_peer_id)
|
||||
|
||||
async def fetch_updates(self, *, offset: int | None) -> AsyncIterator[dict[str, Any]]:
|
||||
async def fetch_updates(
|
||||
self, *, offset: int | None
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
async for update in self.client.get_updates(offset=offset):
|
||||
yield update
|
||||
|
||||
|
|
|
|||
|
|
@ -106,4 +106,3 @@ async def retry_plaintext_on_bad_markdown(call, *args, **kwargs) -> PlatformSend
|
|||
raise
|
||||
kwargs["parse_mode"] = None
|
||||
return await call(*args, **kwargs)
|
||||
|
||||
|
|
|
|||
|
|
@ -54,7 +54,9 @@ async def handle_start_command(
|
|||
return True
|
||||
|
||||
|
||||
async def handle_help_command(*, adapter: TelegramAdapter, event: ParsedInboundEvent) -> bool:
|
||||
async def handle_help_command(
|
||||
*, adapter: TelegramAdapter, event: ParsedInboundEvent
|
||||
) -> bool:
|
||||
if not event.external_peer_id:
|
||||
return True
|
||||
await adapter.send_message(external_peer_id=event.external_peer_id, text=HELP_TEXT)
|
||||
|
|
@ -114,4 +116,4 @@ class TelegramGatewayCommands(BaseGatewayCommands):
|
|||
adapter=adapter,
|
||||
event=event,
|
||||
dashboard_url=dashboard_url,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -32,9 +32,13 @@ def _split_at_boundary(text: str, max_units: int) -> tuple[str, str]:
|
|||
end -= 1
|
||||
|
||||
candidate = text[:end]
|
||||
boundary = max(candidate.rfind("\n\n"), candidate.rfind(". "), candidate.rfind("\n"))
|
||||
boundary = max(
|
||||
candidate.rfind("\n\n"), candidate.rfind(". "), candidate.rfind("\n")
|
||||
)
|
||||
if boundary > max(200, end // 2):
|
||||
end = boundary + (2 if candidate[boundary : boundary + 2] in {"\n\n", ". "} else 1)
|
||||
end = boundary + (
|
||||
2 if candidate[boundary : boundary + 2] in {"\n\n", ". "} else 1
|
||||
)
|
||||
|
||||
return text[:end], text[end:]
|
||||
|
||||
|
|
@ -56,4 +60,3 @@ def chunk_message(
|
|||
chunks.append(chunk)
|
||||
return chunks
|
||||
return split_text_message(text, max_chars=max_units)
|
||||
|
||||
|
|
|
|||
|
|
@ -49,7 +49,9 @@ class TelegramStreamTranslator(BaseStreamTranslator):
|
|||
async def translate(self, events: AsyncIterator[GatewayStreamEvent]) -> None:
|
||||
async for event in events:
|
||||
if event.type in {"text-delta", "text_delta", "text"}:
|
||||
self._buffer += str(event.data.get("text") or event.data.get("delta") or "")
|
||||
self._buffer += str(
|
||||
event.data.get("text") or event.data.get("delta") or ""
|
||||
)
|
||||
await self._maybe_flush()
|
||||
elif event.type in {"data-interrupt-request", "interrupt"}:
|
||||
await self._handle_hitl_interrupt()
|
||||
|
|
@ -159,7 +161,9 @@ class TelegramStreamTranslator(BaseStreamTranslator):
|
|||
)
|
||||
if chat_wait:
|
||||
record_gateway_rate_limit_hit(bucket="tg:chat")
|
||||
global_wait = await wait_for_token("tg:global", capacity=25, refill_per_sec=25.0)
|
||||
global_wait = await wait_for_token(
|
||||
"tg:global", capacity=25, refill_per_sec=25.0
|
||||
)
|
||||
if global_wait:
|
||||
record_gateway_rate_limit_hit(bucket="tg:global")
|
||||
|
||||
|
|
@ -168,4 +172,3 @@ class TelegramStreamTranslator(BaseStreamTranslator):
|
|||
await self._flush(final=False)
|
||||
await self._send_text(HITL_UNSUPPORTED_MESSAGE)
|
||||
record_gateway_hitl_aborted(platform="telegram")
|
||||
|
||||
|
|
|
|||
|
|
@ -36,5 +36,6 @@ def release_thread_lock(thread_id: int) -> None:
|
|||
try:
|
||||
_redis().delete(_lock_key(thread_id))
|
||||
except redis.RedisError as exc:
|
||||
logger.warning("Failed to release gateway thread lock for %s: %s", thread_id, exc)
|
||||
|
||||
logger.warning(
|
||||
"Failed to release gateway thread lock for %s: %s", thread_id, exc
|
||||
)
|
||||
|
|
|
|||
|
|
@ -36,7 +36,8 @@ class WhatsAppBaileysAdapter(BasePlatformAdapter):
|
|||
external_user_id=sender_id or None,
|
||||
text=str(body) if body is not None else None,
|
||||
raw_payload=raw_payload,
|
||||
display_name=str(raw_payload.get("chatName") or sender_id or chat_id) or None,
|
||||
display_name=str(raw_payload.get("chatName") or sender_id or chat_id)
|
||||
or None,
|
||||
username=None,
|
||||
metadata={
|
||||
"sender_id": sender_id,
|
||||
|
|
@ -92,7 +93,9 @@ class WhatsAppBaileysAdapter(BasePlatformAdapter):
|
|||
response.raise_for_status()
|
||||
return response.json()
|
||||
|
||||
async def fetch_updates(self, *, offset: int | None) -> AsyncIterator[dict[str, Any]]:
|
||||
async def fetch_updates(
|
||||
self, *, offset: int | None
|
||||
) -> AsyncIterator[dict[str, Any]]:
|
||||
async with httpx.AsyncClient(timeout=35) as client:
|
||||
response = await client.get(f"{self.bridge_url}/messages")
|
||||
response.raise_for_status()
|
||||
|
|
|
|||
|
|
@ -54,7 +54,9 @@ class WhatsAppCloudAdapter(BasePlatformAdapter):
|
|||
username=None,
|
||||
metadata={
|
||||
"phone_number_id": _metadata(raw_payload).get("phone_number_id"),
|
||||
"display_phone_number": _metadata(raw_payload).get("display_phone_number"),
|
||||
"display_phone_number": _metadata(raw_payload).get(
|
||||
"display_phone_number"
|
||||
),
|
||||
"timestamp": message.get("timestamp"),
|
||||
"message_type": message.get("type"),
|
||||
},
|
||||
|
|
@ -96,7 +98,9 @@ def _changes(raw_payload: dict[str, Any]) -> list[dict[str, Any]]:
|
|||
for entry in raw_payload.get("entry") or []:
|
||||
if isinstance(entry, dict):
|
||||
changes.extend(
|
||||
change for change in (entry.get("changes") or []) if isinstance(change, dict)
|
||||
change
|
||||
for change in (entry.get("changes") or [])
|
||||
if isinstance(change, dict)
|
||||
)
|
||||
return changes
|
||||
|
||||
|
|
|
|||
|
|
@ -18,8 +18,7 @@ class WhatsAppCredentials(TypedDict, total=False):
|
|||
|
||||
def load_system_whatsapp_credentials() -> WhatsAppCredentials:
|
||||
if not (
|
||||
config.WHATSAPP_SHARED_BUSINESS_TOKEN
|
||||
and config.WHATSAPP_SHARED_PHONE_NUMBER_ID
|
||||
config.WHATSAPP_SHARED_BUSINESS_TOKEN and config.WHATSAPP_SHARED_PHONE_NUMBER_ID
|
||||
):
|
||||
raise RuntimeError("whatsapp_system_credentials_not_configured")
|
||||
|
||||
|
|
|
|||
|
|
@ -41,7 +41,9 @@ class WhatsAppCloudStreamTranslator(BaseStreamTranslator):
|
|||
if event.type in {"text-delta", "text_delta", "text"}:
|
||||
if not self._typing_sent:
|
||||
await self._send_typing_indicator()
|
||||
self._buffer += str(event.data.get("text") or event.data.get("delta") or "")
|
||||
self._buffer += str(
|
||||
event.data.get("text") or event.data.get("delta") or ""
|
||||
)
|
||||
elif event.type in {"data-interrupt-request", "interrupt"}:
|
||||
await self._handle_hitl_interrupt()
|
||||
return
|
||||
|
|
|
|||
|
|
@ -42,7 +42,9 @@ class WhatsAppBaileysStreamTranslator(BaseStreamTranslator):
|
|||
await self._send_typing_indicator()
|
||||
async for event in events:
|
||||
if event.type in {"text-delta", "text_delta", "text"}:
|
||||
self._buffer += str(event.data.get("text") or event.data.get("delta") or "")
|
||||
self._buffer += str(
|
||||
event.data.get("text") or event.data.get("delta") or ""
|
||||
)
|
||||
await self._maybe_flush()
|
||||
elif event.type in {"data-interrupt-request", "interrupt"}:
|
||||
await self._handle_hitl_interrupt()
|
||||
|
|
@ -86,7 +88,9 @@ class WhatsAppBaileysStreamTranslator(BaseStreamTranslator):
|
|||
if not isinstance(self.adapter, WhatsAppBaileysAdapter):
|
||||
return
|
||||
try:
|
||||
await self.adapter.send_typing_indicator(external_peer_id=self.external_peer_id)
|
||||
await self.adapter.send_typing_indicator(
|
||||
external_peer_id=self.external_peer_id
|
||||
)
|
||||
record_gateway_outbound(platform="whatsapp", kind="typing", status="sent")
|
||||
except Exception:
|
||||
logger.debug("WhatsApp Baileys typing indicator failed", exc_info=True)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue