refactor(gateway): AST-neutral repack of hosted-room call/def layouts (<=110 cols)

This commit is contained in:
Teknium
2026-09-02 18:29:57 -07:00
parent 1d185f1a37
commit efff31ccd0
9 changed files with 416 additions and 899 deletions
+84 -181
View File
@@ -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<source>[1-9][0-9]*)\.r(?P<round>[0-2])\."
r"p(?P<position>[0-5])\.s(?P<seen>[1-9][0-9]*)\."
r"m(?P<member>[0-9a-f]{24})$"
)
r"m(?P<member>[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))
+70 -143
View File
@@ -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))
+6 -12
View File
@@ -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"]
+9 -24
View File
@@ -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()})
+28 -57
View File
@@ -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
+46 -98
View File
@@ -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<?",
(event["room_id"], thread_id, int(cutoff["seq"])),
)
(event["room_id"], thread_id, int(cutoff["seq"])))
def _backfill_transcript(self, conn: sqlite3.Connection, *, room_id: str, through_seq: int) -> 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,))
+23 -50
View File
@@ -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}
+139 -300
View File
File diff suppressed because it is too large Load Diff
+11 -34
View File
@@ -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