refactor(gateway): AST-neutral repack of hosted-room call/def layouts (<=110 cols)
This commit is contained in:
@@ -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
@@ -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))
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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,))
|
||||
|
||||
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user