146 lines
4.3 KiB
Python
146 lines
4.3 KiB
Python
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
|