From efff31ccd047ba3ecbbe00c7465c8d901feca465 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:29:57 -0700 Subject: [PATCH] refactor(gateway): AST-neutral repack of hosted-room call/def layouts (<=110 cols) --- gateway/hosted_room_discussion.py | 265 +++++--------- gateway/hosted_room_driver.py | 213 ++++------- gateway/hosted_room_execution_policy.py | 18 +- gateway/hosted_room_links.py | 33 +- gateway/hosted_room_peer.py | 85 ++--- gateway/hosted_room_policy_checkpoint.py | 144 +++----- gateway/hosted_room_replicas.py | 73 ++-- gateway/hosted_rooms.py | 439 +++++++---------------- gateway/hosted_rooms_common.py | 45 +-- 9 files changed, 416 insertions(+), 899 deletions(-) diff --git a/gateway/hosted_room_discussion.py b/gateway/hosted_room_discussion.py index 07f0e5f2de..af9aaae697 100644 --- a/gateway/hosted_room_discussion.py +++ b/gateway/hosted_room_discussion.py @@ -42,8 +42,7 @@ _MENTION_RE = re.compile(r"@([A-Za-z0-9][A-Za-z0-9._:-]*)", re.IGNORECASE) _TURN_ID_RE = re.compile( r"^d(?P[1-9][0-9]*)\.r(?P[0-2])\." r"p(?P[0-5])\.s(?P[1-9][0-9]*)\." - r"m(?P[0-9a-f]{24})$" -) + r"m(?P[0-9a-f]{24})$") _LOCAL_TARGET_FIELDS = frozenset({"kind", "profile"}) _PEER_TARGET_FIELDS = frozenset({"kind", "peer_id", "installation_id", "profile", "capability_digest"}) @@ -51,31 +50,25 @@ _REMOTE_MEMBER_FIELDS = frozenset({ "connectionId", "connectionKind", "connectionLabel", "connection_id", "connection_kind", "connection_label", "remoteSource", "route", "sourceMissing", "sourceReachable", "sourceScoped", "targetProfile", - "target_profile", -}) + "target_profile"}) _USER_PAYLOAD_FIELDS = frozenset({"text", "thread_id"}) -_TURN_COORDINATE_FIELDS = frozenset({ - "discussion_event_id", "member_id", "member_index", "round_index", - "task_id", "thread_id", "turn_id", -}) +_TURN_COORDINATE_FIELDS = frozenset( + {"discussion_event_id", "member_id", "member_index", "round_index", "task_id", "thread_id", "turn_id", }) _MEMBER_MESSAGE_FIELDS = _TURN_COORDINATE_FIELDS | {"text"} _TERMINAL_COMMON_FIELDS = _TURN_COORDINATE_FIELDS | {"seen_through_seq"} _TERMINAL_EXTRA_FIELDS = { "turn.settled": frozenset({"message_event_id", "passed"}), "turn.failed": frozenset({"error"}), "turn.cancelled": frozenset({"reason"}), - "turn.deferred": frozenset({"execution_generation", "reason"}), -} + "turn.deferred": frozenset({"execution_generation", "reason"})} _TERMINAL_OPTIONAL_FIELDS = {"turn.failed": frozenset({"reason_code"})} _TERMINAL_EVENT_KINDS = frozenset(_TERMINAL_EXTRA_FIELDS) # Gateway-authored control events: kind -> (exact payload fields, identifier fields). _GATEWAY_EVENT_FIELDS = { "room.activity": ( frozenset({"status", "reason_code", "thread_id", "discussion_event_id"}), - ("reason_code", "thread_id", "discussion_event_id"), - ), - "room.stop_requested": (frozenset({"cancel_id"}), ("cancel_id",)), -} + ("reason_code", "thread_id", "discussion_event_id")), + "room.stop_requested": (frozenset({"cancel_id"}), ("cancel_id",))} _EPOCH_STAMPED_KINDS = _TERMINAL_EVENT_KINDS | {"message.member", *_GATEWAY_EVENT_FIELDS} @@ -155,8 +148,7 @@ class EventPlan: "room_id": room_id, "event_id": self.event_id, "kind": self.kind, "actor": dict(self.actor), "payload": dict(self.payload), "authority_gateway_id": self.authority_gateway_id, - "authority_epoch": self.authority_epoch, - } + "authority_epoch": self.authority_epoch} @dataclass(frozen=True) @@ -179,8 +171,7 @@ class _ValidatedEvent: _identifier = partial( - common.identifier, error=DiscussionValidationError, max_chars=driver.MAX_IDENTIFIER_CHARS -) + common.identifier, error=DiscussionValidationError, max_chars=driver.MAX_IDENTIFIER_CHARS) _exact_fields = partial(common.exact_fields, error=DiscussionValidationError) _bounded_int = partial(common.bounded_int, error=DiscussionValidationError) @@ -209,13 +200,11 @@ def validate_user_payload(value: Any) -> dict[str, Any]: def _validate_member_target( - value: Any, *, profile: str, known_profiles: set[str], index: int -) -> dict[str, Any]: + value: Any, *, profile: str, known_profiles: set[str], index: int) -> dict[str, Any]: if value is None: if profile not in known_profiles: raise DiscussionValidationError( - f"member {index} profile '{profile}' is not local to this gateway" - ) + f"member {index} profile '{profile}' is not local to this gateway") return {"kind": "local", "profile": profile} if not isinstance(value, Mapping): raise DiscussionValidationError(f"member {index} target must be an object") @@ -224,19 +213,16 @@ def _validate_member_target( raise DiscussionValidationError(f"member {index} target kind must be local or peer") target = _exact_fields( value, label=f"member {index} {kind} target", - required=_LOCAL_TARGET_FIELDS if kind == "local" else _PEER_TARGET_FIELDS, - ) + required=_LOCAL_TARGET_FIELDS if kind == "local" else _PEER_TARGET_FIELDS) target_profile = _identifier(target["profile"], label=f"member {index} target profile") if kind == "local": if target_profile != profile or profile not in known_profiles: raise DiscussionValidationError( - f"member {index} local target does not match a local profile" - ) + f"member {index} local target does not match a local profile") return {"kind": "local", "profile": profile} if target_profile != profile: raise DiscussionValidationError( - f"member {index} peer target profile does not match member profile" - ) + f"member {index} peer target profile does not match member profile") capability_digest = target["capability_digest"] if not isinstance(capability_digest, str) or not re.fullmatch(r"[0-9a-f]{64}", capability_digest): raise DiscussionValidationError(f"member {index} capability_digest must be a sha256 digest") @@ -245,8 +231,7 @@ def _validate_member_target( "peer_id": _identifier(target["peer_id"], label=f"member {index} peer_id"), "installation_id": _identifier(target["installation_id"], label=f"member {index} installation_id"), "profile": target_profile, - "capability_digest": capability_digest, - } + "capability_digest": capability_digest} def validate_roster(value: Any, *, local_profiles: Iterable[str]) -> tuple[DiscussionMember, ...]: @@ -256,8 +241,7 @@ def validate_roster(value: Any, *, local_profiles: Iterable[str]) -> tuple[Discu if not MIN_DISCUSSION_MEMBERS <= len(value) <= MAX_DISCUSSION_MEMBERS: raise DiscussionValidationError( f"members must contain between {MIN_DISCUSSION_MEMBERS} and " - f"{MAX_DISCUSSION_MEMBERS} entries" - ) + f"{MAX_DISCUSSION_MEMBERS} entries") known_profiles = {_identifier(profile, label="local profile") for profile in local_profiles} members: list[DiscussionMember] = [] @@ -271,19 +255,16 @@ def validate_roster(value: Any, *, local_profiles: Iterable[str]) -> tuple[Discu remote_fields = frozenset(raw) & _REMOTE_MEMBER_FIELDS if remote_fields: raise DiscussionValidationError( - f"member {index} contains cross-gateway fields: {', '.join(sorted(remote_fields))}" - ) + f"member {index} contains cross-gateway fields: {', '.join(sorted(remote_fields))}") member = _exact_fields( raw, label=f"member {index}", required=frozenset({"member_id", "profile", "handle"}), - optional=frozenset({"display_name", "target"}), - ) + optional=frozenset({"display_name", "target"})) member_id = _identifier(member["member_id"], label=f"member {index} id") profile = _identifier(member["profile"], label=f"member {index} profile") handle = _identifier(member["handle"], label=f"member {index} handle") target = _validate_member_target( - member.get("target"), profile=profile, known_profiles=known_profiles, index=index - ) + member.get("target"), profile=profile, known_profiles=known_profiles, index=index) display_name = member.get("display_name", "") if not isinstance(display_name, str): raise DiscussionValidationError(f"member {index} display_name must be a string") @@ -294,21 +275,18 @@ def validate_roster(value: Any, *, local_profiles: Iterable[str]) -> tuple[Discu target_key = compact_json(target, ensure_ascii=False).casefold() target_message = ( "member profiles must be unique" if target.get("kind") == "local" - else "member targets must be unique" - ) + else "member targets must be unique") for key, seen, message in ( (target_key, targets, target_message), (handle.casefold(), handles, "member handles must be unique and cannot reserve @all or @everyone"), - (member_id.casefold(), member_ids, "member ids must be unique"), - ): + (member_id.casefold(), member_ids, "member ids must be unique")): if key in seen: raise DiscussionValidationError(message) seen.add(key) members.append(DiscussionMember( member_id=member_id, profile=profile, handle=handle, - display_name=display_name, target=target, - )) + display_name=display_name, target=target)) return tuple(members) @@ -329,9 +307,7 @@ def validate_room(value: Any, *, local_profiles: Iterable[str]) -> DiscussionRoo authority_epoch = _positive_int(value.get("authority_epoch"), label="authority_epoch") members = validate_roster(value.get("members"), local_profiles=local_profiles) return DiscussionRoom( - room_id=room_id, name=name, members=members, - gateway_id=gateway_id, authority_epoch=authority_epoch, - ) + room_id=room_id, name=name, members=members, gateway_id=gateway_id, authority_epoch=authority_epoch) def is_pass_text(value: Any) -> bool: @@ -360,8 +336,7 @@ def resolve_mentions( def _unaddressed_member_mentions( - messages: Sequence[_ValidatedEvent], room: DiscussionRoom -) -> tuple[DiscussionMember, ...]: + messages: Sequence[_ValidatedEvent], room: DiscussionRoom) -> tuple[DiscussionMember, ...]: """Return peers explicitly cited by a Bot and not heard from afterward.""" cited_at: dict[str, int] = {} last_post_at: dict[str, int] = {} @@ -377,8 +352,7 @@ def _unaddressed_member_mentions( return tuple( member for member in room.members if member.member_id in cited_at - and last_post_at.get(member.member_id, 0) <= cited_at[member.member_id] - ) + and last_post_at.get(member.member_id, 0) <= cited_at[member.member_id]) def _require_gateway_actor(actor: Mapping[str, Any], room: DiscussionRoom, message: str) -> None: @@ -429,8 +403,7 @@ def _validate_member_message( "kind": "member", "id": member.member_id, "profile": member.profile, - "connection_id": peer.get("peer_id") if peer else None, - } + "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 @@ -442,8 +415,7 @@ def _validate_terminal_event( _exact_fields( payload, label=f"{kind} payload", required=_TERMINAL_COMMON_FIELDS | _TERMINAL_EXTRA_FIELDS[kind], - optional=_TERMINAL_OPTIONAL_FIELDS.get(kind, frozenset()), - ) + optional=_TERMINAL_OPTIONAL_FIELDS.get(kind, frozenset())) _validate_turn_coordinates(payload, room) _positive_int(payload.get("seen_through_seq"), label="seen_through_seq") if actor.get("connection_id") is not None: @@ -467,9 +439,7 @@ def _validate_terminal_event( from tools.bot_failure_reasons import ALL_REASONS if payload["reason_code"] not in ALL_REASONS: - raise DiscussionValidationError( - "turn.failed reason_code must use the shared failure vocabulary" - ) + raise DiscussionValidationError("turn.failed reason_code must use the shared failure vocabulary") return payload @@ -490,8 +460,7 @@ _EVENT_VALIDATORS = { "message.user": _validate_user_event, "message.member": _validate_member_message, **dict.fromkeys(_TERMINAL_EVENT_KINDS, _validate_terminal_event), - **dict.fromkeys(_GATEWAY_EVENT_FIELDS, _validate_gateway_event), -} + **dict.fromkeys(_GATEWAY_EVENT_FIELDS, _validate_gateway_event)} def _validate_event(raw: Any, *, room: DiscussionRoom, previous_seq: int) -> _ValidatedEvent: @@ -507,8 +476,7 @@ def _validate_event(raw: Any, *, room: DiscussionRoom, previous_seq: int) -> _Va 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"), - ): + (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: @@ -520,8 +488,7 @@ def _validate_event(raw: Any, *, room: DiscussionRoom, previous_seq: int) -> _Va def _validated_events( - events: Sequence[Mapping[str, Any]], *, room: DiscussionRoom -) -> tuple[_ValidatedEvent, ...]: + events: Sequence[Mapping[str, Any]], *, room: DiscussionRoom) -> tuple[_ValidatedEvent, ...]: validated: list[_ValidatedEvent] = [] previous_seq = 0 event_ids: set[str] = set() @@ -555,14 +522,12 @@ def _derive_member_watermarks(events: Sequence[_ValidatedEvent]) -> dict[tuple[s if previous is not None: if previous.kind != "turn.deferred": raise DiscussionValidationError( - f"task '{task_id}' has more than one terminal room event" - ) + f"task '{task_id}' has more than one terminal room event") if event.kind == "turn.deferred" and int( event.payload["execution_generation"] ) <= int(previous.payload["execution_generation"]): raise DiscussionValidationError( - f"task '{task_id}' deferral generation did not advance" - ) + f"task '{task_id}' deferral generation did not advance") terminal_by_task[task_id] = event key = (str(event.payload["thread_id"]), str(event.payload["member_id"])) watermark = int(event.payload["seen_through_seq"]) @@ -570,8 +535,7 @@ def _derive_member_watermarks(events: Sequence[_ValidatedEvent]) -> dict[tuple[s message = messages_by_id.get(str(event.payload["message_event_id"])) if message is None or any( message.payload.get(field) != event.payload.get(field) - for field in ("task_id", "member_id", "thread_id") - ): + for field in ("task_id", "member_id", "thread_id")): raise DiscussionValidationError("turn.settled references no matching member message") watermark = max(watermark, message.seq) watermarks[key] = max(watermarks.get(key, 0), watermark) @@ -616,28 +580,23 @@ def _truncate_utf8_text(value: Any, *, max_bytes: int, suffix: str = "") -> str: def _build_prompt( *, room: DiscussionRoom, member: DiscussionMember, messages: Sequence[_ValidatedEvent], - watermark: int, seen_through_seq: int, -) -> str: + watermark: int, seen_through_seq: int) -> str: delta = [event for event in messages if watermark < event.seq <= seen_through_seq][ - -MAX_DISCUSSION_DELTA_LINES: - ] + -MAX_DISCUSSION_DELTA_LINES:] peers = ", ".join( - f"@{candidate.handle}" for candidate in room.members if candidate.member_id != member.member_id - ) + f"@{candidate.handle}" for candidate in room.members if candidate.member_id != member.member_id) opening = [ f'[Discussion: "{room.name}"] You are @{member.handle}, one participant ' f"with {peers or 'no other members'} and the user.", "", - "New messages in this thread since your last turn (oldest first):", - ] + "New messages in this thread since your last turn (oldest first):"] rules = [ "", "Rules for this Discussion:", "- Reply with one conversational message only when you have something new worth adding.", '- If you have nothing new to add, reply with exactly "(pass)".', "- Mention a teammate by handle to pull them into the next round; do not repeat points already made.", - "- Never reveal content from private conversations. Your reply is published verbatim.", - ] + "- Never reveal content from private conversations. Your reply is published verbatim."] fixed_bytes = len("\n".join([*opening, *rules]).encode("utf-8")) available = max(0, driver.MAX_PROMPT_BYTES - fixed_bytes - 1) selected: list[str] = [] @@ -664,12 +623,10 @@ def _build_prompt( def _make_task_plan( *, room: DiscussionRoom, discussion_event: _ValidatedEvent, member: DiscussionMember, - member_index: int, round_index: int, seen_through_seq: int, prompt: str, -) -> DiscussionTaskPlan: + member_index: int, round_index: int, seen_through_seq: int, prompt: str) -> DiscussionTaskPlan: turn_id = ( f"d{discussion_event.seq}.r{round_index}.p{member_index}." - f"s{seen_through_seq}.m{_member_digest(member)}" - ) + f"s{seen_through_seq}.m{_member_digest(member)}") seed = compact_json({ "discussion_event_id": discussion_event.event_id, "member_id": member.member_id, @@ -679,36 +636,29 @@ def _make_task_plan( "round_index": round_index, "seen_through_seq": seen_through_seq, "source_event_seq": discussion_event.seq, - "thread_id": discussion_event.payload["thread_id"], - }) + "thread_id": discussion_event.payload["thread_id"]}) task_id = f"dtask:{hashlib.sha256(seed.encode('utf-8')).hexdigest()[:48]}" identity = driver.TaskIdentity( room_id=room.room_id, task_id=task_id, - thread_id=str(discussion_event.payload["thread_id"]), turn_id=turn_id, - ) + thread_id=str(discussion_event.payload["thread_id"]), turn_id=turn_id) payload = { "target_member_id": member.member_id, "target_profile": member.profile, "prompt": prompt, - "source_event_seq": discussion_event.seq, - } + "source_event_seq": discussion_event.seq} return DiscussionTaskPlan( - identity=identity, payload=payload, discussion_event_id=discussion_event.event_id, - member=member, member_index=member_index, round_index=round_index, - seen_through_seq=seen_through_seq, - ) + identity=identity, payload=payload, discussion_event_id=discussion_event.event_id, member=member, + member_index=member_index, round_index=round_index, seen_through_seq=seen_through_seq) def _pending_discussion(validated: Sequence[_ValidatedEvent]) -> _ValidatedEvent | None: """Oldest latest-per-thread user message not stopped and not yet completed.""" stopped_through_seq = max( - (event.seq for event in validated if event.kind == "room.stop_requested"), default=0 - ) + (event.seq for event in validated if event.kind == "room.stop_requested"), default=0) completed_discussion_ids = { str(event.payload["discussion_event_id"]) for event in validated - if event.kind == "room.activity" and event.payload.get("status") in {"settled", "bounded"} - } + if event.kind == "room.activity" and event.payload.get("status") in {"settled", "bounded"}} latest_by_thread: dict[str, _ValidatedEvent] = {} for event in validated: if event.kind == "message.user": @@ -716,8 +666,7 @@ def _pending_discussion(validated: Sequence[_ValidatedEvent]) -> _ValidatedEvent pending = [ event for event in sorted(latest_by_thread.values(), key=lambda item: item.seq) - if event.seq > stopped_through_seq and event.event_id not in completed_discussion_ids - ] + if event.seq > stopped_through_seq and event.event_id not in completed_discussion_ids] return pending[0] if pending else None @@ -729,8 +678,7 @@ def _thread_messages( committed_member_message_ids = { str(event.payload["message_event_id"]) for event in validated - if event.kind == "turn.settled" and event.payload.get("message_event_id") is not None - } + if event.kind == "turn.settled" and event.payload.get("message_event_id") is not None} # Publication writes the visible member message before the terminal event. # A crash in that gap leaves the message in the log, but it is not committed # policy input yet: ignoring it reproduces the original task coordinates so @@ -741,15 +689,12 @@ def _thread_messages( if event.payload.get("thread_id") == thread_id and ( event.kind == "message.user" - or (event.kind == "message.member" and event.event_id in committed_member_message_ids) - ) - ) + or (event.kind == "message.member" and event.event_id in committed_member_message_ids))) discussion_messages = tuple(event for event in thread_messages if event.seq >= discussion.seq) member_messages = tuple( event for event in thread_messages if event.kind == "message.member" - and event.payload.get("discussion_event_id") == discussion.event_id - ) + and event.payload.get("discussion_event_id") == discussion.event_id) return thread_messages, discussion_messages, member_messages @@ -759,20 +704,15 @@ def _effective_watermarks( watermarks = { (str(thread_id), str(member_id)): int(value) for (thread_id, member_id), value in (initial_watermarks or {}).items() - if int(value) >= 0 - } + if int(value) >= 0} for key, value in _derive_member_watermarks(validated).items(): watermarks[key] = max(watermarks.get(key, 0), value) return watermarks def plan_next_task( - room_value: Any, - events: Sequence[Mapping[str, Any]], - *, - local_profiles: Iterable[str], - initial_watermarks: Mapping[tuple[str, str], int] | None = None, -) -> DiscussionDecision: + room_value: Any, events: Sequence[Mapping[str, Any]], *, local_profiles: Iterable[str], + initial_watermarks: Mapping[tuple[str, str], int] | None = None) -> DiscussionDecision: """Replay the complete room log and return at most one next member task.""" room = validate_room(room_value, local_profiles=local_profiles) validated = _validated_events(events, room=room) @@ -782,8 +722,7 @@ def plan_next_task( thread_id = str(discussion.payload["thread_id"]) decide = partial( DiscussionDecision, discussion_event_id=discussion.event_id, - source_event_seq=discussion.seq, thread_id=thread_id, - ) + source_event_seq=discussion.seq, thread_id=thread_id) thread_messages, discussion_messages, member_messages = _thread_messages(validated, discussion) if len(member_messages) >= MAX_DISCUSSION_MESSAGES: @@ -793,8 +732,7 @@ def plan_next_task( (int(event.payload["round_index"]), str(event.payload["member_id"])): event for event in validated if event.kind in _TERMINAL_EVENT_KINDS - and event.payload.get("discussion_event_id") == discussion.event_id - } + and event.payload.get("discussion_event_id") == discussion.event_id} watermarks = _effective_watermarks(validated, initial_watermarks) seen_through_seq = max(event.seq for event in thread_messages) @@ -807,8 +745,7 @@ def plan_next_task( responders = ( resolve_mentions((str(discussion.payload["text"]),), room.members) if round_index == 0 - else _unaddressed_member_mentions(discussion_messages, room) - ) + else _unaddressed_member_mentions(discussion_messages, room)) for member_index, member in enumerate(_rotate(responders, round_index)): if (round_index, member.member_id) in terminals: continue @@ -817,12 +754,10 @@ def plan_next_task( continue prompt = _build_prompt( room=room, member=member, messages=thread_messages, - watermark=watermark, seen_through_seq=seen_through_seq, - ) + watermark=watermark, seen_through_seq=seen_through_seq) task = _make_task_plan( room=room, discussion_event=discussion, member=member, member_index=member_index, - round_index=round_index, seen_through_seq=seen_through_seq, prompt=prompt, - ) + round_index=round_index, seen_through_seq=seen_through_seq, prompt=prompt) return decide("task", "member_turn", task=task) if not any(int(event.payload["round_index"]) == round_index for event in member_messages): @@ -835,8 +770,7 @@ def plan_next_task( def reconstruct_task_plan( room_value: Any, events: Sequence[Mapping[str, Any]], task: Mapping[str, Any], - *, local_profiles: Iterable[str], -) -> DiscussionTaskPlan: + *, local_profiles: Iterable[str]) -> DiscussionTaskPlan: """Reconstruct and verify one persisted driver task after a restart.""" room = validate_room(room_value, local_profiles=local_profiles) validated = _validated_events(events, room=room) @@ -846,8 +780,7 @@ def reconstruct_task_plan( raise DiscussionReconstructionError("driver task has no valid identity or payload") required_payload = frozenset({"target_profile", "prompt", "source_event_seq"}) if not required_payload <= frozenset(payload) or ( - frozenset(payload) - required_payload - {"target_member_id"} - ): + frozenset(payload) - required_payload - {"target_member_id"}): raise DiscussionReconstructionError("driver task payload shape changed") match = _TURN_ID_RE.fullmatch(identity.turn_id) if match is None: @@ -855,9 +788,7 @@ def reconstruct_task_plan( source_event_seq = int(match.group("source")) if payload.get("source_event_seq") != source_event_seq: raise DiscussionReconstructionError("task source event does not match turn_id") - discussion = next( - (e for e in validated if e.seq == source_event_seq and e.kind == "message.user"), None - ) + discussion = next((e for e in validated if e.seq == source_event_seq and e.kind == "message.user"), None) if discussion is None: raise DiscussionReconstructionError("task source user event is missing") if identity.room_id != room.room_id or identity.thread_id != discussion.payload["thread_id"]: @@ -870,11 +801,8 @@ def reconstruct_task_plan( if ( candidate.member_id == target_member_id if target_member_id is not None - else candidate.profile == profile - ) - ), - None, - ) + else candidate.profile == profile)), + None) if member is not None and member.profile != profile: member = None if member is None or _member_digest(member) != match.group("member"): @@ -885,10 +813,8 @@ def reconstruct_task_plan( if len(prompt.encode("utf-8")) > driver.MAX_PROMPT_BYTES: raise DiscussionReconstructionError("task prompt exceeds the driver limit") reconstructed = _make_task_plan( - room=room, discussion_event=discussion, member=member, - member_index=int(match.group("position")), round_index=int(match.group("round")), - seen_through_seq=int(match.group("seen")), prompt=prompt, - ) + room=room, discussion_event=discussion, member=member, member_index=int(match.group("position")), + round_index=int(match.group("round")), seen_through_seq=int(match.group("seen")), prompt=prompt) if reconstructed.identity != identity or dict(reconstructed.payload) != dict(payload): raise DiscussionReconstructionError("driver task failed deterministic reconstruction") return reconstructed @@ -914,8 +840,7 @@ def _settled_effects( ) -> tuple[dict[str, Any], list[EventPlan]]: text = _truncate_utf8_text( _terminal_text(result, field="text", fallback=""), - max_bytes=MAX_MEMBER_TEXT_BYTES, suffix=_TRUNCATED_REPLY_NOTICE, - ) + max_bytes=MAX_MEMBER_TEXT_BYTES, suffix=_TRUNCATED_REPLY_NOTICE) passed = is_pass_text(text) effects: list[EventPlan] = [] if not passed: @@ -935,10 +860,8 @@ def _settled_effects( "task_id": task.identity.task_id, "text": text, "thread_id": task.identity.thread_id, - "turn_id": task.identity.turn_id, - }, - authority_gateway_id=room.gateway_id, authority_epoch=room.authority_epoch, - )) + "turn_id": task.identity.turn_id}, + authority_gateway_id=room.gateway_id, authority_epoch=room.authority_epoch)) return {"message_event_id": None if passed else message_event_id, "passed": passed}, effects @@ -948,25 +871,21 @@ def _failed_effects(result: Any, **_: Any) -> tuple[dict[str, Any], list[EventPl supplied_reason = ( str(result.get("reason_code") or result.get("reason") or "").strip() - if isinstance(result, Mapping) else "" - ) + if isinstance(result, Mapping) else "") reason_code = supplied_reason if supplied_reason in ALL_REASONS else classify_agent_error(error_text) return {"error": error_text, "reason_code": reason_code}, [] def _cancelled_effects( - result: Any, *, newer_same_thread: bool, **_: Any -) -> tuple[dict[str, Any], list[EventPlan]]: + result: Any, *, newer_same_thread: bool, **_: Any) -> tuple[dict[str, Any], list[EventPlan]]: reason = ( "superseded_by_newer_user_event" if newer_same_thread - else _terminal_text(result, field="reason", fallback="member turn cancelled") - ) + else _terminal_text(result, field="reason", fallback="member turn cancelled")) return {"reason": reason}, [] def _deferred_effects( - result: Any, *, execution_generation: int | None, **_: Any -) -> tuple[dict[str, Any], list[EventPlan]]: + result: Any, *, execution_generation: int | None, **_: Any) -> tuple[dict[str, Any], list[EventPlan]]: return { "execution_generation": execution_generation, "reason": _terminal_text(result, field="reason", fallback="member_unavailable"), @@ -977,19 +896,12 @@ _TERMINAL_EFFECTS = { "settled": _settled_effects, "failed": _failed_effects, "cancelled": _cancelled_effects, - "deferred": _deferred_effects, -} + "deferred": _deferred_effects} def plan_publication( - room_value: Any, - events: Sequence[Mapping[str, Any]], - task: DiscussionTaskPlan, - *, - status: TerminalKind, - result: Any = None, - execution_generation: int | None = None, - local_profiles: Iterable[str], + room_value: Any, events: Sequence[Mapping[str, Any]], task: DiscussionTaskPlan, *, status: TerminalKind, + result: Any = None, execution_generation: int | None = None, local_profiles: Iterable[str], ) -> PublicationPlan: """Plan idempotent room effects for one terminal driver task. @@ -1013,16 +925,12 @@ def plan_publication( event.kind == "message.user" and event.seq > task.seen_through_seq and event.payload.get("thread_id") == task.identity.thread_id - for event in validated - ) - effective_status: TerminalKind = ( - "cancelled" if newer_same_thread and status != "deferred" else status - ) + for event in validated) + effective_status: TerminalKind = ("cancelled" if newer_same_thread and status != "deferred" else status) digest = task.identity.task_id.removeprefix("dtask:") terminal_event_id = ( f"ddeferred:{digest}:g{execution_generation}" - if effective_status == "deferred" else f"dterminal:{digest}" - ) + if effective_status == "deferred" else f"dterminal:{digest}") common_payload = { "discussion_event_id": task.discussion_event_id, "member_id": task.member.member_id, @@ -1031,19 +939,14 @@ def plan_publication( "seen_through_seq": task.seen_through_seq, "task_id": task.identity.task_id, "thread_id": task.identity.thread_id, - "turn_id": task.identity.turn_id, - } + "turn_id": task.identity.turn_id} extra, effects = _TERMINAL_EFFECTS[effective_status]( result, task=task, room=room, message_event_id=f"dmessage:{digest}", - newer_same_thread=newer_same_thread, execution_generation=execution_generation, - ) + newer_same_thread=newer_same_thread, execution_generation=execution_generation) terminal_kind = f"turn.{effective_status}" effects.append(EventPlan( event_id=terminal_event_id, kind=terminal_kind, actor={"kind": "gateway", "id": room.gateway_id}, payload={**common_payload, **extra}, - authority_gateway_id=room.gateway_id, authority_epoch=room.authority_epoch, - )) - return PublicationPlan( - task_id=task.identity.task_id, terminal_kind=terminal_kind, events=tuple(effects) - ) + authority_gateway_id=room.gateway_id, authority_epoch=room.authority_epoch)) + return PublicationPlan(task_id=task.identity.task_id, terminal_kind=terminal_kind, events=tuple(effects)) diff --git a/gateway/hosted_room_driver.py b/gateway/hosted_room_driver.py index d2ef4e9206..52105d288b 100644 --- a/gateway/hosted_room_driver.py +++ b/gateway/hosted_room_driver.py @@ -19,8 +19,7 @@ from pathlib import Path from typing import Any, Callable, Literal, get_args from gateway.hosted_rooms_common import ( - bounded_int, canonical_json, compact_json, connect, identifier, table_columns, text, transaction, -) + bounded_int, canonical_json, compact_json, connect, identifier, table_columns, text, transaction) Clock = Callable[[], float] TaskStatus = Literal["queued", "running", "settled", "failed", "cancelled", "indeterminate", "deferred", "stopping"] @@ -39,16 +38,12 @@ _TASK_PAYLOAD_REQUIRED_FIELDS = frozenset({"target_profile", "prompt", "source_e _TASK_PAYLOAD_OPTIONAL_FIELDS = frozenset({"target_member_id"}) _LEASE_COLUMNS = frozenset({ "room_id", "gateway_id", "authority_epoch", "process_generation", - "lease_generation", "expires_at", "acquired_at", "updated_at", "released_at", -}) + "lease_generation", "expires_at", "acquired_at", "updated_at", "released_at"}) _TASK_COLUMN_ORDER = ( - "room_id", "task_id", "thread_id", "turn_id", "source_event_seq", - "payload_json", "payload_digest", "status", "execution_generation", - "cancel_generation", "run_gateway_id", "run_process_generation", - "run_lease_generation", "cancel_id", "settlement_id", "settlement_status", - "result_json", "created_at", "updated_at", "started_at", "terminal_at", - "indeterminate_at", -) + "room_id", "task_id", "thread_id", "turn_id", "source_event_seq", "payload_json", "payload_digest", + "status", "execution_generation", "cancel_generation", "run_gateway_id", "run_process_generation", + "run_lease_generation", "cancel_id", "settlement_id", "settlement_status", "result_json", "created_at", + "updated_at", "started_at", "terminal_at", "indeterminate_at") _TASK_COLUMNS = frozenset(_TASK_COLUMN_ORDER) _TASK_ORDER = "ORDER BY source_event_seq, created_at, task_id" _SELECT_LEASE = "SELECT * FROM hosted_room_driver_leases WHERE room_id=?" @@ -79,17 +74,14 @@ _SETTLE_RUNNING_SQL = _generation_update(_SETTLE_SET, "running") + f" AND {_RUN_ _SETTLE_STOPPING_SQL = _generation_update(_SETTLE_SET, "stopping") _REQUEUE_RUNNING_SQL = _task_update( f"{_REQUEUE_SET}, started_at=NULL, updated_at=?", - f"status='running' AND {_GENERATION_FENCE} AND {_RUN_FENCE}", -) + f"status='running' AND {_GENERATION_FENCE} AND {_RUN_FENCE}") _CANCEL_QUEUED_SQL = _task_update(_CANCEL_SET, "status IN ('queued', 'deferred') AND cancel_generation=?") _BEGIN_STOP_SQL = _task_update( "status='stopping', cancel_generation=?, cancel_id=?, updated_at=?", - "status IN ('running', 'indeterminate') AND cancel_generation=?", -) + "status IN ('running', 'indeterminate') AND cancel_generation=?") _COMPLETE_STOP_SQL = _task_update( "status='cancelled', terminal_at=?, updated_at=?", - "status='stopping' AND cancel_id=? AND cancel_generation=?", -) + "status='stopping' AND cancel_id=? AND cancel_generation=?") # Lease-first recovery transitions: name -> (fenced status, SET clause, generation-guard stale message, # row stale message); the UPDATE is _generation_update(set_clause, status). @@ -100,22 +92,17 @@ _GENERATION_TRANSITIONS = { ), "resolve_cancel": ( "indeterminate", _CANCEL_SET, "indeterminate cancellation proof is stale", - "indeterminate cancellation proof lost its fence", - ), + "indeterminate cancellation proof lost its fence"), "requeue": ( "indeterminate", f"{_REQUEUE_SET}, started_at=NULL, indeterminate_at=NULL, updated_at=?", - _INDETERMINATE_STALE, "indeterminate task changed during requeue", - ), + _INDETERMINATE_STALE, "indeterminate task changed during requeue"), "defer": ( "indeterminate", "status='deferred', result_json=?, terminal_at=?, updated_at=?", - _INDETERMINATE_STALE, "indeterminate task changed during deferral", - ), + _INDETERMINATE_STALE, "indeterminate task changed during deferral"), "requeue_deferred": ( "deferred", f"{_REQUEUE_SET}, result_json=NULL, started_at=NULL, terminal_at=NULL, indeterminate_at=NULL, updated_at=?", - "deferred task generation changed", "deferred task changed during requeue", - ), -} + "deferred task generation changed", "deferred task changed during requeue")} class DriverStateError(ValueError): @@ -154,8 +141,7 @@ _identifier = partial(identifier, error=DriverValidationError, max_chars=MAX_IDE _bounded_int = partial(bounded_int, error=DriverValidationError) _canonical_json = partial( canonical_json, error=DriverValidationError, label="result", max_bytes=MAX_RESULT_JSON_BYTES, - ensure_ascii=True, -) + ensure_ascii=True) def _finite(compute: Callable[[], Any], message: str, *, positive: bool = False) -> float: @@ -196,12 +182,9 @@ def _task_payload(value: Any) -> tuple[dict[str, Any], str, str]: raise DriverValidationError(f"missing payload fields: {', '.join(sorted(missing))}") target_profile = _identifier(value["target_profile"], label="target_profile") prompt = text( - value["prompt"], error=DriverValidationError, label="prompt", max_bytes=MAX_PROMPT_BYTES, - strip=False, - ) + value["prompt"], error=DriverValidationError, label="prompt", max_bytes=MAX_PROMPT_BYTES, strip=False) source_event_seq = _bounded_int( - value["source_event_seq"], message="source_event_seq must be a positive integer", low=1 - ) + value["source_event_seq"], message="source_event_seq must be a positive integer", low=1) normalized = {"target_profile": target_profile, "prompt": prompt, "source_event_seq": source_event_seq} if "target_member_id" in value: normalized["target_member_id"] = _identifier(value["target_member_id"], label="target_member_id") @@ -262,8 +245,7 @@ def _create_task_table(conn: sqlite3.Connection, table: str = "hosted_room_drive settlement_id TEXT, settlement_status TEXT, result_json TEXT, created_at REAL NOT NULL, updated_at REAL NOT NULL, started_at REAL, terminal_at REAL, indeterminate_at REAL, PRIMARY KEY (room_id, task_id), UNIQUE (room_id, thread_id, turn_id), - FOREIGN KEY (room_id) REFERENCES hosted_rooms(room_id))""" - ) + FOREIGN KEY (room_id) REFERENCES hosted_rooms(room_id))""") def _initialize_schema(conn: sqlite3.Connection) -> None: @@ -286,8 +268,7 @@ def _validate_schema(conn: sqlite3.Connection) -> None: if lease_columns != _LEASE_COLUMNS or task_columns != _TASK_COLUMNS: raise DriverStateError( "unsupported unpublished hosted-room driver schema; " - "recreate the driver tables before starting the driver" - ) + "recreate the driver tables before starting the driver") for table in ("hosted_room_driver_leases", "hosted_room_driver_tasks"): foreign_keys = conn.execute(f"PRAGMA foreign_key_list({table})").fetchall() if not any(row[2] == "hosted_rooms" and row[3] == "room_id" and row[4] == "room_id" for row in foreign_keys): @@ -309,8 +290,7 @@ def _schema_objects_exist(conn: sqlite3.Connection) -> bool: def _task_schema_supports_current_statuses(conn: sqlite3.Connection) -> bool: row = conn.execute( - "SELECT sql FROM sqlite_master WHERE type='table' AND name='hosted_room_driver_tasks'" - ).fetchone() + "SELECT sql FROM sqlite_master WHERE type='table' AND name='hosted_room_driver_tasks'").fetchone() sql = str(row[0] or "").lower() if row else "" return "'stopping'" in sql and "'deferred'" in sql @@ -359,17 +339,14 @@ def _transaction(db_path: Path | str): def _lease_from_row(row: sqlite3.Row | dict[str, Any], *, reclaimed: bool = False) -> DriverLease: return DriverLease( - room_id=row["room_id"], gateway_id=row["gateway_id"], - authority_epoch=int(row["authority_epoch"]), process_generation=row["process_generation"], - lease_generation=int(row["lease_generation"]), expires_at=float(row["expires_at"]), - reclaimed=reclaimed, - ) + room_id=row["room_id"], gateway_id=row["gateway_id"], authority_epoch=int(row["authority_epoch"]), + process_generation=row["process_generation"], lease_generation=int(row["lease_generation"]), + expires_at=float(row["expires_at"]), reclaimed=reclaimed) def _task_identity_from_row(row: sqlite3.Row) -> TaskIdentity: return TaskIdentity( - room_id=row["room_id"], task_id=row["task_id"], thread_id=row["thread_id"], turn_id=row["turn_id"] - ) + room_id=row["room_id"], task_id=row["task_id"], thread_id=row["thread_id"], turn_id=row["turn_id"]) def _optional(cast: Callable[[Any], Any]) -> Callable[[Any], Any]: @@ -381,8 +358,7 @@ def _optional(cast: Callable[[Any], Any]) -> Callable[[Any], Any]: _TASK_VIEW_CASTS: dict[str, Callable[[Any], Any]] = { "execution_generation": int, "cancel_generation": int, "run_lease_generation": _optional(int), "result_json": _optional(json.loads), "created_at": float, "updated_at": float, - "started_at": _optional(float), "terminal_at": _optional(float), "indeterminate_at": _optional(float), -} + "started_at": _optional(float), "terminal_at": _optional(float), "indeterminate_at": _optional(float)} def _task_from_row(row: sqlite3.Row, *, idempotent: bool = False) -> dict[str, Any]: @@ -391,8 +367,7 @@ def _task_from_row(row: sqlite3.Row, *, idempotent: bool = False) -> dict[str, A except (TypeError, json.JSONDecodeError, DriverValidationError) as exc: raise TaskConflictError("stored task payload is invalid") from exc if (encoded_payload, payload_digest, payload["source_event_seq"]) != ( - row["payload_json"], row["payload_digest"], int(row["source_event_seq"]) - ): + row["payload_json"], row["payload_digest"], int(row["source_event_seq"])): raise TaskConflictError("stored task payload failed its integrity check") task: dict[str, Any] = {"identity": _task_identity_from_row(row), "payload": payload} for column in _TASK_COLUMN_ORDER[_TASK_COLUMN_ORDER.index("payload_digest"):]: @@ -420,8 +395,7 @@ def _load_active_room(conn: sqlite3.Connection, room_id: str) -> sqlite3.Row: try: row = conn.execute( "SELECT room_id, authority_gateway_id, authority_epoch, disbanded_at FROM hosted_rooms WHERE room_id=?", - (room_id,), - ).fetchone() + (room_id,)).fetchone() except sqlite3.OperationalError as exc: if "no such table" in str(exc).lower(): raise RoomUnavailableError("hosted room does not exist") from exc @@ -434,8 +408,7 @@ def _load_active_room(conn: sqlite3.Connection, room_id: str) -> sqlite3.Row: def _require_room_authority( - conn: sqlite3.Connection, *, room_id: str, gateway_id: str, authority_epoch: int -) -> sqlite3.Row: + conn: sqlite3.Connection, *, room_id: str, gateway_id: str, authority_epoch: int) -> sqlite3.Row: room = _load_active_room(conn, room_id) if room["authority_gateway_id"] != gateway_id or int(room["authority_epoch"]) != authority_epoch: raise StaleLeaseError("hosted room authority changed") @@ -444,8 +417,7 @@ def _require_room_authority( def _require_lease_authority(conn: sqlite3.Connection, lease: DriverLease) -> sqlite3.Row: return _require_room_authority( - conn, room_id=lease.room_id, gateway_id=lease.gateway_id, authority_epoch=lease.authority_epoch - ) + conn, room_id=lease.room_id, gateway_id=lease.gateway_id, authority_epoch=lease.authority_epoch) def _lease_row_matches(row: sqlite3.Row | None, lease: DriverLease) -> bool: @@ -475,8 +447,7 @@ def _cancel_generation(value: int) -> int: def _expected_generations( - lease: DriverLease, identity: TaskIdentity, execution_generation: int, cancel_generation: int -) -> None: + lease: DriverLease, identity: TaskIdentity, execution_generation: int, cancel_generation: int) -> None: _check_same_room(lease, identity) if not isinstance(execution_generation, int) or execution_generation < 1: raise DriverValidationError("expected_execution_generation must be a positive integer") @@ -511,8 +482,7 @@ def _cancel_replay(cancel_id: str, status: str = "cancelled") -> Callable[[sqlit def _generations_match(row: sqlite3.Row, status: str, execution_generation: int, cancel_generation: int) -> bool: return (row["status"], int(row["execution_generation"]), int(row["cancel_generation"])) == ( - status, execution_generation, cancel_generation - ) + status, execution_generation, cancel_generation) def _require_cancel_generation(row: sqlite3.Row, expected_cancel_generation: int) -> None: @@ -524,8 +494,7 @@ def _transition( db_path: Path | str, identity: TaskIdentity, *, sql: str, set_params: tuple[Any, ...], fence_params: tuple[Any, ...], stale: str, now: float, lease: DriverLease | None = None, lease_first: bool = True, replay: Callable[[sqlite3.Row], dict[str, Any] | None] | None = None, - guard: Callable[[sqlite3.Row], None] | None = None, -) -> dict[str, Any]: + guard: Callable[[sqlite3.Row], None] | None = None) -> dict[str, Any]: """Run one fenced task transition: load -> idempotent replay -> lease/fence guard -> UPDATE. ``sql`` binds ``(*set_params, room_id, task_id, *fence_params)`` and must hit exactly one row @@ -554,8 +523,7 @@ def _transition( def _generation_transition( db_path: Path | str, identity: TaskIdentity, lease: DriverLease, name: str, execution_generation: int, cancel_generation: int, *, now: float, set_params: tuple[Any, ...], - replay: Callable[[sqlite3.Row], Any] | None = None, -) -> dict[str, Any]: + replay: Callable[[sqlite3.Row], Any] | None = None) -> dict[str, Any]: """Lease-first transition from ``_GENERATION_TRANSITIONS`` fenced on status + both generations.""" status, set_clause, generation_stale, stale = _GENERATION_TRANSITIONS[name] @@ -566,14 +534,12 @@ def _generation_transition( return _transition( db_path, identity, lease=lease, now=now, replay=replay, guard=guard, sql=_generation_update(set_clause, status), set_params=set_params, - fence_params=(execution_generation, cancel_generation), stale=stale, - ) + fence_params=(execution_generation, cancel_generation), stale=stale) def _run_fence_transition( db_path: Path | str, attempt: TaskAttempt, *, guard_stale: str, - lease_generation: Callable[[Any], int] = int, **transition: Any, -) -> dict[str, Any]: + lease_generation: Callable[[Any], int] = int, **transition: Any) -> dict[str, Any]: """Transition fenced on this attempt's running generation under its exact lease (row guard + SQL fence). ``lease_generation`` casts the stored run_lease_generation: ``int`` raises on NULL, @@ -586,30 +552,25 @@ def _run_fence_transition( _generations_match(row, "running", attempt.execution_generation, attempt.cancel_generation) and row["run_gateway_id"] == lease.gateway_id and row["run_process_generation"] == lease.process_generation - and lease_generation(row["run_lease_generation"]) == lease.lease_generation - ): + and lease_generation(row["run_lease_generation"]) == lease.lease_generation): raise StaleTaskError(guard_stale) return _transition( db_path, attempt.identity, lease=lease, guard=guard, fence_params=( attempt.execution_generation, attempt.cancel_generation, - lease.gateway_id, lease.process_generation, lease.lease_generation, - ), - **transition, - ) + lease.gateway_id, lease.process_generation, lease.lease_generation), + **transition) def acquire_lease( db_path: Path | str, *, room_id: Any, gateway_id: Any, authority_epoch: Any, process_generation: Any, - ttl_seconds: Any, clock: Clock, -) -> DriverLease: + ttl_seconds: Any, clock: Clock) -> DriverLease: """Acquire an empty or expired room lease with a monotonic generation.""" room_id = _identifier(room_id, label="room_id") gateway_id = _identifier(gateway_id, label="gateway_id") authority_epoch = _bounded_int( - authority_epoch, message="authority_epoch must be a positive integer", low=1 - ) + authority_epoch, message="authority_epoch must be a positive integer", low=1) process_generation = _identifier(process_generation, label="process_generation") ttl_seconds = _ttl(ttl_seconds) now = _timestamp(clock) @@ -624,8 +585,7 @@ def acquire_lease( room_id, gateway_id, authority_epoch, process_generation, lease_generation, expires_at, acquired_at, updated_at, released_at ) VALUES (?, ?, ?, ?, 1, ?, ?, ?, NULL)""", - (room_id, gateway_id, authority_epoch, process_generation, expires_at, now, now), - ) + (room_id, gateway_id, authority_epoch, process_generation, expires_at, now, now)) return _lease_from_row(conn.execute(_SELECT_LEASE, (room_id,)).fetchone()) same_authority = row["gateway_id"] == gateway_id and int(row["authority_epoch"]) == authority_epoch @@ -634,8 +594,7 @@ def acquire_lease( renewed_expiry = max(float(row["expires_at"]), expires_at) conn.execute( "UPDATE hosted_room_driver_leases SET expires_at=?, updated_at=? WHERE room_id=? AND lease_generation=?", - (renewed_expiry, now, room_id, int(row["lease_generation"])), - ) + (renewed_expiry, now, room_id, int(row["lease_generation"]))) return _lease_from_row({**dict(row), "expires_at": renewed_expiry}) if same_authority and live: raise LeaseHeldError("room driver lease is held by another generation") @@ -648,9 +607,7 @@ def acquire_lease( gateway_id != ? OR authority_epoch != ? OR released_at IS NOT NULL OR expires_at <= ?)""", ( gateway_id, authority_epoch, process_generation, expires_at, now, now, - room_id, int(row["lease_generation"]), gateway_id, authority_epoch, now, - ), - ) + room_id, int(row["lease_generation"]), gateway_id, authority_epoch, now)) if updated.rowcount != 1: raise LeaseHeldError("room driver lease changed during acquisition") return _lease_from_row(conn.execute(_SELECT_LEASE, (room_id,)).fetchone(), reclaimed=True) @@ -695,8 +652,7 @@ def release_lease(db_path: Path | str, lease: DriverLease, *, clock: Clock) -> d conn.execute( """UPDATE hosted_room_driver_leases SET expires_at=?, updated_at=?, released_at=? WHERE room_id=? AND lease_generation=?""", - (now, now, now, lease.room_id, lease.lease_generation), - ) + (now, now, now, lease.room_id, lease.lease_generation)) current = {**dict(row), "expires_at": now, "updated_at": now, "released_at": now} return {"lease": _lease_from_row(current), "idempotent": False} @@ -716,8 +672,7 @@ def admit_task(db_path: Path | str, identity: TaskIdentity, *, payload: Any, clo return _task_from_row(existing, idempotent=True) turn = conn.execute( "SELECT * FROM hosted_room_driver_tasks WHERE room_id=? AND thread_id=? AND turn_id=?", - (identity.room_id, identity.thread_id, identity.turn_id), - ).fetchone() + (identity.room_id, identity.thread_id, identity.turn_id)).fetchone() if turn is not None: raise TaskConflictError("thread_id and turn_id are already bound to a task") conn.execute( @@ -727,9 +682,7 @@ def admit_task(db_path: Path | str, identity: TaskIdentity, *, payload: Any, clo ) VALUES (?, ?, ?, ?, ?, ?, ?, 'queued', 0, 0, ?, ?)""", ( identity.room_id, identity.task_id, identity.thread_id, identity.turn_id, - normalized_payload["source_event_seq"], payload_json, payload_digest, now, now, - ), - ) + normalized_payload["source_event_seq"], payload_json, payload_digest, now, now)) return _task_from_row(_load_task(conn, identity)) @@ -749,14 +702,12 @@ def start_task( unresolved = conn.execute( f"""SELECT task_id, status FROM hosted_room_driver_tasks WHERE room_id=? AND status IN ('running', 'indeterminate', 'stopping') {_TASK_ORDER} LIMIT 1""", - (identity.room_id,), - ).fetchone() + (identity.room_id,)).fetchone() if unresolved is not None: raise InvalidTaskTransitionError("room recovery must resolve the prior task before starting new work") next_queued = conn.execute( f"SELECT task_id FROM hosted_room_driver_tasks WHERE room_id=? AND status='queued' {_TASK_ORDER} LIMIT 1", - (identity.room_id,), - ).fetchone() + (identity.room_id,)).fetchone() if next_queued is None or next_queued["task_id"] != identity.task_id: raise InvalidTaskTransitionError("task is not next in the hosted room event order") execution_generation = int(row["execution_generation"]) + 1 @@ -767,15 +718,12 @@ def start_task( WHERE room_id=? AND task_id=? AND status='queued' AND cancel_generation=?""", ( execution_generation, lease.gateway_id, lease.process_generation, lease.lease_generation, now, now, - identity.room_id, identity.task_id, expected_cancel_generation, - ), - ) + identity.room_id, identity.task_id, expected_cancel_generation)) if updated.rowcount != 1: raise StaleTaskError("task changed during start") return TaskAttempt( identity=identity, lease=lease, - execution_generation=execution_generation, cancel_generation=expected_cancel_generation, - ) + execution_generation=execution_generation, cancel_generation=expected_cancel_generation) def settle_task( @@ -789,8 +737,7 @@ def settle_task( db_path, attempt, guard_stale="task attempt is stale or cancelled", lease_first=False, now=now, replay=_settlement_replay(settlement_id, status, result_json), sql=_SETTLE_RUNNING_SQL, set_params=(status, settlement_id, status, result_json, now, now), - stale="task changed during settlement", - ) + stale="task changed during settlement") def settle_stopping_task( @@ -808,8 +755,7 @@ def settle_stopping_task( replay=_settlement_replay(settlement_id, status, result_json), sql=_SETTLE_STOPPING_SQL, set_params=(status, settlement_id, status, result_json, now, now), fence_params=(expected_execution_generation, expected_cancel_generation), - stale="task completion lost the stop race", - ) + stale="task completion lost the stop race") def resolve_indeterminate_task( @@ -824,14 +770,12 @@ def resolve_indeterminate_task( return _generation_transition( db_path, identity, lease, "resolve", expected_execution_generation, expected_cancel_generation, now=now, replay=_settlement_replay(settlement_id, status, result_json), - set_params=(status, settlement_id, status, result_json, now, now), - ) + set_params=(status, settlement_id, status, result_json, now, now)) def resolve_indeterminate_cancellation( db_path: Path | str, identity: TaskIdentity, lease: DriverLease, *, expected_execution_generation: int, - expected_cancel_generation: int, cancel_id: Any, clock: Clock, -) -> dict[str, Any]: + expected_cancel_generation: int, cancel_id: Any, clock: Clock) -> dict[str, Any]: """Commit a verified terminal cancellation for an uncertain attempt.""" _expected_generations(lease, identity, expected_execution_generation, expected_cancel_generation) cancel_id = _identifier(cancel_id, label="cancel_id") @@ -844,21 +788,18 @@ def resolve_indeterminate_cancellation( def requeue_indeterminate_task( db_path: Path | str, identity: TaskIdentity, lease: DriverLease, *, expected_execution_generation: int, - expected_cancel_generation: int, clock: Clock, -) -> dict[str, Any]: + expected_cancel_generation: int, clock: Clock) -> dict[str, Any]: """Explicitly retry uncertain work after an operator accepts at-least-once risk.""" _expected_generations(lease, identity, expected_execution_generation, expected_cancel_generation) now = _timestamp(clock) return _generation_transition( db_path, identity, lease, "requeue", expected_execution_generation, expected_cancel_generation, now=now, - set_params=(now,), - ) + set_params=(now,)) def defer_indeterminate_task( db_path: Path | str, identity: TaskIdentity, lease: DriverLease, *, expected_execution_generation: int, - expected_cancel_generation: int, reason: Any, clock: Clock, -) -> dict[str, Any]: + expected_cancel_generation: int, reason: Any, clock: Clock) -> dict[str, Any]: """Fence one uncertain attempt and release later room work.""" _expected_generations(lease, identity, expected_execution_generation, expected_cancel_generation) reason = _identifier(reason, label="defer_reason") @@ -871,21 +812,18 @@ def defer_indeterminate_task( return _generation_transition( db_path, identity, lease, "defer", expected_execution_generation, expected_cancel_generation, now=now, - replay=replay, set_params=(result_json, now, now), - ) + replay=replay, set_params=(result_json, now, now)) def requeue_deferred_task( db_path: Path | str, identity: TaskIdentity, lease: DriverLease, *, expected_execution_generation: int, - expected_cancel_generation: int, clock: Clock, -) -> dict[str, Any]: + expected_cancel_generation: int, clock: Clock) -> dict[str, Any]: """Explicitly retry a fenced deferred turn under a new generation.""" _expected_generations(lease, identity, expected_execution_generation, expected_cancel_generation) now = _timestamp(clock) return _generation_transition( db_path, identity, lease, "requeue_deferred", expected_execution_generation, expected_cancel_generation, - now=now, set_params=(now,), - ) + now=now, set_params=(now,)) def requeue_not_admitted_task(db_path: Path | str, attempt: TaskAttempt, *, clock: Clock) -> dict[str, Any]: @@ -898,15 +836,13 @@ def requeue_not_admitted_task(db_path: Path | str, attempt: TaskAttempt, *, cloc _generations_match(row, "queued", attempt.execution_generation, attempt.cancel_generation) and row["run_gateway_id"] is None and row["run_process_generation"] is None - and row["run_lease_generation"] is None - ) + and row["run_lease_generation"] is None) return _task_from_row(row, idempotent=True) if requeued else None return _run_fence_transition( db_path, attempt, guard_stale="not-admitted task attempt lost its fence", lease_generation=lambda value: int(value or 0), now=now, replay=replay, - sql=_REQUEUE_RUNNING_SQL, set_params=(now,), stale="not-admitted task changed during requeue", - ) + sql=_REQUEUE_RUNNING_SQL, set_params=(now,), stale="not-admitted task changed during requeue") def cancel_task( @@ -927,8 +863,7 @@ def cancel_task( return _transition( db_path, identity, now=now, replay=_cancel_replay(cancel_id), guard=guard, sql=_CANCEL_QUEUED_SQL, set_params=(expected_cancel_generation + 1, cancel_id, now, now), fence_params=(expected_cancel_generation,), - stale="task changed during cancellation", - ) + stale="task changed during cancellation") def begin_task_cancel( @@ -947,8 +882,7 @@ def begin_task_cancel( return _transition( db_path, identity, now=now, replay=_cancel_replay(cancel_id, "stopping"), guard=guard, sql=_BEGIN_STOP_SQL, set_params=(expected_cancel_generation + 1, cancel_id, now), fence_params=(expected_cancel_generation,), - stale="task changed during stop request", - ) + stale="task changed during stop request") def complete_task_cancel( @@ -960,15 +894,13 @@ def complete_task_cancel( def guard(row: sqlite3.Row) -> None: if (row["status"], row["cancel_id"], int(row["cancel_generation"])) != ( - "stopping", cancel_id, expected_cancel_generation - ): + "stopping", cancel_id, expected_cancel_generation): raise StaleTaskError("task stop acknowledgement is stale") return _transition( db_path, identity, now=now, replay=_cancel_replay(cancel_id), guard=guard, sql=_COMPLETE_STOP_SQL, set_params=(now, now), fence_params=(cancel_id, expected_cancel_generation), - stale="task changed during stop acknowledgement", - ) + stale="task changed during stop acknowledgement") def recover_room(db_path: Path | str, lease: DriverLease, *, clock: Clock) -> dict[str, list[TaskIdentity]]: @@ -979,18 +911,15 @@ def recover_room(db_path: Path | str, lease: DriverLease, *, clock: Clock) -> di with _transaction(db_path) as conn: _require_active_lease(conn, lease, now=now) stale_rows = conn.execute( - f"SELECT * FROM hosted_room_driver_tasks WHERE {foreign_running} {_TASK_ORDER}", fence - ).fetchall() + f"SELECT * FROM hosted_room_driver_tasks WHERE {foreign_running} {_TASK_ORDER}", fence).fetchall() if stale_rows: conn.execute( f"""UPDATE hosted_room_driver_tasks SET status='indeterminate', indeterminate_at=?, updated_at=? WHERE {foreign_running}""", - (now, now, *fence), - ) + (now, now, *fence)) return { status: [_task_identity_from_row(row) for row in _tasks_in_order(conn, lease.room_id, status)] - for status in ("queued", "indeterminate") - } + for status in ("queued", "indeterminate")} def _read(db_path: Path | str, query: Callable[[sqlite3.Connection], Any]) -> Any: @@ -1038,8 +967,7 @@ def prune_published_terminal_tasks( WHERE p.room_id=t.room_id AND p.task_id=t.task_id AND p.kind IN ('turn.settled', 'turn.failed', 'turn.cancelled')) ORDER BY t.terminal_at DESC, t.task_id ASC""", - (room_id,), - ).fetchall() + (room_id,)).fetchall() cutoff = now - float(retention_seconds) candidates = [ str(row["task_id"]) @@ -1051,6 +979,5 @@ def prune_published_terminal_tasks( placeholders = ",".join("?" for _ in candidates) deleted = conn.execute( f"DELETE FROM hosted_room_driver_tasks WHERE room_id=? AND task_id IN ({placeholders})", - (room_id, *candidates), - ) + (room_id, *candidates)) return max(0, int(deleted.rowcount)) diff --git a/gateway/hosted_room_execution_policy.py b/gateway/hosted_room_execution_policy.py index 41da1e08cf..ce25dad2e5 100644 --- a/gateway/hosted_room_execution_policy.py +++ b/gateway/hosted_room_execution_policy.py @@ -15,8 +15,7 @@ MAX_POLICY_TOOLSETS = 128 MAX_POLICY_ITERATIONS = (1 << 53) - 1 _IDENTIFIER_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:-]*$") _POLICY_FIELDS = { - "version", "target_profile", "enabled_toolsets", "approval_mode", "max_iterations", "policy_digest", -} + "version", "target_profile", "enabled_toolsets", "approval_mode", "max_iterations", "policy_digest"} class RoomExecutionPolicyError(ValueError): @@ -67,16 +66,14 @@ class RoomExecutionPolicy: if ( isinstance(max_iterations, bool) or not isinstance(max_iterations, int) - or not 1 <= max_iterations <= MAX_POLICY_ITERATIONS - ): + or not 1 <= max_iterations <= MAX_POLICY_ITERATIONS): raise RoomExecutionPolicyError("max_iterations is invalid") unsigned = { "version": POLICY_VERSION, "target_profile": target_profile, "enabled_toolsets": list(toolsets), "approval_mode": approval_mode, - "max_iterations": max_iterations, - } + "max_iterations": max_iterations} supplied = str(value["policy_digest"] or "").strip().lower() if supplied != _policy_digest(unsigned): raise RoomExecutionPolicyError("policy_digest does not match the execution policy") @@ -107,10 +104,8 @@ def execution_policy_mapping(*, target_profile: str, config: Mapping[str, Any] | "target_profile": _identifier(target_profile, field="target_profile"), "enabled_toolsets": toolsets, "approval_mode": ( - "off" if _YOLO_MODE_FROZEN else _normalize_approval_mode(approvals.get("mode", "manual")) - ), - "max_iterations": min(resolve_turn_limit(agent.get("max_turns")), MAX_POLICY_ITERATIONS), - } + "off" if _YOLO_MODE_FROZEN else _normalize_approval_mode(approvals.get("mode", "manual"))), + "max_iterations": min(resolve_turn_limit(agent.get("max_turns")), MAX_POLICY_ITERATIONS)} value = {**unsigned, "policy_digest": _policy_digest(unsigned)} return RoomExecutionPolicy.from_mapping(value).as_mapping() @@ -138,5 +133,4 @@ __all__ = [ "bind_room_execution_policy", "current_room_execution_policy", "execution_policy_mapping", - "reset_room_execution_policy", -] + "reset_room_execution_policy"] diff --git a/gateway/hosted_room_links.py b/gateway/hosted_room_links.py index 8132c25cdd..7daf74cf7b 100644 --- a/gateway/hosted_room_links.py +++ b/gateway/hosted_room_links.py @@ -21,8 +21,7 @@ from gateway.hosted_room_peer import ( GatewayRoomCatalog, HostedRoomPeerError, TransportSecurity, - validate_room_link_url, -) + validate_room_link_url) from gateway.hosted_rooms_common import compact_json, exact_fields @@ -30,15 +29,13 @@ MAX_LINKS = 512 MAX_GRANT_CHARS = 16 * 1024 _LEGACY_FIELDS = { "room_id", "member_id", "target_url", "target_profile", "grant", "catalog", - "cancellation_scope_id", "trace_id", "updated_at", -} + "cancellation_scope_id", "trace_id", "updated_at"} _OPTIONAL_FIELDS = {"transport_security", "status"} # SQLite record columns that map 1:1 onto mapping fields, in record order # (``catalog_json`` is the serialized ``catalog``). _RECORD_FIELDS = ( "room_id", "member_id", "target_url", "target_profile", "grant", "cancellation_scope_id", - "trace_id", "transport_security", "status", "updated_at", -) + "trace_id", "transport_security", "status", "updated_at") _STATUSES = {"ready", "unavailable", "needs_reauthorization"} @@ -86,8 +83,7 @@ class StoredRoomLink: trace_id=_short_string(value["trace_id"], "trace_id"), transport_security=transport_security, # type: ignore[arg-type] status=status, - updated_at=updated_at, - ) + updated_at=updated_at) @classmethod def from_record(cls, value: Mapping[str, Any]) -> "StoredRoomLink": @@ -109,8 +105,7 @@ class StoredRoomLink: _link_fields = partial( exact_fields, label="stored room link", required=_LEGACY_FIELDS, optional=_OPTIONAL_FIELDS, error=HostedRoomPeerError, not_object="stored room link fields are invalid", - missing_fmt="stored room link fields are invalid", unknown_fmt="stored room link fields are invalid", -) + missing_fmt="stored room link fields are invalid", unknown_fmt="stored room link fields are invalid") def _short_string(value: Any, field: str) -> str: @@ -157,21 +152,12 @@ def mark_room_link_status(db_path: Path | str, *, room_id: str, member_id: str, raise HostedRoomPeerError("stored room link status is invalid") return hosted_rooms.update_room_link_status( db_path, room_id=_short_string(room_id, "room_id"), member_id=_short_string(member_id, "member_id"), - status=status, - ) + status=status) def make_stored_link( - *, - room_id: str, - member_id: str, - target_url: str, - target_profile: str, - grant: str, - catalog: GatewayRoomCatalog, - cancellation_scope_id: str, - trace_id: str, -) -> StoredRoomLink: + *, room_id: str, member_id: str, target_url: str, target_profile: str, grant: str, + catalog: GatewayRoomCatalog, cancellation_scope_id: str, trace_id: str) -> StoredRoomLink: target_url, transport_security = validate_room_link_url(target_url) return StoredRoomLink.from_mapping({ "room_id": room_id, @@ -184,5 +170,4 @@ def make_stored_link( "trace_id": trace_id, "transport_security": transport_security, "status": "ready", - "updated_at": time.time(), - }) + "updated_at": time.time()}) diff --git a/gateway/hosted_room_peer.py b/gateway/hosted_room_peer.py index 41675e74d5..042640b6fd 100644 --- a/gateway/hosted_room_peer.py +++ b/gateway/hosted_room_peer.py @@ -131,14 +131,12 @@ def derive_room_grant_secret(api_key: str) -> bytes: def _identifier(value: Any, *, field: str) -> str: return identifier( value, label=field, error=HostedRoomPeerError, max_chars=256, pattern=_IDENTIFIER_RE, - invalid=f"{field} is invalid", - ) + invalid=f"{field} is invalid") def _positive_int(value: Any, *, field: str) -> int: return bounded_int( - value, error=HostedRoomPeerError, message=f"{field} must be a positive integer", low=1 - ) + value, error=HostedRoomPeerError, message=f"{field} must be a positive integer", low=1) def _digest(value: Any, *, field: str) -> str: @@ -149,8 +147,7 @@ def _digest(value: Any, *, field: str) -> str: _exact_fields = partial( exact_fields, error=HostedRoomPeerError, missing_fmt="{label} missing fields: {fields}", - unknown_fmt="{label} unknown fields: {fields}", -) + unknown_fmt="{label} unknown fields: {fields}") def _canonical_json(value: Mapping[str, Any]) -> bytes: @@ -199,8 +196,7 @@ def _parse_endpoint(endpoint: Any) -> tuple[str | None, str | None, TransportSec _CATALOG_FIELDS = { "installation_id", "protocol_versions", "link_modes", "persistent_process", "text", - "attachments", "execution_policy", "catalog_digest", -} + "attachments", "execution_policy", "catalog_digest"} @dataclass(frozen=True) @@ -240,9 +236,7 @@ class GatewayRoomCatalog: persistent_process=value["persistent_process"], text=value["text"], attachments=value["attachments"], execution_policy=policy, catalog_digest=_digest(value["catalog_digest"], field="catalog_digest"), - endpoint_url=endpoint_url, endpoint_reason=endpoint_reason, - transport_security=transport_security, - ) + endpoint_url=endpoint_url, endpoint_reason=endpoint_reason, transport_security=transport_security) unsigned = catalog.as_mapping() del unsigned["catalog_digest"] expected = hashlib.sha256(_canonical_json(unsigned)).hexdigest() @@ -256,8 +250,7 @@ class GatewayRoomCatalog: "installation_id": self.installation_id, "protocol_versions": list(self.protocol_versions), "link_modes": list(self.link_modes), "persistent_process": self.persistent_process, "text": self.text, "attachments": self.attachments, - "execution_policy": self.execution_policy.as_mapping(), "catalog_digest": self.catalog_digest, - } + "execution_policy": self.execution_policy.as_mapping(), "catalog_digest": self.catalog_digest} if self.endpoint_url is not None or self.endpoint_reason is not None: value["endpoint"] = self.endpoint_mapping() return value @@ -266,8 +259,7 @@ class GatewayRoomCatalog: """Return the normalized self-advertised endpoint capability.""" if self.endpoint_url is not None: return { - "available": True, "url": self.endpoint_url, "transport_security": self.transport_security, - } + "available": True, "url": self.endpoint_url, "transport_security": self.transport_security} return {"available": False, "reason": self.endpoint_reason or "not_configured"} @@ -275,18 +267,15 @@ def catalog_mapping( *, installation_id: str, protocol_versions: Iterable[int] = (PROTOCOL_VERSION,), link_modes: Iterable[LinkMode] = ("direct", "pull"), persistent_process: bool, text: bool = True, attachments: bool = False, endpoint: Mapping[str, Any] | None = None, target_profile: str | None = None, - execution_policy: Mapping[str, Any] | None = None, -) -> dict[str, Any]: + execution_policy: Mapping[str, Any] | None = None) -> dict[str, Any]: """Build a canonical catalog mapping with its digest.""" # A Desktop-managed gateway exits with the app: the caller's flag is only an # upper bound, so call sites that still pass ``True`` stay honest. persistent_process = bool(persistent_process and os.getenv("HERMES_DESKTOP") != "1") profile = ( - str(target_profile or "").strip() or (os.getenv("HERMES_PROFILE") or "default").strip() or "default" - ) + str(target_profile or "").strip() or (os.getenv("HERMES_PROFILE") or "default").strip() or "default") checked_policy = RoomExecutionPolicy.from_mapping( - execution_policy or execution_policy_mapping(target_profile=profile) - ) + execution_policy or execution_policy_mapping(target_profile=profile)) # A RoomLink run is initiated by another installation. Process-wide YOLO # mode bypasses the scoped approval ContextVar, so rewriting the advertised # policy cannot make it safe: refuse until manual or smart approvals are on. @@ -295,15 +284,13 @@ def catalog_mapping( value = { "installation_id": _identifier(installation_id, field="installation_id"), "protocol_versions": sorted( - {_positive_int(item, field="protocol_version") for item in protocol_versions} - ), + {_positive_int(item, field="protocol_version") for item in protocol_versions}), # Direct HTTPS/loopback is the only RoomLink transport implemented by # this backend slice. Do not advertise pull/relay placeholders. "link_modes": [mode for mode in dict.fromkeys(link_modes) if mode == "direct"], "persistent_process": persistent_process, "text": bool(text), "attachments": bool(attachments), "execution_policy": checked_policy.as_mapping(), - "endpoint": dict(local_room_link_endpoint() if endpoint is None else endpoint), - } + "endpoint": dict(local_room_link_endpoint() if endpoint is None else endpoint)} value["catalog_digest"] = hashlib.sha256(_canonical_json(value)).hexdigest() GatewayRoomCatalog.from_mapping(value) return value @@ -312,14 +299,12 @@ def catalog_mapping( def local_catalog_mapping( *, installation_id: str, protocol_versions: Iterable[int] = (PROTOCOL_VERSION,), link_modes: Iterable[LinkMode] = ("direct", "pull"), text: bool = True, attachments: bool = False, - target_profile: str | None = None, execution_policy: Mapping[str, Any] | None = None, -) -> dict[str, Any]: + target_profile: str | None = None, execution_policy: Mapping[str, Any] | None = None) -> dict[str, Any]: """Build the one truthful catalog advertised by this local process.""" return catalog_mapping( installation_id=installation_id, protocol_versions=protocol_versions, link_modes=link_modes, persistent_process=True, text=text, attachments=attachments, target_profile=target_profile, - execution_policy=execution_policy, - ) + execution_policy=execution_policy) def local_room_link_endpoint(value: Any | None = None) -> dict[str, Any]: @@ -410,8 +395,7 @@ _DISPATCH_FIELDS: dict[str, Callable[..., Any]] = dict( authority_gateway_id=_identifier, authority_epoch=_positive_int, member_id=_identifier, target_install_id=_identifier, target_profile=_identifier, task_id=_identifier, execution_generation=_positive_int, source_event_seq=_positive_int, cancellation_scope_id=_identifier, - capability_digest=_digest, execution_policy_digest=_digest, trace_id=_identifier, -) + capability_digest=_digest, execution_policy_digest=_digest, trace_id=_identifier) @dataclass(frozen=True) @@ -453,19 +437,16 @@ class HostedMemberDispatch: raise HostedRoomPeerError("prompt_digest does not match prompt") return cls( prompt=prompt, prompt_digest=prompt_digest, - **{name: check(value[name], field=name) for name, check in _DISPATCH_FIELDS.items()}, - ) + **{name: check(value[name], field=name) for name, check in _DISPATCH_FIELDS.items()}) # Grant scope fields, in issue-time validation order; each is validated with # the same checker as the matching dispatch field (``grant_id`` is an identifier). _GRANT_SCOPE = ( "grant_id", "room_id", "home_install_id", "authority_gateway_id", "authority_epoch", - "member_id", "target_install_id", "target_profile", -) + "member_id", "target_install_id", "target_profile") _GRANT_FIELDS = frozenset({ - "version", *_GRANT_SCOPE, "execution_policy_digest", "permissions", "issued_at", "expires_at", -}) + "version", *_GRANT_SCOPE, "execution_policy_digest", "permissions", "issued_at", "expires_at"}) _GRANT_REFRESH_FIELDS = _GRANT_FIELDS | {"status_expires_at"} _GRANT_PERMISSIONS = {"approve", "dispatch", "status", "stop"} MAX_DISPATCH_GRANT_TTL_SECONDS = 24 * 60 * 60 @@ -476,9 +457,8 @@ def issue_room_grant( secret: bytes, *, grant_id: str, room_id: str, home_install_id: str, authority_gateway_id: str, authority_epoch: int, member_id: str, target_install_id: str, target_profile: str, execution_policy_digest: str | None = None, - permissions: Iterable[str] = ("approve", "dispatch", "status", "stop"), - issued_at: float | None = None, ttl_seconds: float = 3600, status_ttl_seconds: float | None = None, - status_expires_at: float | None = None, + permissions: Iterable[str] = ("approve", "dispatch", "status", "stop"), issued_at: float | None = None, + ttl_seconds: float = 3600, status_ttl_seconds: float | None = None, status_expires_at: float | None = None, ) -> str: """Issue a target-verifiable bearer grant scoped to one room member.""" if len(secret) < 32: @@ -487,13 +467,11 @@ def issue_room_grant( bounded_status_expiry = ( now + float(ttl_seconds if status_ttl_seconds is None else status_ttl_seconds) if status_expires_at is None - else float(status_expires_at) - ) + else float(status_expires_at)) if ( not math.isfinite(now) or ttl_seconds <= 0 or ttl_seconds > MAX_DISPATCH_GRANT_TTL_SECONDS or not math.isfinite(bounded_status_expiry) or bounded_status_expiry < now + float(ttl_seconds) - or bounded_status_expiry > now + MAX_STATUS_GRANT_TTL_SECONDS - ): + or bounded_status_expiry > now + MAX_STATUS_GRANT_TTL_SECONDS): raise HostedRoomGrantError("room grant lifetime is invalid") allowed = tuple(sorted(set(permissions))) if not allowed or not set(allowed) <= _GRANT_PERMISSIONS: @@ -501,19 +479,16 @@ def issue_room_grant( scope = { "grant_id": grant_id, "room_id": room_id, "home_install_id": home_install_id, "authority_gateway_id": authority_gateway_id, "authority_epoch": authority_epoch, - "member_id": member_id, "target_install_id": target_install_id, "target_profile": target_profile, - } + "member_id": member_id, "target_install_id": target_install_id, "target_profile": target_profile} payload = { "version": PROTOCOL_VERSION, **{name: _DISPATCH_FIELDS.get(name, _identifier)(scope[name], field=name) for name in _GRANT_SCOPE}, "execution_policy_digest": _digest( execution_policy_digest or execution_policy_mapping(target_profile=target_profile)["policy_digest"], - field="execution_policy_digest", - ), + field="execution_policy_digest"), "permissions": list(allowed), "issued_at": now, "expires_at": now + float(ttl_seconds), - "status_expires_at": bounded_status_expiry, - } + "status_expires_at": bounded_status_expiry} encoded = _canonical_json(payload) signature = hmac.new(secret, encoded, hashlib.sha256).digest() token = f"{_b64encode(encoded)}.{_b64encode(signature)}" @@ -524,8 +499,7 @@ def issue_room_grant( def verify_room_grant( secret: bytes, token: str, dispatch: HostedMemberDispatch, *, permission: str = "dispatch", - now: float | None = None, -) -> dict[str, Any]: + now: float | None = None) -> dict[str, Any]: """Verify one room grant against exact recipient dispatch coordinates.""" payload = decode_room_grant(secret, token, permission=permission, now=now) if payload["version"] != dispatch.protocol_version: @@ -536,9 +510,7 @@ def verify_room_grant( return payload -def decode_room_grant( - secret: bytes, token: str, *, permission: str, now: float | None = None -) -> dict[str, Any]: +def decode_room_grant(secret: bytes, token: str, *, permission: str, now: float | None = None) -> dict[str, Any]: """Verify grant signature, lifetime and operation without a dispatch.""" if not isinstance(token, str) or len(token.encode("utf-8")) > MAX_TOKEN_BYTES: raise HostedRoomGrantError("room grant is invalid") @@ -576,8 +548,7 @@ def decode_room_grant( def room_grant_needs_dispatch_refresh( - token: str, *, now: float | None = None, leeway_seconds: float = 5 * 60 -) -> bool: + token: str, *, now: float | None = None, leeway_seconds: float = 5 * 60) -> bool: """Read only grant timing to schedule target-validated refresh. This deliberately does not establish trust; the target validates the diff --git a/gateway/hosted_room_policy_checkpoint.py b/gateway/hosted_room_policy_checkpoint.py index c3d82d9251..d7c67de2a9 100644 --- a/gateway/hosted_room_policy_checkpoint.py +++ b/gateway/hosted_room_policy_checkpoint.py @@ -52,12 +52,9 @@ _SCHEMA_DDL = ( room_id TEXT PRIMARY KEY, schema_version INTEGER NOT NULL)""", ) -_ROOM_EVENT_COLUMNS = ( - "room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, created_at" -) +_ROOM_EVENT_COLUMNS = ("room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, created_at") _DELETE_ACTIVE_EVENTS_SQL = ( - "DELETE FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=?" -) + "DELETE FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=?") _TRANSCRIPT_EVENTS_SQL = f"""WITH transcript_events(seq) AS ( SELECT seq FROM hosted_room_policy_transcript WHERE room_id=? AND thread_id=? UNION ALL @@ -93,12 +90,10 @@ def _settled_message( """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() + (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, - ) + None) class HostedRoomPolicyCheckpoint: @@ -129,14 +124,11 @@ class HostedRoomPolicyCheckpoint: ) VALUES (?, ?, ?, ?, ?)""", ( event["room_id"], thread_id, discussion_event_id, int(event["seq"]), - json.dumps(dict(event), ensure_ascii=True, sort_keys=True, separators=(",", ":")), - ), - ) + json.dumps(dict(event), ensure_ascii=True, sort_keys=True, separators=(",", ":")))) @staticmethod def _store_transcript_event( - conn: sqlite3.Connection, *, event: Mapping[str, Any], thread_id: str, - settled_seq: int | None = None, + conn: sqlite3.Connection, *, event: Mapping[str, Any], thread_id: str, settled_seq: int | None = None, ) -> None: conn.execute( """INSERT INTO hosted_room_policy_transcript( @@ -144,20 +136,17 @@ class HostedRoomPolicyCheckpoint: ) VALUES (?, ?, ?, ?, ?) ON CONFLICT(room_id, thread_id, seq) DO UPDATE SET settled_seq=COALESCE(excluded.settled_seq, hosted_room_policy_transcript.settled_seq)""", - (event["room_id"], thread_id, int(event["seq"]), str(event["kind"]), settled_seq), - ) + (event["room_id"], thread_id, int(event["seq"]), str(event["kind"]), settled_seq)) if event["kind"] in {"message.user", "message.member"}: cutoff = conn.execute( """SELECT seq FROM hosted_room_policy_transcript WHERE room_id=? AND thread_id=? AND kind IN ('message.user', 'message.member') ORDER BY seq DESC LIMIT 1 OFFSET ?""", - (event["room_id"], thread_id, MAX_THREAD_TRANSCRIPT_EVENTS - 1), - ).fetchone() + (event["room_id"], thread_id, MAX_THREAD_TRANSCRIPT_EVENTS - 1)).fetchone() if cutoff is not None: conn.execute( "DELETE FROM hosted_room_policy_transcript WHERE room_id=? AND thread_id=? AND seq None: """Migrate bounded committed thread history from the durable room log.""" @@ -167,8 +156,7 @@ class HostedRoomPolicyCheckpoint: for row in conn.execute( """SELECT seq, payload_json FROM hosted_room_events WHERE room_id=? AND seq<=? AND kind='turn.settled' ORDER BY seq""", - (room_id, through_seq), - ): + (room_id, through_seq)): message_event_id = str(json.loads(row["payload_json"]).get("message_event_id") or "") if message_event_id: settled_seq_by_message[message_event_id] = int(row["seq"]) @@ -176,8 +164,7 @@ class HostedRoomPolicyCheckpoint: f"""SELECT {_ROOM_EVENT_COLUMNS} FROM hosted_room_events WHERE room_id=? AND seq<=? AND kind IN ('message.user', 'message.member') ORDER BY seq""", - (room_id, through_seq), - ) + (room_id, through_seq)) for row in rows: if row["kind"] == "message.member" and row["event_id"] not in settled_seq_by_message: continue @@ -186,38 +173,31 @@ class HostedRoomPolicyCheckpoint: if thread_id: self._store_transcript_event( conn, event=event, thread_id=thread_id, - settled_seq=settled_seq_by_message.get(str(row["event_id"])), - ) + settled_seq=settled_seq_by_message.get(str(row["event_id"]))) def _discussion_events( self, conn: sqlite3.Connection, *, room_id: str, thread_id: str, discussion_event_id: str, - bound_error: str, - ) -> list[dict[str, Any]]: + bound_error: str) -> list[dict[str, Any]]: """Merge the thread transcript with the active projection, ordered by seq.""" active_rows = conn.execute( """SELECT event_json FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=? ORDER BY seq LIMIT ?""", - (room_id, discussion_event_id, MAX_ACTIVE_POLICY_EVENTS + 1), - ).fetchall() + (room_id, discussion_event_id, MAX_ACTIVE_POLICY_EVENTS + 1)).fetchall() if len(active_rows) > MAX_ACTIVE_POLICY_EVENTS: raise RuntimeError(bound_error) rows = conn.execute( - _TRANSCRIPT_EVENTS_SQL, (room_id, thread_id, room_id, thread_id, room_id) - ).fetchall() + _TRANSCRIPT_EVENTS_SQL, (room_id, thread_id, room_id, thread_id, room_id)).fetchall() events_by_seq = { int(event["seq"]): event for event in ( *(_event_from_room_row(row) for row in rows), - *(json.loads(row["event_json"]) for row in active_rows), - ) - } + *(json.loads(row["event_json"]) for row in active_rows))} return [events_by_seq[seq] for seq in sorted(events_by_seq)] # -- per-kind projection handlers (dispatched by _apply_event) ----------- def _apply_user_message( - self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any] - ) -> None: + self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None: room_id = str(event["room_id"]) thread_id = str(payload.get("thread_id") or "") event_id = str(event.get("event_id") or "") @@ -230,14 +210,12 @@ class HostedRoomPolicyCheckpoint: ON CONFLICT(room_id, thread_id) DO UPDATE SET discussion_event_id=excluded.discussion_event_id, latest_user_seq=excluded.latest_user_seq, completed=0""", - (room_id, thread_id, event_id, int(event["seq"])), - ) + (room_id, thread_id, event_id, int(event["seq"]))) self._store_active_event(conn, event=event, thread_id=thread_id, discussion_event_id=event_id) self._store_transcript_event(conn, event=event, thread_id=thread_id) def _apply_discussion_event( - self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any] - ) -> None: + self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None: """Index member messages and terminal turn outcomes of a known discussion.""" room_id = str(event["room_id"]) seq = int(event["seq"]) @@ -246,26 +224,22 @@ class HostedRoomPolicyCheckpoint: discussion_event_id = str(payload.get("discussion_event_id") or "") source = conn.execute( "SELECT 1 FROM hosted_room_policy_events WHERE room_id=? AND discussion_event_id=? LIMIT 1", - (room_id, discussion_event_id), - ).fetchone() + (room_id, discussion_event_id)).fetchone() if source is None: return self._store_active_event( - conn, event=event, thread_id=thread_id, discussion_event_id=discussion_event_id - ) + conn, event=event, thread_id=thread_id, discussion_event_id=discussion_event_id) if kind not in _TERMINAL_KINDS: return task_id = str(payload.get("task_id") or "") execution_generation = ( - int(payload.get("execution_generation") or 0) if kind == "turn.deferred" else 0 - ) + int(payload.get("execution_generation") or 0) if kind == "turn.deferred" else 0) if task_id: conn.execute( """INSERT OR IGNORE INTO hosted_room_policy_publications( room_id, task_id, kind, execution_generation, seq ) VALUES (?, ?, ?, ?, ?)""", - (room_id, task_id, kind, execution_generation, seq), - ) + (room_id, task_id, kind, execution_generation, seq)) 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"): @@ -280,37 +254,30 @@ class HostedRoomPolicyCheckpoint: ) VALUES (?, ?, ?, ?) ON CONFLICT(room_id, thread_id, member_id) DO UPDATE SET seen_through_seq=MAX(hosted_room_policy_watermarks.seen_through_seq, excluded.seen_through_seq)""", - (room_id, thread_id, member_id, seen_through_seq), - ) + (room_id, thread_id, member_id, seen_through_seq)) def _apply_room_activity( - self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any] - ) -> None: + self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None: room_id = str(event["room_id"]) thread_id = str(payload.get("thread_id") or "") discussion_event_id = str(payload.get("discussion_event_id") or "") conn.execute(_DELETE_ACTIVE_EVENTS_SQL, (room_id, discussion_event_id)) conn.execute( - "DELETE FROM hosted_room_policy_threads WHERE room_id=? AND thread_id=?", - (room_id, thread_id), - ) + "DELETE FROM hosted_room_policy_threads WHERE room_id=? AND thread_id=?", (room_id, thread_id)) def _apply_stop_requested( - self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any] - ) -> None: + self, conn: sqlite3.Connection, event: Mapping[str, Any], payload: Mapping[str, Any]) -> None: conn.execute( """UPDATE hosted_room_policy_cursors SET stopped_through_seq=MAX(stopped_through_seq, ?) WHERE room_id=?""", - (int(event["seq"]), str(event["room_id"])), - ) + (int(event["seq"]), str(event["room_id"]))) _APPLY_BY_KIND: dict[str, Callable[..., None]] = { "message.user": _apply_user_message, "message.member": _apply_discussion_event, **dict.fromkeys(_TERMINAL_KINDS, _apply_discussion_event), "room.activity": _apply_room_activity, - "room.stop_requested": _apply_stop_requested, - } + "room.stop_requested": _apply_stop_requested} def _apply_event(self, conn: sqlite3.Connection, event: Mapping[str, Any]) -> None: handler = self._APPLY_BY_KIND.get(str(event.get("kind") or "")) @@ -326,15 +293,12 @@ class HostedRoomPolicyCheckpoint: """INSERT OR IGNORE INTO hosted_room_policy_cursors( room_id, through_seq, stopped_through_seq, updated_at ) VALUES (?, 0, 0, 0)""", - (room_id,), - ) + (room_id,)) row = conn.execute( - "SELECT through_seq FROM hosted_room_policy_cursors WHERE room_id=?", (room_id,) - ).fetchone() + "SELECT through_seq FROM hosted_room_policy_cursors WHERE room_id=?", (room_id,)).fetchone() cursor = int(row["through_seq"]) transcript_state = conn.execute( - "SELECT schema_version FROM hosted_room_policy_transcript_state WHERE room_id=?", - (room_id,), + "SELECT schema_version FROM hosted_room_policy_transcript_state WHERE room_id=?", (room_id,), ).fetchone() if transcript_state is None or int(transcript_state["schema_version"]) < _TRANSCRIPT_SCHEMA_VERSION: self._backfill_transcript(conn, room_id=room_id, through_seq=cursor) @@ -342,8 +306,7 @@ class HostedRoomPolicyCheckpoint: """INSERT INTO hosted_room_policy_transcript_state(room_id, schema_version) VALUES (?, ?) ON CONFLICT(room_id) DO UPDATE SET schema_version=excluded.schema_version""", - (room_id, _TRANSCRIPT_SCHEMA_VERSION), - ) + (room_id, _TRANSCRIPT_SCHEMA_VERSION)) return cursor def sync(self, *, room_id: str, latest_seq: int) -> int: @@ -356,8 +319,7 @@ class HostedRoomPolicyCheckpoint: while cursor < latest_seq: page = hosted_rooms.read_events( - self.db_path, room_id=room_id, since_seq=cursor, limit=hosted_rooms.MAX_LOG_LIMIT - ) + self.db_path, room_id=room_id, since_seq=cursor, limit=hosted_rooms.MAX_LOG_LIMIT) rows = [event for event in page.get("events", []) if isinstance(event, Mapping)] next_cursor = int(page.get("cursor") or cursor) if not rows or next_cursor <= cursor: @@ -369,8 +331,7 @@ class HostedRoomPolicyCheckpoint: self._apply_event(conn, event) updated = conn.execute( "UPDATE hosted_room_policy_cursors SET through_seq=?, updated_at=? WHERE room_id=?", - (next_cursor, float(rows[-1].get("created_at") or 0), room_id), - ) + (next_cursor, float(rows[-1].get("created_at") or 0), room_id)) if updated.rowcount != 1: raise RuntimeError("room policy cursor disappeared during replay") cursor = next_cursor @@ -381,43 +342,35 @@ class HostedRoomPolicyCheckpoint: through_seq = self.sync(room_id=room_id, latest_seq=latest_seq) with self._connect() as conn: cursor = conn.execute( - "SELECT stopped_through_seq FROM hosted_room_policy_cursors WHERE room_id=?", - (room_id,), + "SELECT stopped_through_seq FROM hosted_room_policy_cursors WHERE room_id=?", (room_id,), ).fetchone() stopped_through_seq = int(cursor["stopped_through_seq"]) thread = conn.execute( """SELECT thread_id, discussion_event_id FROM hosted_room_policy_threads WHERE room_id=? AND completed=0 AND latest_user_seq>? ORDER BY latest_user_seq, thread_id LIMIT 1""", - (room_id, stopped_through_seq), - ).fetchone() + (room_id, stopped_through_seq)).fetchone() if thread is None: return PolicySnapshot( through_seq=through_seq, stopped_through_seq=stopped_through_seq, - events=(), watermarks={}, - ) + events=(), watermarks={}) thread_id = str(thread["thread_id"]) events = self._discussion_events( conn, room_id=room_id, thread_id=thread_id, discussion_event_id=str(thread["discussion_event_id"]), - bound_error="active room policy projection exceeded its bound", - ) + bound_error="active room policy projection exceeded its bound") watermark_rows = conn.execute( """SELECT member_id, seen_through_seq FROM hosted_room_policy_watermarks WHERE room_id=? AND thread_id=?""", - (room_id, thread_id), - ).fetchall() + (room_id, thread_id)).fetchall() return PolicySnapshot( through_seq=through_seq, stopped_through_seq=stopped_through_seq, events=tuple(events), watermarks={ (thread_id, str(row["member_id"])): int(row["seen_through_seq"]) - for row in watermark_rows - }, - ) + for row in watermark_rows}) def publication_exists( - self, *, room_id: str, task_id: str, status: str, execution_generation: int - ) -> bool: + self, *, room_id: str, task_id: str, status: str, execution_generation: int) -> bool: """Return whether one exact driver outcome is already in the room log.""" if status == "deferred": sql = """SELECT 1 FROM hosted_room_policy_publications @@ -435,25 +388,20 @@ class HostedRoomPolicyCheckpoint: with self._connect() as conn: source = conn.execute( "SELECT discussion_event_id, thread_id FROM hosted_room_policy_events WHERE room_id=? AND seq=?", - (room_id, source_event_seq), - ).fetchone() + (room_id, source_event_seq)).fetchone() if source is None: return [] return self._discussion_events( conn, room_id=room_id, thread_id=str(source["thread_id"]), discussion_event_id=str(source["discussion_event_id"]), - bound_error="task policy projection exceeded its bound", - ) + bound_error="task policy projection exceeded its bound") def compact_completed(self, *, room_id: str) -> None: """Drop any completed projections left by an interrupted sync.""" with self._connect() as conn: completed = conn.execute( "SELECT discussion_event_id FROM hosted_room_policy_threads WHERE room_id=? AND completed=1", - (room_id,), - ).fetchall() + (room_id,)).fetchall() for row in completed: conn.execute(_DELETE_ACTIVE_EVENTS_SQL, (room_id, str(row["discussion_event_id"]))) - conn.execute( - "DELETE FROM hosted_room_policy_threads WHERE room_id=? AND completed=1", (room_id,) - ) + conn.execute("DELETE FROM hosted_room_policy_threads WHERE room_id=? AND completed=1", (room_id,)) diff --git a/gateway/hosted_room_replicas.py b/gateway/hosted_room_replicas.py index 08ffa915a3..8441fbed73 100644 --- a/gateway/hosted_room_replicas.py +++ b/gateway/hosted_room_replicas.py @@ -33,8 +33,7 @@ from gateway.hosted_rooms import ( _validate_identifier, _validate_members, _validate_room_name, - local_authority_gateway_id, -) + local_authority_gateway_id) from gateway.hosted_rooms_common import bounded_int, clock, utf8_len MAX_REPLICA_ROOMS = 256 @@ -109,8 +108,7 @@ def _event_bytes(event: dict[str, Any]) -> int: return utf8_len( str(event["event_id"]), str(event["kind"]), json.dumps(event["actor"], ensure_ascii=False, separators=(",", ":")), - json.dumps(event["payload"], ensure_ascii=False, separators=(",", ":")), - ) + json.dumps(event["payload"], ensure_ascii=False, separators=(",", ":"))) def _validate_page(page: Any) -> tuple[list[dict[str, Any]], dict[str, Any]]: @@ -123,8 +121,7 @@ def _validate_page(page: Any) -> tuple[list[dict[str, Any]], dict[str, Any]]: if not isinstance(authority, dict): raise ReplicaError("page.authority is required for replication") gateway_id = _validate_identifier( - authority.get("gateway_id"), label="page.authority.gateway_id", max_chars=MAX_ACTOR_ID_CHARS - ) + authority.get("gateway_id"), label="page.authority.gateway_id", max_chars=MAX_ACTOR_ID_CHARS) epoch = _positive_int(authority.get("epoch"), message="page.authority.epoch must be a positive integer") previous_seq: int | None = None for event in events: @@ -149,8 +146,7 @@ def _replica_row_state(conn: sqlite3.Connection, room_id: str) -> tuple[sqlite3. row = conn.execute( """SELECT authority_gateway_id, authority_epoch, last_seq, latest_seq, event_bytes FROM hosted_room_replicas WHERE room_id=?""", - (room_id,), - ).fetchone() + (room_id,)).fetchone() if row is None: count = conn.execute("SELECT COUNT(*) FROM hosted_room_replicas").fetchone()[0] if int(count) >= MAX_REPLICA_ROOMS: @@ -161,8 +157,7 @@ def _replica_row_state(conn: sqlite3.Connection, room_id: str) -> tuple[sqlite3. def _store_replica( conn: sqlite3.Connection, *, is_new: bool, room_id: str, room_name: str, members_json: str, - authority: dict[str, Any], new_last: int, latest_seq: int, added_bytes: int, now: float, -) -> None: + authority: dict[str, Any], new_last: int, latest_seq: int, added_bytes: int, now: float) -> None: """INSERT the replica row for a new room, else UPDATE it (event_bytes accumulates).""" values = (room_name, members_json, authority["gateway_id"], authority["epoch"], new_last, max(latest_seq, new_last)) if is_new: @@ -170,15 +165,13 @@ def _store_replica( """INSERT INTO hosted_room_replicas (room_id, name, members_json, authority_gateway_id, authority_epoch, last_seq, latest_seq, event_bytes, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", - (room_id, *values, added_bytes, now, now), - ) + (room_id, *values, added_bytes, now, now)) else: conn.execute( """UPDATE hosted_room_replicas SET name=?, members_json=?, authority_gateway_id=?, authority_epoch=?, last_seq=?, latest_seq=?, event_bytes=event_bytes+?, updated_at=? WHERE room_id=?""", - (*values, added_bytes, now, room_id), - ) + (*values, added_bytes, now, room_id)) def ingest_page( @@ -214,9 +207,7 @@ def ingest_page( _INSERT_REPLICA_EVENT, ( room_id, int(event["seq"]), event["event_id"], event["kind"], actor_json, - event.get("authority_epoch"), payload_json, float(event.get("created_at") or now), - ), - ) + event.get("authority_epoch"), payload_json, float(event.get("created_at") or now))) added_bytes += size new_last = int(new_events[-1]["seq"]) if new_events else last_seq latest_seq = page.get("latest_seq") @@ -224,15 +215,13 @@ def ingest_page( latest_seq = new_last _store_replica( conn, is_new=row is None, room_id=room_id, room_name=room_name, members_json=members_json, - authority=authority, new_last=new_last, latest_seq=latest_seq, added_bytes=added_bytes, now=now, - ) + authority=authority, new_last=new_last, latest_seq=latest_seq, added_bytes=added_bytes, now=now) return { "room_id": room_id, "stored_seq": new_last, "ingested": len(new_events), "authority": authority, - "caught_up": new_last >= max(latest_seq, new_last), - } + "caught_up": new_last >= max(latest_seq, new_last)} def replica_state(db_path: Path | str, *, room_id: Any) -> dict[str, Any]: @@ -251,8 +240,7 @@ def replica_state(db_path: Path | str, *, room_id: Any) -> dict[str, Any]: "latest_seq": int(row["latest_seq"]), "event_bytes": int(row["event_bytes"]), "created_at": float(row["created_at"]), - "updated_at": float(row["updated_at"]), - } + "updated_at": float(row["updated_at"])} def promote_replica( @@ -292,8 +280,7 @@ def promote_replica( "authority_gateway_id": local_gateway, "authority_epoch": target_epoch, "promoted_from_replica": True, - "reason": reason, - }) + "reason": reason}) claim_bytes = utf8_len(claim_event_id, "authority.claimed", claim_actor_json, claim_payload_json) conn.execute( @@ -303,22 +290,17 @@ def promote_replica( VALUES (?, ?, ?, ?, ?, ?, ?, 1, ?, ?, NULL)""", ( room_id, replica["name"], replica["members_json"], local_gateway, target_epoch, - claim_seq + 1, int(replica["event_bytes"]) + claim_bytes, now, now, - ), - ) + claim_seq + 1, int(replica["event_bytes"]) + claim_bytes, now, now)) conn.execute( f"""INSERT INTO hosted_room_events {_EVENT_COLUMNS} SELECT room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, created_at FROM hosted_room_replica_events WHERE room_id=?""", - (room_id,), - ) + (room_id,)) conn.execute( _INSERT_ROOM_EVENT, ( room_id, claim_seq, claim_event_id, "authority.claimed", - claim_actor_json, target_epoch, claim_payload_json, now, - ), - ) + claim_actor_json, target_epoch, claim_payload_json, now)) conn.execute("DELETE FROM hosted_room_replica_events WHERE room_id=?", (room_id,)) conn.execute("DELETE FROM hosted_room_replicas WHERE room_id=?", (room_id,)) return { @@ -328,8 +310,7 @@ def promote_replica( "previous_gateway_id": previous_gateway, "previous_epoch": previous_epoch, "claim_seq": claim_seq, - "latest_seq": claim_seq, - } + "latest_seq": claim_seq} def demote_room( @@ -344,8 +325,7 @@ def demote_room( """ room_id = _room_id(room_id) observed_gateway_id = _validate_identifier( - observed_gateway_id, label="observed_gateway_id", max_chars=MAX_ACTOR_ID_CHARS - ) + observed_gateway_id, label="observed_gateway_id", max_chars=MAX_ACTOR_ID_CHARS) observed_epoch = _positive_int(observed_epoch, message="observed_epoch must be a positive integer") now = clock(now) local_gateway = local_authority_gateway_id() @@ -354,8 +334,7 @@ def demote_room( row = conn.execute( """SELECT authority_gateway_id, authority_epoch, next_seq FROM hosted_rooms WHERE room_id=? AND disbanded_at IS NULL""", - (room_id,), - ).fetchone() + (room_id,)).fetchone() if row is None: raise ReplicaError("room not found in the local authoritative store") current_gateway = str(row["authority_gateway_id"]) @@ -365,8 +344,7 @@ def demote_room( "room_id": room_id, "authority_gateway_id": current_gateway, "authority_epoch": current_epoch, - "idempotent": True, - } + "idempotent": True} if observed_epoch <= current_epoch: raise ReplicaEpochRegressionError("observed epoch does not supersede the stored authority") if current_gateway != local_gateway: @@ -374,24 +352,19 @@ def demote_room( lost_actor_json, lost_payload_json = _control_event_json({ "previous_gateway_id": current_gateway, "authority_gateway_id": observed_gateway_id, - "authority_epoch": observed_epoch, - }) + "authority_epoch": observed_epoch}) conn.execute( _INSERT_ROOM_EVENT, ( room_id, int(row["next_seq"]), f"system:authority-lost:{observed_epoch}", - "authority.lost", lost_actor_json, observed_epoch, lost_payload_json, now, - ), - ) + "authority.lost", lost_actor_json, observed_epoch, lost_payload_json, now)) conn.execute( """UPDATE hosted_rooms SET authority_gateway_id=?, authority_epoch=?, next_seq=next_seq+1, revision=revision+1, updated_at=? WHERE room_id=?""", - (observed_gateway_id, observed_epoch, now, room_id), - ) + (observed_gateway_id, observed_epoch, now, room_id)) return { "room_id": room_id, "authority_gateway_id": observed_gateway_id, "authority_epoch": observed_epoch, - "idempotent": False, - } + "idempotent": False} diff --git a/gateway/hosted_rooms.py b/gateway/hosted_rooms.py index d7bf97b3b0..58cfaa8ee8 100644 --- a/gateway/hosted_rooms.py +++ b/gateway/hosted_rooms.py @@ -21,8 +21,7 @@ from typing import Any, Mapping, NoReturn from gateway.hosted_rooms_common import ( bounded_int, canonical_json, clock as _now, compact_json, connect, identifier, open_sqlite, - table_columns, table_exists, transaction, -) + table_columns, table_exists, transaction) PROTOCOL_VERSION = 2 MAX_ROOM_ID_CHARS = 128 @@ -52,28 +51,23 @@ _JOURNAL_MODE_LOCK_RETRIES = 8 _EVENT_KIND_RE = re.compile(r"^[a-z][a-z0-9_.-]*$") _CONTROL_EVENT_KINDS = frozenset({ - "authority.claimed", "authority.lost", "room.disbanded", "room.stop_requested", -}) + "authority.claimed", "authority.lost", "room.disbanded", "room.stop_requested"}) _EVENT_KINDS_BY_ACTOR = { "user": frozenset({"message.user"}), "member": frozenset({"message.member"}), "gateway": frozenset({ "member.unavailable", "room.activity", "room.stop_requested", "turn.deferred", - "turn.reassigned", "turn.cancelled", "turn.failed", "turn.settled", "turn.started", - }), + "turn.reassigned", "turn.cancelled", "turn.failed", "turn.settled", "turn.started"}), "system": frozenset({ "authority.claimed", "authority.lost", "room.created", "room.disbanded", - "room.members_changed", "room.renamed", - }), -} + "room.members_changed", "room.renamed"})} _ACTOR_FIELDS = frozenset({"kind", "id", "display_name", "profile", "connection_id"}) # --- schema ----------------------------------------------------------------- _REMOTE_RUN_IDENTITY_COLUMNS = ( "room_id", "home_install_id", "authority_gateway_id", "authority_epoch", "member_id", - "target_install_id", "target_profile", "task_id", "execution_generation", -) + "target_install_id", "target_profile", "task_id", "execution_generation") _REMOTE_RUNS_BODY = """ room_id TEXT NOT NULL, home_install_id TEXT NOT NULL, @@ -163,50 +157,36 @@ _SCHEMA_DDL = ( _REQUIRED_COLUMNS = tuple( ( re.search(r"EXISTS (\w+)", ddl).group(1), - frozenset(re.findall(r"^\s*(\w+) (?:TEXT|INTEGER|REAL)\b", ddl.split("(", 1)[1], re.M)), - ) - for ddl in _SCHEMA_DDL -) + frozenset(re.findall(r"^\s*(\w+) (?:TEXT|INTEGER|REAL)\b", ddl.split("(", 1)[1], re.M))) + for ddl in _SCHEMA_DDL) _REMOTE_RUN_SCHEMA_COLUMNS = _REQUIRED_COLUMNS[4][1] # --- SQL fragments (statement text must stay byte-stable after whitespace normalisation) --- -_EVENT_COLUMNS = ( - "room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, created_at" -) +_EVENT_COLUMNS = ("room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, created_at") _SELECT_EVENT = f"SELECT {_EVENT_COLUMNS} FROM hosted_room_events WHERE room_id=? AND event_id=?" _INSERT_EVENT = ( - f"INSERT INTO hosted_room_events ({_EVENT_COLUMNS}) VALUES (?, ?, ?, ?, ?, ?, ?, ?)" -) + f"INSERT INTO hosted_room_events ({_EVENT_COLUMNS}) VALUES (?, ?, ?, ?, ?, ?, ?, ?)") _ROOM_COLUMNS = ( "room_id, name, members_json, authority_gateway_id, authority_epoch, next_seq, revision," - " created_at, updated_at, disbanded_at" -) + " created_at, updated_at, disbanded_at") _ROOM_COLUMNS_WITH_BYTES = ( "room_id, name, members_json, authority_gateway_id, authority_epoch, next_seq, event_bytes," - " revision, created_at, updated_at, disbanded_at" -) + " revision, created_at, updated_at, disbanded_at") _SELECT_ROOM = f"SELECT {_ROOM_COLUMNS} FROM hosted_rooms WHERE room_id=?" _SELECT_ROOM_WITH_BYTES = f"SELECT {_ROOM_COLUMNS_WITH_BYTES} FROM hosted_rooms WHERE room_id=?" _SUM_EVENT_BYTES = "SELECT COALESCE(SUM(event_bytes), 0) FROM hosted_rooms" -_INSERT_RETIRED = ( - "INSERT OR IGNORE INTO hosted_room_retired_ids (room_id, retired_at) VALUES (?, ?)" -) +_INSERT_RETIRED = ("INSERT OR IGNORE INTO hosted_room_retired_ids (room_id, retired_at) VALUES (?, ?)") _RETIRE_FROM_ROOMS = ( "INSERT OR IGNORE INTO hosted_room_retired_ids (room_id, retired_at)" - " SELECT room_id, disbanded_at FROM hosted_rooms WHERE {where}" -) + " SELECT room_id, disbanded_at FROM hosted_rooms WHERE {where}") _LINK_COLUMNS = ( "room_id", "member_id", "target_url", "target_profile", "grant", "catalog_json", - "cancellation_scope_id", "trace_id", "transport_security", "status", "updated_at", -) + "cancellation_scope_id", "trace_id", "transport_security", "status", "updated_at") _REMOTE_RUN_WHERE = " AND ".join(f"{column}=?" for column in _REMOTE_RUN_IDENTITY_COLUMNS) -_LIVE_RESERVATION_WHERE = ( - "WHERE room_id=? AND target_profile=? AND expires_at>? AND revoked_at IS NULL" -) +_LIVE_RESERVATION_WHERE = ("WHERE room_id=? AND target_profile=? AND expires_at>? AND revoked_at IS NULL") _SELECT_LIVE_RESERVATION = ( - f"SELECT 1 FROM hosted_room_peer_reservations {_LIVE_RESERVATION_WHERE} LIMIT 1" -) + f"SELECT 1 FROM hosted_room_peer_reservations {_LIVE_RESERVATION_WHERE} LIMIT 1") class HostedRoomError(ValueError): @@ -247,7 +227,6 @@ class AuthoritySupersededError(AuthorityConflictError): # --- validation --------------------------------------------------------------- - _canonical_json = partial(canonical_json, error=HostedRoomError, ensure_ascii=False) _validate_identifier = partial(identifier, error=HostedRoomError) _room_id = partial(_validate_identifier, label="room_id", max_chars=MAX_ROOM_ID_CHARS) @@ -262,8 +241,7 @@ def _actor_id(value: Any, label: str) -> str: def _require_positive_int(value: Any, label: str) -> int: return bounded_int( - value, error=HostedRoomError, message=f"{label} must be a positive integer", low=1 - ) + value, error=HostedRoomError, message=f"{label} must be a positive integer", low=1) def _bounded_limit(value: Any, maximum: int) -> int: @@ -283,8 +261,7 @@ def _claim_payload_json(previous_gateway_id: str, new_gateway_id: str, epoch: in return _payload_json({ "previous_gateway_id": previous_gateway_id, "authority_gateway_id": new_gateway_id, - "authority_epoch": epoch, - }) + "authority_epoch": epoch}) def user_event_id(client_event_id: Any) -> str: @@ -295,8 +272,7 @@ def user_event_id(client_event_id: Any) -> str: _validate_room_name = partial( _validate_identifier, label="name", max_chars=MAX_ROOM_NAME_CHARS, pattern=None, - invalid="invalid room name", -) + invalid="invalid room name") def _validate_members(value: Any) -> tuple[list[dict[str, Any]], str]: @@ -335,8 +311,7 @@ def _legacy_members_match(existing_json: str, proposed: list[dict[str, Any]]) -> _validate_event_kind = partial( _validate_identifier, label="kind", max_chars=MAX_EVENT_KIND_CHARS, pattern=_EVENT_KIND_RE, - invalid="invalid event kind", -) + invalid="invalid event kind") def _optional_actor_field(actor: dict[str, Any], field: str, max_chars: int) -> str: @@ -366,8 +341,7 @@ def _validate_actor(value: Any, *, kind: str) -> tuple[dict[str, str], str]: for field, max_chars in ( ("display_name", MAX_ACTOR_LABEL_CHARS), ("profile", MAX_ACTOR_ID_CHARS), - ("connection_id", MAX_ACTOR_ID_CHARS), - ): + ("connection_id", MAX_ACTOR_ID_CHARS)): field_value = _optional_actor_field(value, field, max_chars) if field_value: actor[field] = field_value @@ -385,8 +359,7 @@ def _primary_key_columns(conn: sqlite3.Connection, table: str) -> tuple[str, ... def _remote_run_schema_current(conn: sqlite3.Connection, columns: frozenset[str]) -> bool: return ( _REMOTE_RUN_SCHEMA_COLUMNS.issubset(columns) - and _primary_key_columns(conn, "hosted_room_remote_runs") == _REMOTE_RUN_IDENTITY_COLUMNS - ) + and _primary_key_columns(conn, "hosted_room_remote_runs") == _REMOTE_RUN_IDENTITY_COLUMNS) def _migrate_remote_run_schema(conn: sqlite3.Connection) -> None: @@ -410,8 +383,7 @@ def _migrate_remote_run_schema(conn: sqlite3.Connection) -> None: SELECT room_id, {home}, {gateway}, {epoch}, member_id, task_id, execution_generation, target_install_id, target_profile, run_id, session_id, created_at, updated_at - FROM hosted_room_remote_runs""" - ) + FROM hosted_room_remote_runs""") conn.execute("DROP TABLE hosted_room_remote_runs") conn.execute("ALTER TABLE hosted_room_remote_runs_migrating RENAME TO hosted_room_remote_runs") @@ -424,26 +396,20 @@ _LEGACY_ACTOR_JSON = _system_actor_json("legacy").replace("'", "''") _LEGACY_COLUMN_DDL = ( ( "hosted_rooms", "authority_gateway_id", - "ALTER TABLE hosted_rooms ADD COLUMN authority_gateway_id TEXT NOT NULL DEFAULT 'legacy'", - ), + "ALTER TABLE hosted_rooms ADD COLUMN authority_gateway_id TEXT NOT NULL DEFAULT 'legacy'"), ( "hosted_rooms", "authority_epoch", - "ALTER TABLE hosted_rooms ADD COLUMN authority_epoch INTEGER NOT NULL DEFAULT 1", - ), + "ALTER TABLE hosted_rooms ADD COLUMN authority_epoch INTEGER NOT NULL DEFAULT 1"), ( "hosted_rooms", "event_bytes", - "ALTER TABLE hosted_rooms ADD COLUMN event_bytes INTEGER NOT NULL DEFAULT 0", - ), + "ALTER TABLE hosted_rooms ADD COLUMN event_bytes INTEGER NOT NULL DEFAULT 0"), ( "hosted_room_events", "actor_json", "ALTER TABLE hosted_room_events " - f"ADD COLUMN actor_json TEXT NOT NULL DEFAULT '{_LEGACY_ACTOR_JSON}'", - ), + f"ADD COLUMN actor_json TEXT NOT NULL DEFAULT '{_LEGACY_ACTOR_JSON}'"), ( "hosted_room_events", "authority_epoch", - "ALTER TABLE hosted_room_events ADD COLUMN authority_epoch INTEGER", - ), -) + "ALTER TABLE hosted_room_events ADD COLUMN authority_epoch INTEGER")) def _migrate_legacy_columns(conn: sqlite3.Connection) -> None: @@ -481,9 +447,7 @@ def _initialize_schema(conn: sqlite3.Connection) -> None: conn.execute(_RETIRE_FROM_ROOMS.format(where="disbanded_at IS NOT NULL")) _migrate_remote_run_schema(conn) conn.execute( - "CREATE INDEX IF NOT EXISTS idx_hosted_room_events_cursor " - "ON hosted_room_events(room_id, seq)" - ) + "CREATE INDEX IF NOT EXISTS idx_hosted_room_events_cursor ON hosted_room_events(room_id, seq)") if not _schema_is_current(conn): raise HostedRoomError("hosted room schema migration did not complete") @@ -497,8 +461,7 @@ def _schema_is_current(conn: sqlite3.Connection) -> bool: if table == "hosted_room_remote_runs" and not _remote_run_schema_current(conn, columns): return False index = conn.execute( - "SELECT 1 FROM sqlite_master WHERE type='index' AND name='idx_hosted_room_events_cursor'" - ).fetchone() + "SELECT 1 FROM sqlite_master WHERE type='index' AND name='idx_hosted_room_events_cursor'").fetchone() return index is not None @@ -524,8 +487,7 @@ def local_authority_gateway_id() -> str: def _connect(db_path: Path | str) -> sqlite3.Connection: return connect( db_path, db_label="state.db (hosted_rooms)", ready=_schema_is_current, - initialize=lambda conn: _initialize_schema(conn), lock_retries=_JOURNAL_MODE_LOCK_RETRIES, - ) + initialize=lambda conn: _initialize_schema(conn), lock_retries=_JOURNAL_MODE_LOCK_RETRIES) def _read_connection(db_path: Path | str) -> sqlite3.Connection: @@ -554,13 +516,9 @@ def _raise_room_not_found(conn: sqlite3.Connection, room_id: str) -> NoReturn: # A retained disband tombstone still has replayable history. The # caller simply did not opt into reading disbanded rooms. raise RoomNotFoundError("hosted room not found") - retired = conn.execute( - "SELECT 1 FROM hosted_room_retired_ids WHERE room_id=?", (room_id,) - ).fetchone() + retired = conn.execute("SELECT 1 FROM hosted_room_retired_ids WHERE room_id=?", (room_id,)).fetchone() if retired is not None: - raise RoomHistoryExpiredError( - "Group Chat history expired; room_id remains permanently retired" - ) + raise RoomHistoryExpiredError("Group Chat history expired; room_id remains permanently retired") raise RoomNotFoundError("hosted room not found") @@ -582,8 +540,7 @@ def _room_from_row(row: sqlite3.Row, *, idempotent: bool = False) -> dict[str, A "revision": int(row["revision"]), "created_at": float(row["created_at"]), "updated_at": float(row["updated_at"]), - "idempotent": idempotent, - } + "idempotent": idempotent} # sqlite3.Row: ``x in row`` scans values, so ``.keys()`` is load-bearing. if "disbanded_at" in row.keys() and row["disbanded_at"] is not None: # noqa: SIM118 room["disbanded_at"] = float(row["disbanded_at"]) @@ -603,8 +560,7 @@ def _event_from_row(row: sqlite3.Row, *, idempotent: bool = False) -> dict[str, "authority_epoch": int(epoch) if epoch is not None else None, "payload": json.loads(row["payload_json"]), "created_at": float(row["created_at"]), - "idempotent": idempotent, - } + "idempotent": idempotent} def _load_event(conn: sqlite3.Connection, room_id: str, event_id: str) -> sqlite3.Row | None: @@ -622,20 +578,17 @@ def _gateway_event_bytes(conn: sqlite3.Connection) -> int: def _insert_event( conn: sqlite3.Connection, room: sqlite3.Row, room_id: str, seq: int, event_id: str, kind: str, - actor_json: str, epoch: int, payload_json: str, now: float, *, allow_control: bool = False, -) -> int: + actor_json: str, epoch: int, payload_json: str, now: float, *, allow_control: bool = False) -> int: """Capacity-check then INSERT one event at ``seq``; returns its accounted bytes.""" event_bytes = _prepare_event( - conn, room, event_id, kind, actor_json, payload_json, allow_control=allow_control - ) + conn, room, event_id, kind, actor_json, payload_json, allow_control=allow_control) conn.execute(_INSERT_EVENT, (room_id, seq, event_id, kind, actor_json, epoch, payload_json, now)) return event_bytes def _prepare_event( conn: sqlite3.Connection, room: sqlite3.Row, event_id: str, kind: str, actor_json: str, - payload_json: str, *, allow_control: bool = False, -) -> int: + payload_json: str, *, allow_control: bool = False) -> int: """Size one pending event and enforce per-room and gateway capacity; returns its bytes.""" additional_bytes = len((event_id + kind + actor_json + payload_json).encode("utf-8")) count_reserve = CONTROL_EVENT_COUNT_RESERVE if allow_control else 0 @@ -643,22 +596,18 @@ def _prepare_event( gateway_byte_limit = MAX_GATEWAY_EVENT_BYTES + byte_reserve if int(room["next_seq"]) - 1 >= MAX_EVENTS_PER_ROOM + count_reserve: raise HostedRoomError( - "This Group Chat reached its history limit. Start a new Group Chat to continue." - ) + "This Group Chat reached its history limit. Start a new Group Chat to continue.") if int(room["event_bytes"]) + additional_bytes > MAX_ROOM_EVENT_BYTES + byte_reserve: raise HostedRoomError( - "This Group Chat reached its storage limit. Start a new Group Chat to continue." - ) + "This Group Chat reached its storage limit. Start a new Group Chat to continue.") gateway_bytes = _gateway_event_bytes(conn) if gateway_bytes + additional_bytes > gateway_byte_limit: _prune_disbanded_rooms_locked( - conn, now=None, max_gateway_event_bytes=max(0, gateway_byte_limit - additional_bytes) - ) + conn, now=None, max_gateway_event_bytes=max(0, gateway_byte_limit - additional_bytes)) gateway_bytes = _gateway_event_bytes(conn) if gateway_bytes + additional_bytes > gateway_byte_limit: raise HostedRoomError( - "Group Chat storage is full on this host. Delete an old Group Chat and try again." - ) + "Group Chat storage is full on this host. Delete an old Group Chat and try again.") return additional_bytes @@ -666,12 +615,10 @@ def _prepare_event( # Deleted in this order when a disbanded room's payload is purged. _DEPENDENT_TABLES = ( - "hosted_room_policy_transcript_state", "hosted_room_policy_transcript", - "hosted_room_policy_publications", "hosted_room_policy_watermarks", - "hosted_room_policy_events", "hosted_room_policy_threads", "hosted_room_policy_cursors", - "hosted_room_driver_tasks", "hosted_room_driver_leases", "hosted_room_remote_runs", - "hosted_room_links", "hosted_room_peer_reservations", "hosted_room_events", -) + "hosted_room_policy_transcript_state", "hosted_room_policy_transcript", "hosted_room_policy_publications", + "hosted_room_policy_watermarks", "hosted_room_policy_events", "hosted_room_policy_threads", + "hosted_room_policy_cursors", "hosted_room_driver_tasks", "hosted_room_driver_leases", + "hosted_room_remote_runs", "hosted_room_links", "hosted_room_peer_reservations", "hosted_room_events") def _room_ids(conn: sqlite3.Connection, sql: str, params: tuple[Any, ...]) -> list[str]: @@ -679,22 +626,19 @@ def _room_ids(conn: sqlite3.Connection, sql: str, params: tuple[Any, ...]) -> li def _prune_disbanded_rooms_locked( - conn: sqlite3.Connection, *, now: float | None, max_gateway_event_bytes: int | None = None -) -> int: + conn: sqlite3.Connection, *, now: float | None, max_gateway_event_bytes: int | None = None) -> int: candidates: set[str] = set() if now is not None: candidates.update(_room_ids( conn, """SELECT room_id FROM hosted_rooms WHERE disbanded_at IS NOT NULL AND disbanded_at<=?""", - (now - DISBANDED_ROOM_RETENTION_SECONDS,), - )) + (now - DISBANDED_ROOM_RETENTION_SECONDS,))) candidates.update(_room_ids( conn, """SELECT room_id FROM hosted_rooms WHERE disbanded_at IS NOT NULL ORDER BY disbanded_at DESC, room_id ASC LIMIT -1 OFFSET ?""", - (MAX_DISBANDED_ROOM_TOMBSTONES,), - )) + (MAX_DISBANDED_ROOM_TOMBSTONES,))) if max_gateway_event_bytes is not None: retained_bytes = _gateway_event_bytes(conn) if retained_bytes > max_gateway_event_bytes: @@ -712,10 +656,8 @@ def _prune_disbanded_rooms_locked( room_ids = tuple(sorted(candidates)) conn.execute( _RETIRE_FROM_ROOMS.format( - where=f"room_id IN ({placeholders}) AND disbanded_at IS NOT NULL" - ), - room_ids, - ) + where=f"room_id IN ({placeholders}) AND disbanded_at IS NOT NULL"), + room_ids) for table in _DEPENDENT_TABLES: if table_exists(conn, table): conn.execute(f"DELETE FROM {table} WHERE room_id IN ({placeholders})", room_ids) @@ -746,15 +688,12 @@ def list_room_link_records(db_path: Path | str) -> list[dict[str, Any]]: return [dict(row) for row in rows] -def upsert_room_link_record( - db_path: Path | str, *, record: Mapping[str, Any], max_links: int -) -> None: +def upsert_room_link_record(db_path: Path | str, *, record: Mapping[str, Any], max_links: int) -> None: """Atomically insert or replace one private RoomLink record.""" with _transaction(db_path, immediate=True) as conn: existing = conn.execute( "SELECT 1 FROM hosted_room_links WHERE room_id=? AND member_id=?", - (record["room_id"], record["member_id"]), - ).fetchone() + (record["room_id"], record["member_id"])).fetchone() if existing is None: count = int(conn.execute("SELECT COUNT(*) FROM hosted_room_links").fetchone()[0]) if count >= max_links: @@ -775,19 +714,16 @@ def upsert_room_link_record( transport_security=excluded.transport_security, status=excluded.status, updated_at=excluded.updated_at""", - tuple(record[column] for column in _LINK_COLUMNS), - ) + tuple(record[column] for column in _LINK_COLUMNS)) def update_room_link_status( - db_path: Path | str, *, room_id: str, member_id: str, status: str, now: float | None = None -) -> bool: + db_path: Path | str, *, room_id: str, member_id: str, status: str, now: float | None = None) -> bool: """Persist a non-secret route health classification.""" with _transaction(db_path, immediate=True) as conn: cursor = conn.execute( "UPDATE hosted_room_links SET status=?, updated_at=? WHERE room_id=? AND member_id=?", - (status, _now(now), room_id, member_id), - ) + (status, _now(now), room_id, member_id)) return cursor.rowcount == 1 @@ -803,17 +739,14 @@ def _room_grant_scope_key(claims: Mapping[str, Any]) -> str: key: str(claims.get(key) or "") for key in ( "room_id", "home_install_id", "authority_gateway_id", "authority_epoch", - "member_id", "target_install_id", "target_profile", - ) - } + "member_id", "target_install_id", "target_profile")} if not all(fields.values()): raise HostedRoomError("room grant scope is incomplete") return hashlib.sha256(compact_json(fields).encode("utf-8")).hexdigest() def revoke_room_grant_scope( - db_path: Path | str, *, claims: Mapping[str, Any], expires_at: float, now: float | None = None -) -> None: + db_path: Path | str, *, claims: Mapping[str, Any], expires_at: float, now: float | None = None) -> None: """Revoke every grant issued at or before now for one exact room scope.""" scope_key = _room_grant_scope_key(claims) timestamp = _now(now) @@ -831,8 +764,7 @@ def revoke_room_grant_scope( excluded.expires_at), revoked_before=MAX(hosted_room_revoked_grants.revoked_before, excluded.revoked_before)""", - (scope_key, expiry, timestamp), - ) + (scope_key, expiry, timestamp)) conn.execute( """UPDATE hosted_room_peer_reservations SET revoked_at=?, updated_at=? WHERE room_id=? AND member_id=? AND target_profile=? AND authority_gateway_id=? @@ -842,20 +774,16 @@ def revoke_room_grant_scope( str(claims.get("room_id") or ""), str(claims.get("member_id") or ""), str(claims.get("target_profile") or ""), str(claims.get("authority_gateway_id") or ""), - int(claims.get("authority_epoch") or 0), - ), - ) + int(claims.get("authority_epoch") or 0))) def _reservation_claims(claims: Mapping[str, Any]) -> tuple[str, str, str, str, int]: """Validate (room_id, member_id, target_profile, authority_gateway_id, authority_epoch).""" values = ( - _room_id(claims.get("room_id")), - _actor_id(claims.get("member_id"), "member_id"), + _room_id(claims.get("room_id")), _actor_id(claims.get("member_id"), "member_id"), _actor_id(claims.get("target_profile"), "target_profile"), _actor_id(claims.get("authority_gateway_id"), "authority_gateway_id"), - int(claims.get("authority_epoch") or 0), - ) + int(claims.get("authority_epoch") or 0)) if values[4] < 1: raise HostedRoomError("authority_epoch must be positive") return values @@ -864,14 +792,11 @@ def _reservation_claims(claims: Mapping[str, Any]) -> tuple[str, str, str, str, def _reservation_superseded(row: sqlite3.Row, gateway_id: str, epoch: int) -> bool: """A newer epoch, or the same epoch under another gateway, outranks this claim.""" row_epoch = int(row["authority_epoch"]) - return row_epoch > epoch or ( - row_epoch == epoch and str(row["authority_gateway_id"]) != gateway_id - ) + return row_epoch > epoch or (row_epoch == epoch and str(row["authority_gateway_id"]) != gateway_id) def reserve_peer_room( - db_path: Path | str, *, claims: Mapping[str, Any], expires_at: float, now: float | None = None -) -> None: + db_path: Path | str, *, claims: Mapping[str, Any], expires_at: float, now: float | None = None) -> None: """Fence direct Desktop prompts before the first peer run is admitted.""" timestamp = _now(now) expiry = float(expires_at) @@ -884,20 +809,17 @@ def reserve_peer_room( authority_rows = conn.execute( f"""SELECT authority_gateway_id, authority_epoch FROM hosted_room_peer_reservations {_LIVE_RESERVATION_WHERE}""", - (room_id, target_profile, timestamp), - ).fetchall() + (room_id, target_profile, timestamp)).fetchall() if any(_reservation_superseded(row, gateway_id, epoch) for row in authority_rows): raise AuthorityConflictError("peer room reservation authority changed") conn.execute( """UPDATE hosted_room_peer_reservations SET revoked_at=?, updated_at=? WHERE room_id=? AND target_profile=? AND authority_epoch sqlite3.Row | None: @@ -922,8 +843,7 @@ def _read_one(db_path: Path | str, sql: str, params: tuple[Any, ...]) -> sqlite3 def peer_room_is_reserved( - db_path: Path | str, *, room_id: str, target_profile: str, now: float | None = None -) -> bool: + db_path: Path | str, *, room_id: str, target_profile: str, now: float | None = None) -> bool: """Return whether a live target-side RoomLink reservation fences Desktop.""" timestamp = _now(now) params = (_room_id(room_id), _actor_id(target_profile, "target_profile"), timestamp) @@ -931,8 +851,7 @@ def peer_room_is_reserved( def peer_room_grant_is_current( - db_path: Path | str, *, claims: Mapping[str, Any], now: float | None = None -) -> bool: + db_path: Path | str, *, claims: Mapping[str, Any], now: float | None = None) -> bool: """Require a grant to match the target's current live reservation.""" timestamp = _now(now) values = _reservation_claims(claims) @@ -941,14 +860,11 @@ def peer_room_grant_is_current( """SELECT 1 FROM hosted_room_peer_reservations WHERE room_id=? AND member_id=? AND target_profile=? AND authority_gateway_id=? AND authority_epoch=? AND expires_at>? AND revoked_at IS NULL LIMIT 1""", - (*values, timestamp), - ) + (*values, timestamp)) return row is not None -def room_grant_is_revoked( - db_path: Path | str, *, claims: Mapping[str, Any], now: float | None = None -) -> bool: +def room_grant_is_revoked(db_path: Path | str, *, claims: Mapping[str, Any], now: float | None = None) -> bool: """Return whether a grant predates its exact scope's revocation fence.""" timestamp = _now(now) scope_key = _room_grant_scope_key(claims) @@ -957,8 +873,7 @@ def room_grant_is_revoked( db_path, """SELECT revoked_before FROM hosted_room_revoked_grants WHERE scope_key=? AND expires_at>?""", - (scope_key, timestamp), - ) + (scope_key, timestamp)) return row is not None and issued_at <= float(row["revoked_before"]) @@ -967,15 +882,13 @@ def _remote_run_identity(record: Mapping[str, Any]) -> tuple[Any, ...]: def upsert_remote_run_receipt( - db_path: Path | str, *, record: Mapping[str, Any], now: float | None = None -) -> None: + db_path: Path | str, *, record: Mapping[str, Any], now: float | None = None) -> None: """Durably bind one logical peer task attempt to its remote run handle.""" timestamp = _now(now) identity = _remote_run_identity(record) with _transaction(db_path, immediate=True) as conn: existing = conn.execute( - f"SELECT * FROM hosted_room_remote_runs WHERE {_REMOTE_RUN_WHERE}", identity - ).fetchone() + f"SELECT * FROM hosted_room_remote_runs WHERE {_REMOTE_RUN_WHERE}", identity).fetchone() immutable = (*identity, record["run_id"], record["session_id"]) if existing is not None: stored = (*_remote_run_identity(existing), existing["run_id"], existing["session_id"]) @@ -983,8 +896,7 @@ def upsert_remote_run_receipt( raise HostedRoomError("remote run receipt conflicts with its logical task") conn.execute( f"UPDATE hosted_room_remote_runs SET updated_at=? WHERE {_REMOTE_RUN_WHERE}", - (timestamp, *identity), - ) + (timestamp, *identity)) return conn.execute( """INSERT INTO hosted_room_remote_runs( @@ -993,29 +905,24 @@ def upsert_remote_run_receipt( target_profile, task_id, execution_generation, run_id, session_id, created_at, updated_at ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)""", - (*immutable, timestamp, timestamp), - ) + (*immutable, timestamp, timestamp)) def list_remote_run_receipts( db_path: Path | str, *, room_id: str | None = None, target_profile: str | None = None, - session_id: str | None = None, -) -> list[dict[str, Any]]: + session_id: str | None = None) -> list[dict[str, Any]]: """Return remote run handles in durable task order.""" filters = [ (column, value) for column, value in ( - ("room_id", room_id), ("target_profile", target_profile), ("session_id", session_id) - ) - if value is not None - ] + ("room_id", room_id), ("target_profile", target_profile), ("session_id", session_id)) + if value is not None] where = f" WHERE {' AND '.join(f'{column}=?' for column, _ in filters)}" if filters else "" with _transaction(db_path) as conn: rows = conn.execute( "SELECT * FROM hosted_room_remote_runs" + where + " ORDER BY created_at, task_id, execution_generation", - [value for _, value in filters], - ).fetchall() + [value for _, value in filters]).fetchall() return [dict(row) for row in rows] @@ -1024,8 +931,7 @@ def remote_run_receipt(db_path: Path | str, *, record: Mapping[str, Any]) -> dic row = _read_one( db_path, f"SELECT * FROM hosted_room_remote_runs WHERE {_REMOTE_RUN_WHERE}", - _remote_run_identity(record), - ) + _remote_run_identity(record)) return dict(row) if row is not None else None @@ -1034,8 +940,7 @@ def remote_run_receipt(db_path: Path | str, *, record: Mapping[str, Any]) -> dic def _adopt_legacy_room( conn: sqlite3.Connection, existing: sqlite3.Row, *, room_id: str, members_json: str, - authority_gateway_id: str, now: float, -) -> dict[str, Any]: + authority_gateway_id: str, now: float) -> dict[str, Any]: """Claim a 'legacy'-authority room for a real gateway with a fenced claim event.""" target_epoch = int(existing["authority_epoch"]) + 1 seq = int(existing["next_seq"]) @@ -1043,8 +948,7 @@ def _adopt_legacy_room( payload_json = _claim_payload_json("legacy", authority_gateway_id, target_epoch) claim_bytes = _insert_event( conn, existing, room_id, seq, "system:authority-adopted", "authority.claimed", actor_json, - target_epoch, payload_json, now, allow_control=True, - ) + target_epoch, payload_json, now, allow_control=True) adopted = conn.execute( """UPDATE hosted_rooms SET members_json=?, authority_gateway_id=?, authority_epoch=?, @@ -1053,9 +957,7 @@ def _adopt_legacy_room( AND disbanded_at IS NULL""", ( members_json, authority_gateway_id, target_epoch, claim_bytes, now, room_id, - int(existing["authority_epoch"]), seq, - ), - ) + int(existing["authority_epoch"]), seq)) if adopted.rowcount != 1: raise AuthorityConflictError("legacy room adoption lost its fence") existing = _reload(conn, _SELECT_ROOM, (room_id,), "adopted room could not be reloaded") @@ -1063,16 +965,14 @@ def _adopt_legacy_room( result["adopted"] = True claim_event = _reload( conn, _SELECT_EVENT, (room_id, "system:authority-adopted"), - "legacy adoption event could not be reloaded", - ) + "legacy adoption event could not be reloaded") result["claim_event"] = _event_from_row(claim_event) return result def create_room( db_path: Path | str, *, room_id: Any, name: Any, members: Any, authority_gateway_id: Any, - now: float | None = None, -) -> dict[str, Any]: + now: float | None = None) -> dict[str, Any]: """Create a room, or return the identical existing room idempotently.""" room_id = _room_id(room_id) name = _validate_room_name(name) @@ -1082,60 +982,48 @@ def create_room( with _transaction(db_path, immediate=True) as conn: if conn.execute( - "SELECT 1 FROM hosted_room_retired_ids WHERE room_id=?", (room_id,) - ).fetchone(): + "SELECT 1 FROM hosted_room_retired_ids WHERE room_id=?", (room_id,)).fetchone(): raise RoomConflictError("room_id belongs to a disbanded room") existing = conn.execute(_SELECT_ROOM_WITH_BYTES, (room_id,)).fetchone() if existing is not None: if existing["disbanded_at"] is not None: raise RoomConflictError("room_id belongs to a disbanded room") legacy_adoption = ( - existing["authority_gateway_id"] == "legacy" and authority_gateway_id != "legacy" - ) + existing["authority_gateway_id"] == "legacy" and authority_gateway_id != "legacy") members_match = existing["members_json"] == members_json or ( legacy_adoption - and _legacy_members_match(existing["members_json"], normalized_members) - ) + and _legacy_members_match(existing["members_json"], normalized_members)) if existing["name"] != name or not members_match: raise RoomConflictError("room_id already exists with different state") if legacy_adoption: return _adopt_legacy_room( conn, existing, room_id=room_id, members_json=members_json, - authority_gateway_id=authority_gateway_id, now=now, - ) + authority_gateway_id=authority_gateway_id, now=now) if existing["authority_gateway_id"] != authority_gateway_id: raise RoomConflictError("room_id already belongs to a different authority") return _room_from_row(existing, idempotent=True) active_rooms = int( - conn.execute( - "SELECT COUNT(*) FROM hosted_rooms WHERE disbanded_at IS NULL" - ).fetchone()[0] - ) + conn.execute("SELECT COUNT(*) FROM hosted_rooms WHERE disbanded_at IS NULL").fetchone()[0]) if active_rooms >= MAX_ACTIVE_ROOMS: - raise HostedRoomError( - "This host has too many active Group Chats. Delete one and try again." - ) + raise HostedRoomError("This host has too many active Group Chats. Delete one and try again.") conn.execute( f"""INSERT INTO hosted_rooms ({_ROOM_COLUMNS_WITH_BYTES}) VALUES (?, ?, ?, ?, 1, 1, 0, 1, ?, ?, NULL)""", - (room_id, name, members_json, authority_gateway_id, now, now), - ) + (room_id, name, members_json, authority_gateway_id, now, now)) row = _reload( conn, """SELECT room_id, name, members_json, authority_gateway_id, authority_epoch, revision, created_at, updated_at FROM hosted_rooms WHERE room_id=?""", (room_id,), - "created room could not be reloaded", - ) + "created room could not be reloaded") result = _room_from_row(row) result["members"] = normalized_members return result def list_rooms( - db_path: Path | str, *, include_disbanded: bool = False, limit: int = MAX_ROOM_LIST_LIMIT, - offset: int = 0, + db_path: Path | str, *, include_disbanded: bool = False, limit: int = MAX_ROOM_LIST_LIMIT, offset: int = 0, ) -> list[dict[str, Any]]: """Return one bounded read-only page ordered by most recent change.""" limit = _bounded_limit(limit, MAX_ROOM_LIST_LIMIT) @@ -1145,8 +1033,7 @@ def list_rooms( rows = conn.execute( f"""SELECT {_ROOM_COLUMNS} FROM hosted_rooms WHERE disbanded_at IS NULL OR ? ORDER BY updated_at DESC, room_id ASC LIMIT ? OFFSET ?""", - (int(include_disbanded), limit, offset), - ).fetchall() + (int(include_disbanded), limit, offset)).fetchall() finally: conn.close() return [_room_from_row(row) for row in rows] @@ -1183,12 +1070,9 @@ def rename_room( """UPDATE hosted_rooms SET name=?, next_seq=?, event_bytes=event_bytes+?, revision=revision+1, updated_at=? WHERE room_id=?""", - (name, seq + 1, event_bytes, now, room_id), - ) + (name, seq + 1, event_bytes, now, room_id)) conn.execute( - _INSERT_EVENT, - (room_id, seq, event_id, "room.renamed", actor_json, epoch, payload_json, now), - ) + _INSERT_EVENT, (room_id, seq, event_id, "room.renamed", actor_json, epoch, payload_json, now)) updated = conn.execute(_SELECT_ROOM, (room_id,)).fetchone() event = _load_event(conn, room_id, event_id) result = _room_from_row(updated) @@ -1228,33 +1112,28 @@ def append_event( room = conn.execute( """SELECT next_seq, event_bytes, authority_gateway_id, authority_epoch FROM hosted_rooms WHERE room_id=? AND disbanded_at IS NULL""", - (room_id,), - ).fetchone() + (room_id,)).fetchone() if room is None: _raise_room_not_found(conn, room_id) if ( room["authority_gateway_id"] != authority_gateway_id - or int(room["authority_epoch"]) != authority_epoch - ): + or int(room["authority_epoch"]) != authority_epoch): raise AuthorityConflictError("stale hosted room authority") seq = int(room["next_seq"]) event_bytes = _insert_event( conn, room, room_id, seq, event_id, kind, actor_json, authority_epoch, payload_json, now, - allow_control=kind in _CONTROL_EVENT_KINDS, - ) + allow_control=kind in _CONTROL_EVENT_KINDS) advanced = conn.execute( """UPDATE hosted_rooms SET next_seq=?, event_bytes=event_bytes+?, updated_at=? WHERE room_id=? AND next_seq=?""", - (seq + 1, event_bytes, now, room_id, seq), - ) + (seq + 1, event_bytes, now, room_id, seq)) if advanced.rowcount != 1: raise RuntimeError("hosted room sequence advance lost its write fence") row = _reload( conn, f"SELECT {_EVENT_COLUMNS} FROM hosted_room_events WHERE room_id=? AND seq=?", (room_id, seq), - "appended event could not be reloaded", - ) + "appended event could not be reloaded") result = _event_from_row(row) result["actor"] = normalized_actor return result @@ -1268,8 +1147,7 @@ def _probe(path: Path, table: str, query: str, params: tuple[Any, ...], unavaila conn = sqlite3.connect(path, timeout=0.05) try: table_row = conn.execute( - f"SELECT 1 FROM sqlite_master WHERE type='table' AND name='{table}' LIMIT 1" - ).fetchone() + f"SELECT 1 FROM sqlite_master WHERE type='table' AND name='{table}' LIMIT 1").fetchone() if table_row is None: return False return conn.execute(query, params).fetchone() is not None @@ -1290,41 +1168,33 @@ def probe_hosted_room(db_path: Path | str, *, room_id: Any) -> bool: return _probe( Path(db_path), "hosted_rooms", "SELECT 1 FROM hosted_rooms WHERE room_id=? AND disbanded_at IS NULL LIMIT 1", - (checked_room_id,), "hosted room ownership is temporarily unavailable", - ) + (checked_room_id,), "hosted room ownership is temporarily unavailable") def probe_peer_room_reservation( - db_path: Path | str, *, room_id: Any, target_profile: Any, now: float | None = None -) -> bool: + db_path: Path | str, *, room_id: Any, target_profile: Any, now: float | None = None) -> bool: """Check a peer reservation without creating or migrating shared state.""" checked_room_id = _room_id(room_id) checked_profile = _actor_id(target_profile, "target_profile") return _probe( Path(db_path), "hosted_room_peer_reservations", _SELECT_LIVE_RESERVATION, - (checked_room_id, checked_profile, _now(now)), - "peer room ownership is temporarily unavailable", - ) + (checked_room_id, checked_profile, _now(now)), "peer room ownership is temporarily unavailable") -def room_state( - db_path: Path | str, *, room_id: Any, include_disbanded: bool = False -) -> dict[str, Any]: +def room_state(db_path: Path | str, *, room_id: Any, include_disbanded: bool = False) -> dict[str, Any]: """Return durable replay and authority state for one room.""" room_id = _room_id(room_id) with _transaction(db_path) as conn: row = conn.execute( f"""SELECT {_ROOM_COLUMNS} FROM hosted_rooms WHERE room_id=? AND (disbanded_at IS NULL OR ?)""", - (room_id, int(include_disbanded)), - ).fetchone() + (room_id, int(include_disbanded))).fetchone() if row is None: _raise_room_not_found(conn, room_id) claim_row = conn.execute( f"""SELECT {_EVENT_COLUMNS} FROM hosted_room_events WHERE room_id=? AND kind='authority.claimed' AND authority_epoch=? ORDER BY seq DESC LIMIT 1""", - (room_id, int(row["authority_epoch"])), - ).fetchone() + (room_id, int(row["authority_epoch"]))).fetchone() state = _room_from_row(row) if claim_row is not None: state["authority_claim"] = _event_from_row(claim_row) @@ -1332,8 +1202,7 @@ def room_state( def request_room_stop( - db_path: Path | str, *, room_id: Any, cancel_id: Any, expected_gateway_id: Any, - expected_epoch: Any, + db_path: Path | str, *, room_id: Any, cancel_id: Any, expected_gateway_id: Any, expected_epoch: Any, ) -> dict[str, Any]: """Append an idempotent fence that supersedes earlier user turns.""" cancel_id = _validate_identifier(cancel_id, label="cancel_id", max_chars=MAX_EVENT_ID_CHARS) @@ -1341,27 +1210,23 @@ def request_room_stop( return append_event( db_path, room_id=room_id, event_id=f"room-stop:{digest}", kind="room.stop_requested", actor={"kind": "gateway", "id": expected_gateway_id}, payload={"cancel_id": cancel_id}, - authority_gateway_id=expected_gateway_id, authority_epoch=expected_epoch, - ) + authority_gateway_id=expected_gateway_id, authority_epoch=expected_epoch) def _append_authority_claim( conn: sqlite3.Connection, row: sqlite3.Row, *, room_id: str, event_id: str, expected_gateway_id: str, expected_epoch: int, new_gateway_id: str, target_epoch: int, - actor_json: str, payload_json: str, now: float, -) -> sqlite3.Row | None: + actor_json: str, payload_json: str, now: float) -> sqlite3.Row | None: """Insert the claim event and CAS the room's authority; returns the stored claim event.""" claim_bytes = _insert_event( conn, row, room_id, int(row["next_seq"]), event_id, "authority.claimed", actor_json, - target_epoch, payload_json, now, allow_control=True, - ) + target_epoch, payload_json, now, allow_control=True) updated = conn.execute( """UPDATE hosted_rooms SET authority_gateway_id=?, authority_epoch=authority_epoch+1, next_seq=next_seq+1, event_bytes=event_bytes+?, revision=revision+1, updated_at=? WHERE room_id=? AND disbanded_at IS NULL AND authority_gateway_id=? AND authority_epoch=?""", - (new_gateway_id, claim_bytes, now, room_id, expected_gateway_id, expected_epoch), - ) + (new_gateway_id, claim_bytes, now, room_id, expected_gateway_id, expected_epoch)) if updated.rowcount != 1: raise AuthorityConflictError("hosted room authority changed") return _load_event(conn, room_id, event_id) @@ -1369,8 +1234,7 @@ def _append_authority_claim( def claim_authority( db_path: Path | str, *, room_id: Any, expected_gateway_id: Any, expected_epoch: Any, - new_gateway_id: Any, event_id: Any, now: float | None = None, -) -> dict[str, Any]: + new_gateway_id: Any, event_id: Any, now: float | None = None) -> dict[str, Any]: """Fence a verified authority transfer with a compare-and-swap epoch. This storage primitive does not decide *when* takeover is safe. A @@ -1391,8 +1255,7 @@ def claim_authority( row = conn.execute( """SELECT authority_gateway_id, authority_epoch, next_seq, event_bytes FROM hosted_rooms WHERE room_id=? AND disbanded_at IS NULL""", - (room_id,), - ).fetchone() + (room_id,)).fetchone() if row is None: _raise_room_not_found(conn, room_id) current_gateway = str(row["authority_gateway_id"]) @@ -1401,8 +1264,7 @@ def claim_authority( idempotent = existing_event is not None if idempotent: if _event_content(existing_event) != ( - "authority.claimed", claim_actor_json, target_epoch, claim_payload_json - ): + "authority.claimed", claim_actor_json, target_epoch, claim_payload_json): raise EventConflictError("event_id already exists with different content") if current_gateway != new_gateway_id or current_epoch != target_epoch: raise AuthoritySupersededError("authority claim succeeded but was later superseded") @@ -1410,18 +1272,15 @@ def claim_authority( raise AuthorityConflictError("hosted room authority changed") else: existing_event = _append_authority_claim( - conn, row, room_id=room_id, event_id=event_id, - expected_gateway_id=expected_gateway_id, expected_epoch=expected_epoch, - new_gateway_id=new_gateway_id, target_epoch=target_epoch, - actor_json=claim_actor_json, payload_json=claim_payload_json, now=now, - ) + conn, row, room_id=room_id, event_id=event_id, expected_gateway_id=expected_gateway_id, + expected_epoch=expected_epoch, new_gateway_id=new_gateway_id, target_epoch=target_epoch, + actor_json=claim_actor_json, payload_json=claim_payload_json, now=now) state_row = _reload( conn, """SELECT room_id, name, members_json, authority_gateway_id, authority_epoch, next_seq, revision, created_at, updated_at FROM hosted_rooms WHERE room_id=?""", (room_id,), - "claimed room could not be reloaded", - ) + "claimed room could not be reloaded") state = _room_from_row(state_row, idempotent=idempotent) if existing_event is None: # pragma: no cover - both claim paths set it raise RuntimeError("authority claim event could not be reloaded") @@ -1430,19 +1289,16 @@ def claim_authority( def _disband_replay( - conn: sqlite3.Connection, room_id: str, room: sqlite3.Row | None -) -> dict[str, Any] | None: + conn: sqlite3.Connection, room_id: str, room: sqlite3.Row | None) -> dict[str, Any] | None: """Idempotent replay for a retired or already-disbanded room; None when the room is live.""" if room is None: retired = conn.execute( - "SELECT retired_at FROM hosted_room_retired_ids WHERE room_id=?", (room_id,) - ).fetchone() + "SELECT retired_at FROM hosted_room_retired_ids WHERE room_id=?", (room_id,)).fetchone() if retired is None: raise RoomNotFoundError("hosted room not found") return { "room_id": room_id, "disbanded_at": float(retired["retired_at"]), - "idempotent": True, "history_expired": True, - } + "idempotent": True, "history_expired": True} if room["disbanded_at"] is None: return None conn.execute(_INSERT_RETIRED, (room_id, float(room["disbanded_at"]))) @@ -1455,8 +1311,7 @@ def _disband_replay( def disband_room( db_path: Path | str, *, room_id: Any, expected_gateway_id: Any, expected_epoch: Any, - now: float | None = None, -) -> dict[str, Any]: + now: float | None = None) -> dict[str, Any]: """Tombstone a room id permanently and idempotently.""" room_id = _room_id(room_id) expected_gateway_id = _actor_id(expected_gateway_id, "expected_gateway_id") @@ -1467,48 +1322,37 @@ def disband_room( room = conn.execute( """SELECT authority_gateway_id, authority_epoch, next_seq, event_bytes, disbanded_at FROM hosted_rooms WHERE room_id=?""", - (room_id,), - ).fetchone() + (room_id,)).fetchone() replay = _disband_replay(conn, room_id, room) if replay is not None: return replay if ( str(room["authority_gateway_id"]) != expected_gateway_id - or int(room["authority_epoch"]) != expected_epoch - ): + or int(room["authority_epoch"]) != expected_epoch): raise AuthorityConflictError("stale hosted room authority") disband_bytes = _insert_event( conn, room, room_id, int(room["next_seq"]), "system:room-disbanded", "room.disbanded", _system_actor_json("room-control"), int(room["authority_epoch"]), - _payload_json({"room_id": room_id}), now, allow_control=True, - ) + _payload_json({"room_id": room_id}), now, allow_control=True) updated = conn.execute( """UPDATE hosted_rooms SET disbanded_at=?, updated_at=?, revision=revision+1, next_seq=next_seq+1, event_bytes=event_bytes+? WHERE room_id=? AND disbanded_at IS NULL AND authority_gateway_id=? AND authority_epoch=?""", - (now, now, disband_bytes, room_id, expected_gateway_id, expected_epoch), - ) + (now, now, disband_bytes, room_id, expected_gateway_id, expected_epoch)) if updated.rowcount != 1: raise RoomConflictError("hosted room disband lost its fence") conn.execute(_INSERT_RETIRED, (room_id, now)) event = _reload( conn, _SELECT_EVENT, (room_id, "system:room-disbanded"), - "room disband event could not be reloaded", - ) - _prune_disbanded_rooms_locked( - conn, now=now, max_gateway_event_bytes=MAX_GATEWAY_EVENT_BYTES - ) - return { - "room_id": room_id, "disbanded_at": now, "idempotent": False, - "event": _event_from_row(event), - } + "room disband event could not be reloaded") + _prune_disbanded_rooms_locked(conn, now=now, max_gateway_event_bytes=MAX_GATEWAY_EVENT_BYTES) + return {"room_id": room_id, "disbanded_at": now, "idempotent": False, "event": _event_from_row(event), } def read_events( - db_path: Path | str, *, room_id: Any, since_seq: Any = 0, limit: Any = 100, - include_disbanded: bool = False, + db_path: Path | str, *, room_id: Any, since_seq: Any = 0, limit: Any = 100, include_disbanded: bool = False, ) -> dict[str, Any]: """Read a monotonic room-log delta after ``since_seq``.""" room_id = _room_id(room_id) @@ -1519,14 +1363,11 @@ def read_events( room = conn.execute( """SELECT next_seq, authority_gateway_id, authority_epoch FROM hosted_rooms WHERE room_id=? AND (disbanded_at IS NULL OR ?)""", - (room_id, int(include_disbanded)), - ).fetchone() + (room_id, int(include_disbanded))).fetchone() if room is None: _raise_room_not_found(conn, room_id) latest_seq = int(room["next_seq"]) - 1 - authority = { - "gateway_id": str(room["authority_gateway_id"]), "epoch": int(room["authority_epoch"]) - } + authority = {"gateway_id": str(room["authority_gateway_id"]), "epoch": int(room["authority_epoch"])} if since_seq > latest_seq: raise HostedRoomError("since_seq is ahead of the hosted room log") rows = conn.execute( @@ -1546,16 +1387,14 @@ def read_events( FROM candidates WHERE cumulative_bytes<=? ORDER BY seq ASC""", - (room_id, since_seq, limit, MAX_LOG_PAGE_BYTES), - ).fetchall() + (room_id, since_seq, limit, MAX_LOG_PAGE_BYTES)).fetchall() events = [_event_from_row(row) for row in rows] def build_page(page_events: list[dict[str, Any]]) -> dict[str, Any]: cursor = page_events[-1]["seq"] if page_events else since_seq return { "events": page_events, "cursor": cursor, "latest_seq": latest_seq, - "has_more": cursor < latest_seq, "authority": authority, - } + "has_more": cursor < latest_seq, "authority": authority} def page_bytes(page: dict[str, Any]) -> int: return len(json.dumps(page, ensure_ascii=False, separators=(",", ":")).encode("utf-8")) diff --git a/gateway/hosted_rooms_common.py b/gateway/hosted_rooms_common.py index 7fbdee4aef..a02c771666 100644 --- a/gateway/hosted_rooms_common.py +++ b/gateway/hosted_rooms_common.py @@ -23,14 +23,8 @@ IDENTIFIER_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:-]*$") def identifier( - value: Any, - *, - label: str, - error: type[Exception], - max_chars: int = 128, - pattern: re.Pattern[str] | None = IDENTIFIER_RE, - invalid: str | None = None, -) -> str: + value: Any, *, label: str, error: type[Exception], max_chars: int = 128, + pattern: re.Pattern[str] | None = IDENTIFIER_RE, invalid: str | None = None) -> str: """Strip and validate a bounded string; ``pattern=None`` skips the shape check.""" if not isinstance(value, str): raise error(f"{label} must be a string") @@ -40,16 +34,13 @@ def identifier( return value -def bounded_int( - value: Any, *, error: type[Exception], message: str, low: int = 0, high: int | None = None -) -> int: +def bounded_int(value: Any, *, error: type[Exception], message: str, low: int = 0, high: int | None = None) -> int: """Reject bools, non-ints and values outside ``[low, high]`` (``message`` is the exact text).""" if ( isinstance(value, bool) or not isinstance(value, int) or value < low - or (high is not None and value > high) - ): + or (high is not None and value > high)): raise error(message) return value @@ -59,16 +50,10 @@ non_negative_int = partial(bounded_int, low=0) def exact_fields( - value: Any, - *, - label: str, - required: frozenset[str] | set[str], - optional: frozenset[str] | set[str] = frozenset(), - error: type[Exception], - not_object: str | None = None, + value: Any, *, label: str, required: frozenset[str] | set[str], + optional: frozenset[str] | set[str] = frozenset(), error: type[Exception], not_object: str | None = None, missing_fmt: str = "{label} is missing fields: {fields}", - unknown_fmt: str = "{label} has unknown fields: {fields}", -) -> Mapping[str, Any]: + unknown_fmt: str = "{label} has unknown fields: {fields}") -> Mapping[str, Any]: """Require exactly ``required`` (+ any ``optional``) keys; formats name the offenders sorted.""" if not isinstance(value, Mapping): raise error(not_object or f"{label} must be an object") @@ -101,8 +86,7 @@ def compact_json(value: Any, *, ensure_ascii: bool = True) -> str: def canonical_json( - value: Any, *, error: type[Exception], label: str, max_bytes: int, ensure_ascii: bool -) -> str: + value: Any, *, error: type[Exception], label: str, max_bytes: int, ensure_ascii: bool) -> str: """``compact_json`` bounded by ``max_bytes`` of UTF-8; unserializable input raises ``error``.""" try: encoded = compact_json(value, ensure_ascii=ensure_ascii) @@ -131,13 +115,8 @@ def open_sqlite(path: Path | str, *, timeout: float = 10) -> sqlite3.Connection: def connect( - db_path: Path | str, - *, - db_label: str, - ready: Callable[[sqlite3.Connection], bool], - initialize: Callable[[sqlite3.Connection], None], - lock_retries: int = 1, -) -> sqlite3.Connection: + db_path: Path | str, *, db_label: str, ready: Callable[[sqlite3.Connection], bool], + initialize: Callable[[sqlite3.Connection], None], lock_retries: int = 1) -> sqlite3.Connection: """Open the shared root store: WAL, foreign keys, then ``initialize`` in one IMMEDIATE txn if not ``ready``. Multiple profile gateways share this database, so every draft-schema transition @@ -174,9 +153,7 @@ def connect( def table_exists(conn: sqlite3.Connection, table: str) -> bool: - row = conn.execute( - "SELECT 1 FROM sqlite_master WHERE type='table' AND name=?", (table,) - ).fetchone() + row = conn.execute("SELECT 1 FROM sqlite_master WHERE type='table' AND name=?", (table,)).fetchone() return row is not None