chore: linting

This commit is contained in:
DESKTOP-RTLN3BA\$punk 2026-06-09 00:42:26 -07:00
parent 0a012dbc79
commit ce952d2ad1
127 changed files with 821 additions and 517 deletions

View file

@ -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()

View file

@ -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)

View file

@ -52,4 +52,3 @@ async def assert_authorization_invariant(
await _fail(session, binding, f"rbac_{exc.status_code}")
return user

View file

@ -1,2 +1 @@
"""Base gateway interfaces."""

View file

@ -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")

View file

@ -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()

View file

@ -25,4 +25,3 @@ class BaseStreamTranslator(ABC):
@abstractmethod
async def translate(self, events: AsyncIterator[GatewayStreamEvent]) -> None:
"""Consume agent stream events and emit platform messages."""

View file

@ -64,4 +64,3 @@ def resume_binding(binding: ExternalChatBinding) -> None:
binding.state = ExternalChatBindingState.BOUND
binding.suspended_at = None
binding.suspended_reason = None

View file

@ -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

View file

@ -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

View file

@ -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:

View file

@ -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]

View file

@ -51,4 +51,3 @@ async def persist_inbound_event(
)
result = await session.execute(stmt)
return result.scalar_one_or_none()

View file

@ -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()

View file

@ -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

View file

@ -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

View file

@ -133,4 +133,3 @@ async def wait_for_token(
if wait_ms > 0:
await asyncio.sleep(wait_ms / 1000)
return wait_ms

View file

@ -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}"
)

View file

@ -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}
)

View file

@ -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(

View file

@ -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,

View file

@ -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:

View file

@ -1,2 +1 @@
"""Telegram gateway adapter."""

View file

@ -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

View file

@ -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)

View file

@ -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,
)
)

View file

@ -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)

View file

@ -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")

View file

@ -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
)

View file

@ -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()

View file

@ -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

View file

@ -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")

View file

@ -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

View file

@ -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)