from __future__ import annotations from dataclasses import dataclass from enum import StrEnum from slack_sdk.errors import SlackApiError from slack_sdk.web.async_client import AsyncWebClient from yuxi.channels.models import DeliveryResult from yuxi.utils.logging_config import logger class StreamState(StrEnum): INIT = "init" STREAMING = "streaming" STOPPED = "stopped" ERROR = "error" _BENIGN_FINALIZE_ERRORS = frozenset( { "user_not_found", "team_not_found", "missing_recipient_user_id", "not_in_channel", "channel_not_found", } ) class SlackStreamNotDeliveredError(Exception): pass @dataclass class NativeStreamSession: channel: str team_id: str = "" user_id: str = "" message_ts: str = "" accumulated_text: str = "" state: StreamState = StreamState.INIT error: str = "" _MISSING_OVERRIDE = object() class NativeStreamManager: def __init__(self, client: AsyncWebClient): self._client = client self._sessions: dict[str, NativeStreamSession] = {} async def start_stream( self, channel: str, *, team_id: str = "", user_id: str = "", text: str = "", ) -> DeliveryResult: session_key = f"{channel}:{team_id}:{user_id}" existing = self._sessions.get(session_key) if existing and existing.state == StreamState.STREAMING: return await self.append_stream(channel, text, team_id=team_id, user_id=user_id) try: params: dict = { "channel": channel, "text": text or "...", } if user_id: params["user_id"] = user_id if team_id: params["team_id"] = team_id result = await self._client.api_call( api_method="chat.startStream", http_verb="POST", params=params, ) ok = result.get("ok", False) ts = result.get("ts", "") session = NativeStreamSession( channel=channel, team_id=team_id, user_id=user_id, message_ts=ts, accumulated_text=text, state=StreamState.STREAMING if ok else StreamState.ERROR, error=result.get("error", ""), ) self._sessions[session_key] = session if ok: logger.debug(f"Native stream started: channel={channel}, ts={ts}") return DeliveryResult(success=True, message_id=ts) logger.warning(f"Native stream start failed: {result.get('error')}") return DeliveryResult(success=False, error=result.get("error", "Unknown error")) except SlackApiError as e: err = e.response.get("error", str(e)) logger.error(f"Native stream start error: {err}") return DeliveryResult(success=False, error=err) async def append_stream( self, channel: str, text: str, *, team_id: str = "", user_id: str = "", ts: str = "", ) -> DeliveryResult: session_key = f"{channel}:{team_id}:{user_id}" session = self._sessions.get(session_key) if not session or session.state != StreamState.STREAMING: return await self.start_stream(channel, team_id=team_id, user_id=user_id, text=text) try: result = await self._client.api_call( api_method="chat.appendStream", http_verb="POST", params={ "channel": channel, "ts": ts or session.message_ts, "text": text, }, ) ok = result.get("ok", False) if ok: session.accumulated_text = text return DeliveryResult( success=ok, error=result.get("error") if not ok else None, ) except SlackApiError as e: err = e.response.get("error", str(e)) logger.error(f"Native stream append error: {err}") return DeliveryResult(success=False, error=err) async def stop_stream( self, channel: str, *, team_id: str = "", user_id: str = "", ts: str = "", final_text: str = _MISSING_OVERRIDE, ) -> DeliveryResult: session_key = f"{channel}:{team_id}:{user_id}" session = self._sessions.pop(session_key, None) try: result = await self._client.api_call( api_method="chat.stopStream", http_verb="POST", params={ "channel": channel, "ts": ts or (session.message_ts if session else ""), "text": final_text if final_text is not _MISSING_OVERRIDE else (session.accumulated_text if session else ""), }, ) ok = result.get("ok", False) err = result.get("error", "") if not ok and not _is_benign_stop_error(err): logger.error(f"Native stream stop error: {err}") raise SlackStreamNotDeliveredError(f"Stream not delivered: {err}") return DeliveryResult( success=True, error=err if err and err not in _BENIGN_FINALIZE_ERRORS else None, ) except SlackApiError as e: err = e.response.get("error", str(e)) if _is_benign_stop_error(err): return DeliveryResult(success=True) logger.error(f"Native stream stop API error: {err}") raise SlackStreamNotDeliveredError(f"Stream not delivered: {err}") from e def _is_benign_stop_error(error: str) -> bool: return error in _BENIGN_FINALIZE_ERRORS