Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
40 changes: 40 additions & 0 deletions loopx/chat.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,46 @@ def __init__(self, message: str, *, receipt: dict[str, Any]) -> None:
self.receipt = receipt


def require_matching_replay(
existing: Mapping[str, Any], *, identity: str, request: Mapping[str, Any]
) -> None:
if any(existing.get(field) != value for field, value in request.items()):
raise ValueError(f"{identity} already belongs to a different request")


def resolve_attached_completion_replay(
messages: Iterable[Mapping[str, Any]],
*,
turn_id: str,
completion_id: str,
response_message: str,
) -> str | None:
"""Return the canonical message id to append, or None for a valid replay."""

rows = list(messages)
turn_rows = [
row
for row in rows
if row.get("role") == "agent" and row.get("turn_id") == turn_id
]
if len(turn_rows) > 1:
raise ValueError("attached completion transcript identity is ambiguous")
current_id = f"attached.{turn_id}.completed"
if not turn_rows:
if any(row.get("message_id") == current_id for row in rows):
raise ValueError("attached completion transcript identity conflicts")
return current_id
row = turn_rows[0]
if row.get("origin") != "attached_host" or row.get("message_id") not in {
current_id,
f"attached.{completion_id}",
}:
raise ValueError("attached completion transcript identity conflicts")
if row.get("text") != response_message:
raise ValueError("attached completion transcript conflicts with response")
return None


def _stable_digest(payload: dict[str, Any], *, length: int = 24) -> str:
stable = json.dumps(payload, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
return hashlib.sha256(stable.encode("utf-8")).hexdigest()[:length]
Expand Down
30 changes: 13 additions & 17 deletions loopx/chat_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from typing import Any
import uuid

from .chat import require_matching_replay, resolve_attached_completion_replay
from .file_lock import exclusive_file_lock


Expand Down Expand Up @@ -87,13 +88,6 @@ def _session_channel(payload: dict[str, Any]) -> str:
return f"goal.{payload.get('goal_id')}"


def _require_matching_replay(
existing: dict[str, Any], *, identity: str, request: dict[str, Any]
) -> None:
if any(existing.get(field) != value for field, value in request.items()):
raise ValueError(f"{identity} already belongs to a different request")


def _atomic_write_json(path: Path, payload: dict[str, Any], *, preserve_mode: bool = False) -> None:
previous_mode = path.stat().st_mode & 0o777 if preserve_mode and path.exists() else 0o600
temporary = path.with_name(f".{path.name}.{uuid.uuid4().hex}.tmp")
Expand Down Expand Up @@ -581,7 +575,7 @@ def create_ingress_receipt(
):
existing = _read_json(path)
if existing.get("schema_version") == CHAT_INGRESS_SCHEMA_VERSION:
_require_matching_replay(
require_matching_replay(
existing,
identity="client_ingress_id",
request={"mode": _opaque_id(mode, field="mode"), "message": str(message)},
Expand Down Expand Up @@ -660,7 +654,7 @@ def create_turn(
)
if original is None:
raise ValueError("client_turn_id original request is unavailable")
_require_matching_replay(
require_matching_replay(
{**existing, "attachments": original.get("attachments") or None},
identity="client_turn_id",
request={
Expand Down Expand Up @@ -746,7 +740,7 @@ def create_queued_turn(
):
existing = self.turn_for_client(session_id, client_id)
if existing is not None:
_require_matching_replay(
require_matching_replay(
existing,
identity="client_turn_id",
request={
Expand Down Expand Up @@ -1127,14 +1121,16 @@ def finalize_attached_turn_completion(
field="completion_id",
)
completed_at = str(turn.get("completed_at") or utc_now())
self.append_message(
session_id,
role="agent",
text=str(response.get("message") or ""),
turn_id=turn_id,
origin="attached_host",
message_id=f"attached.{completion_id}",
message_id = resolve_attached_completion_replay(
self.messages(session_id), turn_id=turn_id,
completion_id=completion_id,
response_message=str(response.get("message") or ""),
)
if message_id:
self.append_message(
session_id, role="agent", text=str(response.get("message") or ""),
turn_id=turn_id, origin="attached_host", message_id=message_id,
)
self.append_completed_response_events(
session_id,
turn_id,
Expand Down
159 changes: 159 additions & 0 deletions tests/test_attached_session_broker.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,6 +180,80 @@ def complete(_index: int) -> dict[str, object]:
assert len(agent_messages) == 1


def test_completion_id_reuse_across_turns_preserves_each_response(
tmp_path: Path,
) -> None:
store = ChatSessionStore(tmp_path)
session_id = str(_bind(store)["session"]["session_id"])

for index, response in enumerate(("first response", "second response"), 1):
queued, created = store.create_queued_turn(
session_id,
client_turn_id=f"completion-scope-{index}",
message=f"request {index}",
origin="web",
)
assert created
claim_attached_agent_turn(
store=store,
session_id=session_id,
host_surface=HOST_SURFACE,
host_session_id=HOST_SESSION_ID,
claim_id=f"completion-scope-claim-{index}",
)
complete_attached_agent_turn(
store=store,
session_id=session_id,
turn_id=str(queued["turn_id"]),
host_surface=HOST_SURFACE,
host_session_id=HOST_SESSION_ID,
claim_id=f"completion-scope-claim-{index}",
completion_id="host-local-completion",
response={"message": response},
)

agent_messages = [
item for item in store.messages(session_id) if item["role"] == "agent"
]
assert [item["text"] for item in agent_messages] == [
"first response",
"second response",
]
assert agent_messages[0]["turn_id"] != agent_messages[1]["turn_id"]


def test_maximum_length_completion_id_does_not_strand_turn(tmp_path: Path) -> None:
store = ChatSessionStore(tmp_path)
session_id = str(_bind(store)["session"]["session_id"])
queued, _created = store.create_queued_turn(
session_id,
client_turn_id="maximum-completion-id",
message="complete with a valid opaque id",
origin="web",
)
claim_attached_agent_turn(
store=store,
session_id=session_id,
host_surface=HOST_SURFACE,
host_session_id=HOST_SESSION_ID,
claim_id="maximum-completion-id-claim",
)

completed = complete_attached_agent_turn(
store=store,
session_id=session_id,
turn_id=str(queued["turn_id"]),
host_surface=HOST_SURFACE,
host_session_id=HOST_SESSION_ID,
claim_id="maximum-completion-id-claim",
completion_id="c" * 160,
response={"message": "durable response"},
)

assert completed["completed"] is True
assert store.load_turn(session_id, str(queued["turn_id"]))["status"] == "completed" # type: ignore[index]


def test_attached_host_can_wait_for_queue_wakeup(tmp_path: Path) -> None:
store = ChatSessionStore(tmp_path)
session_id = str(_bind(store)["session"]["session_id"])
Expand Down Expand Up @@ -562,12 +636,19 @@ def test_attached_completion_replays_closeout_after_restart(
host_session_id=HOST_SESSION_ID,
claim_id="attached-recovery-claim",
)
append_message = store.append_message
append_events = store.append_completed_response_events

def append_legacy_message(*args: object, **kwargs: object) -> dict[str, object]:
if kwargs.get("role") == "agent":
kwargs["message_id"] = "attached.attached-recovery-completion"
return append_message(*args, **kwargs) # type: ignore[arg-type]

def crash_after_events(*args: object, **kwargs: object) -> None:
append_events(*args, **kwargs) # type: ignore[arg-type]
raise RuntimeError("crash after attached completion events")

monkeypatch.setattr(store, "append_message", append_legacy_message)
monkeypatch.setattr(store, "append_completed_response_events", crash_after_events)
with pytest.raises(RuntimeError, match="crash after attached completion events"):
complete_attached_agent_turn(
Expand Down Expand Up @@ -600,6 +681,84 @@ def crash_after_events(*args: object, **kwargs: object) -> None:
] == ["turn.completed"]


@pytest.mark.parametrize(
"message_specs",
[
[("reserved attached response", "other_runtime", "other")],
[("different runtime response", "other_runtime", "other")],
[("reserved attached response", "attached_host", "other")],
[
("reserved attached response", "attached_host", "legacy"),
("reserved attached response", "attached_host", "current"),
],
],
ids=["equal-wrong-origin", "different-wrong-origin", "wrong-id", "multiple"],
)
def test_attached_completion_recovery_rejects_ambiguous_agent_messages(
tmp_path: Path,
monkeypatch: pytest.MonkeyPatch,
message_specs: list[tuple[str, str, str]],
) -> None:
store = ChatSessionStore(tmp_path)
session_id = str(_bind(store)["session"]["session_id"])
turn, _created = store.create_queued_turn(
session_id,
client_turn_id="ambiguous-completion-transcript",
message="preserve attached completion provenance",
origin="web",
)
turn_id = str(turn["turn_id"])
claim_attached_agent_turn(
store=store,
session_id=session_id,
host_surface=HOST_SURFACE,
host_session_id=HOST_SESSION_ID,
claim_id="wrong-origin-claim",
)
finalize = store.finalize_attached_turn_completion

def crash_before_closeout(*_args: object, **_kwargs: object) -> None:
raise RuntimeError("crash before attached completion closeout")

monkeypatch.setattr(store, "finalize_attached_turn_completion", crash_before_closeout)
with pytest.raises(RuntimeError, match="crash before attached completion closeout"):
complete_attached_agent_turn(
store=store,
session_id=session_id,
turn_id=turn_id,
host_surface=HOST_SURFACE,
host_session_id=HOST_SESSION_ID,
claim_id="wrong-origin-claim",
completion_id="wrong-origin-completion",
response={"message": "reserved attached response"},
)
monkeypatch.setattr(store, "finalize_attached_turn_completion", finalize)
message_ids = {
"other": f"other-runtime-{turn_id}",
"legacy": "attached.wrong-origin-completion",
"current": f"attached.{turn_id}.completed",
}
for text, origin, identity in message_specs:
store.append_message(
session_id,
role="agent",
text=text,
turn_id=turn_id,
origin=origin,
message_id=message_ids[identity],
)

with pytest.raises(ValueError, match="attached completion transcript identity"):
store.finalize_attached_turn_completion(session_id, turn_id)

assert store.load_turn(session_id, turn_id)["status"] == "completing" # type: ignore[index]
assert store.load_session(session_id)["active_turn_id"] == turn_id # type: ignore[index]
assert not any(
event["kind"] == "turn.completed"
for event in store.events_after(session_id, turn_id, None)
)


def test_chat_events_reconcile_writes_from_another_store_instance(tmp_path: Path) -> None:
server_store = ChatSessionStore(tmp_path)
session_id = str(_bind(server_store)["session"]["session_id"])
Expand Down