from __future__ import annotations from typing import TYPE_CHECKING, Any from yuxi.channels.models import DeliveryResult from yuxi.utils.logging_config import logger if TYPE_CHECKING: from nio import AsyncClient _POLL_KIND_MAP = { "m.text": "open", "disclosed": "disclosed", } async def create_poll( client: AsyncClient, room_id: str, question: str, answers: list[str], kind: str = "open", max_selections: int = 1, ) -> DeliveryResult: content = { "org.matrix.msc3381.poll.start": { "question": {"body": question, "kind": kind, "msgtype": "m.text"}, "answers": [{"id": f"answer_{i}", "org.matrix.msc1767.text": ans} for i, ans in enumerate(answers)], "max_selections": max_selections, }, "msgtype": "m.text", "body": f"Poll: {question}\n" + "\n".join(f"{i + 1}. {ans}" for i, ans in enumerate(answers)), } try: resp = await client.room_send( room_id=room_id, message_type="m.room.message", content=content, ) return DeliveryResult(success=True, message_id=resp.event_id) except Exception as e: logger.error(f"Matrix poll creation failed: {e}") return DeliveryResult(success=False, error=str(e)) async def send_poll_response( client: AsyncClient, room_id: str, poll_event_id: str, answer_ids: list[str], ) -> DeliveryResult: content = { "m.relates_to": { "rel_type": "m.reference", "event_id": poll_event_id, }, "org.matrix.msc3381.poll.response": { "answers": answer_ids, }, } try: resp = await client.room_send( room_id=room_id, message_type="m.room.message", content=content, ) return DeliveryResult(success=True, message_id=resp.event_id) except Exception as e: logger.error(f"Matrix poll vote failed: {e}") return DeliveryResult(success=False, error=str(e)) async def end_poll( client: AsyncClient, room_id: str, poll_event_id: str, summary: str = "", ) -> DeliveryResult: content = { "m.relates_to": { "rel_type": "m.reference", "event_id": poll_event_id, }, "org.matrix.msc3381.poll.end": {}, } if summary: content["body"] = summary try: resp = await client.room_send( room_id=room_id, message_type="m.room.message", content=content, ) return DeliveryResult(success=True, message_id=resp.event_id) except Exception as e: logger.error(f"Matrix poll end failed: {e}") return DeliveryResult(success=False, error=str(e)) def summarize_poll_responses(responses: list[dict[str, Any]]) -> dict[str, dict[str, int]]: tally: dict[str, dict[str, int]] = {} for resp in responses: p = resp.get("poll_response", {}) event_id = resp.get("event_id", "") answers = p.get("answers", []) for ans in answers: if event_id not in tally: tally[event_id] = {} tally[event_id][ans] = tally[event_id].get(ans, 0) + 1 return tally def extract_poll_start(event_source: dict[str, Any]) -> dict[str, Any] | None: content = event_source.get("content", {}) poll_data = content.get("org.matrix.msc3381.poll.start") if not poll_data: return None return { "question": poll_data.get("question", {}).get("body", ""), "kind": poll_data.get("question", {}).get("kind", "open"), "answers": [ {"id": a.get("id", ""), "text": a.get("org.matrix.msc1767.text", "")} for a in poll_data.get("answers", []) ], "max_selections": poll_data.get("max_selections", 1), } def extract_poll_response(event_source: dict[str, Any]) -> dict[str, Any] | None: content = event_source.get("content", {}) poll_data = content.get("org.matrix.msc3381.poll.response") if not poll_data: return None relates_to = content.get("m.relates_to", {}) return { "event_id": relates_to.get("event_id", ""), "answers": poll_data.get("answers", []), } def is_poll_end(event_source: dict[str, Any]) -> bool: content = event_source.get("content", {}) return "org.matrix.msc3381.poll.end" in content