diff --git a/gateway/hosted_room_discussion.py b/gateway/hosted_room_discussion.py index b4aac19951..3118e1c6ad 100644 --- a/gateway/hosted_room_discussion.py +++ b/gateway/hosted_room_discussion.py @@ -432,13 +432,13 @@ def _validate_member_message( raise DiscussionValidationError("message.member text must be a non-pass string") member = _member_by_id(room, payload.get("member_id")) peer = _peer_target(member) - expected_connection = peer.get("peer_id") if peer else None - if ( - actor.get("kind") != "member" - or actor.get("id") != member.member_id - or actor.get("profile") != member.profile - or actor.get("connection_id") != expected_connection - ): + expected = { + "kind": "member", + "id": member.member_id, + "profile": member.profile, + "connection_id": peer.get("peer_id") if peer else None, + } + if any(actor.get(key) != value for key, value in expected.items()): raise DiscussionValidationError("message.member actor does not match roster") return payload @@ -510,15 +510,14 @@ def _validate_event(raw: Any, *, room: DiscussionRoom, previous_seq: int) -> _Va if seq <= previous_seq: raise DiscussionValidationError("room events must be in strict sequence order") event_id = _identifier(raw.get("event_id"), label="event_id") - kind = raw.get("kind") - if not isinstance(kind, str): - raise DiscussionValidationError("event kind must be a string") - actor = raw.get("actor") - if not isinstance(actor, Mapping): - raise DiscussionValidationError("event actor must be an object") - payload = raw.get("payload") - if not isinstance(payload, Mapping): - raise DiscussionValidationError("event payload must be an object") + kind, actor, payload = raw.get("kind"), raw.get("actor"), raw.get("payload") + for value, expected, message in ( + (kind, str, "event kind must be a string"), + (actor, Mapping, "event actor must be an object"), + (payload, Mapping, "event payload must be an object"), + ): + if not isinstance(value, expected): + raise DiscussionValidationError(message) if kind in _EPOCH_STAMPED_KINDS and raw.get("authority_epoch") != room.authority_epoch: raise DiscussionValidationError(f"{kind} authority epoch does not match the room") validator = _EVENT_VALIDATORS.get(kind) @@ -576,11 +575,9 @@ def _derive_member_watermarks(events: Sequence[_ValidatedEvent]) -> dict[tuple[s watermark = int(event.payload["seen_through_seq"]) if event.kind == "turn.settled" and not event.payload["passed"]: message = messages_by_id.get(str(event.payload["message_event_id"])) - if ( - message is None - or message.payload.get("task_id") != task_id - or message.payload.get("member_id") != event.payload.get("member_id") - or message.payload.get("thread_id") != event.payload.get("thread_id") + if message is None or any( + message.payload.get(field) != event.payload.get(field) + for field in ("task_id", "member_id", "thread_id") ): raise DiscussionValidationError("turn.settled references no matching member message") watermark = max(watermark, message.seq) @@ -1021,12 +1018,9 @@ def plan_publication( raise DiscussionValidationError("task member is not in the frozen roster") if status not in _TERMINAL_EFFECTS: raise DiscussionValidationError("invalid terminal publication status") - if status == "deferred" and ( - isinstance(execution_generation, bool) - or not isinstance(execution_generation, int) - or execution_generation < 1 - ): - raise DiscussionValidationError("deferred publication requires an execution generation") + if status == "deferred": + message = "deferred publication requires an execution generation" + common.positive_int(execution_generation, error=DiscussionValidationError, message=message) newer_same_thread = any( event.kind == "message.user" diff --git a/gateway/hosted_room_peer.py b/gateway/hosted_room_peer.py index a1e547ee2b..31c302456f 100644 --- a/gateway/hosted_room_peer.py +++ b/gateway/hosted_room_peer.py @@ -19,7 +19,7 @@ import stat import time import urllib.parse from dataclasses import asdict, dataclass -from functools import lru_cache +from functools import lru_cache, partial from pathlib import Path from typing import Any, Callable, Iterable, Literal, Mapping @@ -146,13 +146,10 @@ def _digest(value: Any, *, field: str) -> str: return value -def _exact_fields( - value: Mapping[str, Any], *, required: set[str], optional: set[str] = frozenset(), label: str -) -> None: - exact_fields( - value, label=label, required=required, optional=optional, error=HostedRoomPeerError, - missing_fmt="{label} missing fields: {fields}", unknown_fmt="{label} unknown fields: {fields}", - ) +_exact_fields = partial( + exact_fields, error=HostedRoomPeerError, missing_fmt="{label} missing fields: {fields}", + unknown_fmt="{label} unknown fields: {fields}", +) def _canonical_json(value: Mapping[str, Any]) -> bytes: diff --git a/gateway/hosted_room_policy_checkpoint.py b/gateway/hosted_room_policy_checkpoint.py index 0d55c694a9..a763492db9 100644 --- a/gateway/hosted_room_policy_checkpoint.py +++ b/gateway/hosted_room_policy_checkpoint.py @@ -92,6 +92,20 @@ def _require_room(conn: sqlite3.Connection, room_id: str) -> None: raise hosted_rooms.RoomNotFoundError("hosted room not found") +def _settled_message( + conn: sqlite3.Connection, room_id: str, discussion_event_id: str, message_event_id: Any +) -> dict[str, Any] | None: + """Return the indexed member message a ``turn.settled`` event committed, if it is in the projection.""" + rows = conn.execute( + "SELECT seq, event_json FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=?", + (room_id, discussion_event_id), + ).fetchall() + return next( + (m for m in (json.loads(row["event_json"]) for row in rows) if m.get("event_id") == message_event_id), + None, + ) + + class HostedRoomPolicyCheckpoint: """Incrementally index room policy without compacting visible history.""" @@ -256,22 +270,10 @@ class HostedRoomPolicyCheckpoint: member_id = str(payload.get("member_id") or "") seen_through_seq = int(payload.get("seen_through_seq") or 0) if kind == "turn.settled" and payload.get("message_event_id"): - messages = conn.execute( - "SELECT seq, event_json FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=?", - (room_id, discussion_event_id), - ).fetchall() - committed = next( - ( - message for message in (json.loads(row["event_json"]) for row in messages) - if message.get("event_id") == payload["message_event_id"] - ), - None, - ) + committed = _settled_message(conn, room_id, discussion_event_id, payload["message_event_id"]) if committed is not None: seen_through_seq = max(seen_through_seq, int(committed["seq"])) - self._store_transcript_event( - conn, event=committed, thread_id=thread_id, settled_seq=int(event["seq"]) - ) + self._store_transcript_event(conn, event=committed, thread_id=thread_id, settled_seq=seq) if member_id and seen_through_seq > 0: conn.execute( """INSERT INTO hosted_room_policy_watermarks(