From 96c104c9035872786ee32f48c77e20103dd0cfe0 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Thu, 3 Sep 2026 09:32:40 -0700 Subject: [PATCH] review-fix(public-api): restore from_platform_entry, SessionManager.remove_session/cleanup, is_relay_media_url, SessionTurnLeaseRegistry.__len__ + tests BASE exposed these public names; the simplify refactor dropped them (and deleted/removed their tests) although plugins/connectors import them. Restore each with BASE's signature and body (download() again routes its auth decision through is_relay_media_url), and restore the covering tests ported to the new layout, plus DB-only/task-cwd coverage for remove_session/cleanup, the zero->4096 chunking normalization for from_platform_entry, and len() on empty/populated registries. A/B vs 63279301bcb: /tmp/rf/rev/ab_publicapi.py identical output on both trees. --- acp_adapter/session.py | 51 +++++++++++++++++ gateway/relay/descriptor.py | 41 +++++++++++++ gateway/relay/media.py | 6 +- gateway/turn_lease.py | 3 + tests/acp/test_session.py | 45 +++++++++++++++ .../relay/test_descriptor_from_entry.py | 53 +++++++++++++++++ tests/gateway/relay/test_relay_media.py | 57 +++++++++++++++++++ tests/gateway/test_turn_lease.py | 18 ++++++ 8 files changed, 273 insertions(+), 1 deletion(-) create mode 100644 tests/gateway/relay/test_descriptor_from_entry.py diff --git a/acp_adapter/session.py b/acp_adapter/session.py index 167796e21e..5ba60a984c 100644 --- a/acp_adapter/session.py +++ b/acp_adapter/session.py @@ -95,6 +95,17 @@ def _register_task_cwd(task_id: str, cwd: str) -> None: logger.debug("Failed to register ACP task cwd override", exc_info=True) +def _clear_task_cwd(task_id: str) -> None: + """Remove task-specific cwd overrides for an ACP session.""" + if not task_id: + return + try: + from tools.terminal_tool import clear_task_env_overrides + clear_task_env_overrides(task_id) + except Exception: + logger.debug("Failed to clear ACP task cwd override", exc_info=True) + + def _expand_acp_enabled_toolsets(toolsets: List[str] | None = None, mcp_server_names: List[str] | None = None) -> List[str]: """Return ACP toolsets plus explicit MCP server toolsets for this session.""" @@ -172,6 +183,15 @@ class SessionManager: state = self._sessions.get(session_id) return state if state is not None else self._restore(session_id) + def remove_session(self, session_id: str) -> bool: + """Remove a session from memory and database. Returns True if it existed.""" + with self._lock: + existed = self._sessions.pop(session_id, None) is not None + db_existed = self._delete_persisted(session_id) + if existed or db_existed: + _clear_task_cwd(session_id) + return existed or db_existed + def fork_session(self, session_id: str, cwd: str = ".") -> Optional[SessionState]: """Deep-copy a session's history into a new session.""" cwd = _translate_acp_cwd(cwd) @@ -238,6 +258,26 @@ class SessionManager: self._persist(state) return state + def cleanup(self) -> None: + """Remove all sessions (memory and database) and clear task-specific cwd overrides.""" + with self._lock: + session_ids = list(self._sessions.keys()) + self._sessions.clear() + for session_id in session_ids: + _clear_task_cwd(session_id) + self._delete_persisted(session_id) + # Also remove any DB-only ACP sessions not currently in memory. + db = self._get_db() + if db is not None: + try: + rows = db.search_sessions(source="acp", limit=10000) + for row in rows: + sid = row["id"] + _clear_task_cwd(sid) + db.delete_session(sid) + except Exception: + logger.debug("Failed to cleanup ACP sessions from DB", exc_info=True) + def save_session(self, session_id: str) -> None: """Persist a session; called by the server after prompt completion, history-mutating slash commands, and model switches.""" @@ -349,6 +389,17 @@ class SessionManager: logger.info("Restored ACP session %s from DB (%d messages)", session_id, len(history)) return state + def _delete_persisted(self, session_id: str) -> bool: + """Delete a session from the database. Returns True if it existed.""" + db = self._get_db() + if db is None: + return False + try: + return db.delete_session(session_id) + except Exception: + logger.debug("Failed to delete ACP session %s from DB", session_id, exc_info=True) + return False + # ---- internal ----------------------------------------------------------- def _make_agent(self, *, session_id: str, cwd: str, model: str | None = None, diff --git a/gateway/relay/descriptor.py b/gateway/relay/descriptor.py index 92e9afad34..a6793a4144 100644 --- a/gateway/relay/descriptor.py +++ b/gateway/relay/descriptor.py @@ -93,3 +93,44 @@ class CapabilityDescriptor: else () ) return cls(**filtered) + + @classmethod + def from_platform_entry( + cls, + entry, + *, + len_unit: str = "chars", + supports_draft_streaming: bool = False, + supports_edit: bool = True, + supports_threads: bool = False, + markdown_dialect: str = "plain", + ) -> "CapabilityDescriptor": + """Project a ``gateway.platform_registry.PlatformEntry`` into a descriptor. + + Demonstrates the descriptor is a *subset/projection* of what + ``PlatformEntry`` already encodes, not a parallel concept: ``label``, + ``max_message_length``, ``emoji``, ``platform_hint``, ``pii_safe`` and + the platform name come straight off the entry. The runtime capability + bits that ``PlatformEntry`` does NOT encode (length unit, draft/edit/ + thread/markdown behavior) are supplied by the caller — in production + the connector fills these from the live adapter's capability methods. + + ``max_message_length`` of 0 on a ``PlatformEntry`` means "no limit"; + we map that to the stream_consumer default of 4096 so the descriptor + always carries a concrete chunking bound. + """ + max_len = getattr(entry, "max_message_length", 0) or 4096 + return cls( + contract_version=CONTRACT_VERSION, + platform=entry.name, + label=entry.label, + max_message_length=max_len, + supports_draft_streaming=supports_draft_streaming, + supports_edit=supports_edit, + supports_threads=supports_threads, + markdown_dialect=markdown_dialect, + len_unit=len_unit, + emoji=getattr(entry, "emoji", "\U0001f50c"), + platform_hint=getattr(entry, "platform_hint", ""), + pii_safe=getattr(entry, "pii_safe", False), + ) diff --git a/gateway/relay/media.py b/gateway/relay/media.py index 468ec82c1e..4b957f4165 100644 --- a/gateway/relay/media.py +++ b/gateway/relay/media.py @@ -65,6 +65,10 @@ class RelayMediaClient: def _bearer(self) -> str: return make_upgrade_token(self._gateway_id, self._secret) + def is_relay_media_url(self, url: str) -> bool: + """Is ``url`` a connector re-host reference (needs our bearer to GET)?""" + return "/relay/media/" in (url or "") + async def upload( self, file_path: str, *, mime: Optional[str] = None, filename: Optional[str] = None ) -> Optional[str]: @@ -113,7 +117,7 @@ class RelayMediaClient: """ if not url: return None - needs_auth = "/relay/media/" in url + needs_auth = self.is_relay_media_url(url) if needs_auth and not self.enabled: return None headers = {"User-Agent": _MEDIA_USER_AGENT} diff --git a/gateway/turn_lease.py b/gateway/turn_lease.py index 1f65614e66..687665b312 100644 --- a/gateway/turn_lease.py +++ b/gateway/turn_lease.py @@ -76,6 +76,9 @@ class SessionTurnLeaseRegistry: self._leases: Dict[str, _SessionLease] = {} self._max_entries = max(1, int(max_entries)) + def __len__(self) -> int: + return len(self._leases) + def _get_or_create(self, session_id: str) -> _SessionLease: if (lease := self._leases.get(session_id)) is None: self._evict_idle() diff --git a/tests/acp/test_session.py b/tests/acp/test_session.py index dc6180fbea..d0bd735e71 100644 --- a/tests/acp/test_session.py +++ b/tests/acp/test_session.py @@ -240,6 +240,51 @@ class TestListAndCleanup: assert messages[0]["content"] == "original" assert isinstance(messages[0].get("timestamp"), (int, float)) + def test_cleanup_clears_all(self, manager): + s1 = manager.create_session() + s2 = manager.create_session() + s1.history.append({"role": "user", "content": "one"}) + s2.history.append({"role": "user", "content": "two"}) + assert len(manager.list_sessions()) == 2 + manager.cleanup() + assert manager.list_sessions() == [] + + def test_cleanup_removes_db_only_sessions_and_clears_task_cwd(self, manager, monkeypatch): + """cleanup() must also purge ACP sessions that live only in the DB (e.g. + left over from a previous process) and drop every task-cwd override.""" + cleared: list[str] = [] + monkeypatch.setattr("tools.terminal_tool.clear_task_env_overrides", cleared.append) + live = manager.create_session() + db = manager._get_db() + db.create_session(session_id="acp-db-only", source="acp", model="test") + db.create_session(session_id="cli-keep", source="cli", model="test") + manager.cleanup() + assert manager._sessions == {} + assert db.get_session(live.session_id) is None + assert db.get_session("acp-db-only") is None + assert db.get_session("cli-keep") is not None # non-ACP sessions untouched + assert set(cleared) == {live.session_id, "acp-db-only"} + + def test_remove_session(self, manager): + state = manager.create_session() + assert manager.remove_session(state.session_id) is True + assert manager.get_session(state.session_id) is None + # Removing again returns False + assert manager.remove_session(state.session_id) is False + + def test_remove_session_db_only_and_task_cwd(self, manager, monkeypatch): + """remove_session() handles a DB-only session (True, row deleted, cwd + override cleared) and leaves overrides alone for unknown ids (False).""" + cleared: list[str] = [] + monkeypatch.setattr("tools.terminal_tool.clear_task_env_overrides", cleared.append) + db = manager._get_db() + db.create_session(session_id="acp-db-only", source="acp", model="test") + assert manager.remove_session("acp-db-only") is True + assert db.get_session("acp-db-only") is None + assert cleared == ["acp-db-only"] + assert manager.remove_session("never-existed") is False + assert cleared == ["acp-db-only"] + # --------------------------------------------------------------------------- # persistence — sessions survive process restarts (via SessionDB) diff --git a/tests/gateway/relay/test_descriptor_from_entry.py b/tests/gateway/relay/test_descriptor_from_entry.py new file mode 100644 index 0000000000..5261a470fe --- /dev/null +++ b/tests/gateway/relay/test_descriptor_from_entry.py @@ -0,0 +1,53 @@ +"""Descriptor <- PlatformEntry projection (relay Phase 0, Task 0.3). + +Proves the CapabilityDescriptor is a projection of the existing PlatformEntry, +not a parallel concept: the entry's label/limit/emoji/hint/pii fields carry +straight through. +""" + +from gateway.platform_registry import PlatformEntry +from gateway.relay.descriptor import CONTRACT_VERSION, CapabilityDescriptor + + +def _entry(**overrides) -> PlatformEntry: + base = dict( + name="telegram", + label="Telegram", + adapter_factory=lambda cfg: None, + check_fn=lambda: True, + max_message_length=4096, + pii_safe=False, + emoji="\u2708\ufe0f", + platform_hint="You are on Telegram.", + ) + base.update(overrides) + return PlatformEntry(**base) + + +def test_projection_carries_platform_entry_fields(): + d = CapabilityDescriptor.from_platform_entry(_entry(), len_unit="utf16") + assert d.contract_version == CONTRACT_VERSION + assert d.platform == "telegram" + assert d.label == "Telegram" + assert d.max_message_length == 4096 + assert d.emoji == "\u2708\ufe0f" + assert d.platform_hint == "You are on Telegram." + assert d.pii_safe is False + assert d.len_unit == "utf16" + + +def test_projection_defaults_for_runtime_bits(): + # Bits PlatformEntry does not encode take the documented defaults. + d = CapabilityDescriptor.from_platform_entry(_entry()) + assert d.len_unit == "chars" + assert d.supports_draft_streaming is False + assert d.supports_edit is True + assert d.supports_threads is False + assert d.markdown_dialect == "plain" + + +def test_zero_unlimited_length_normalizes_to_chunking_default(): + # PlatformEntry uses 0 for "no limit"; the descriptor must always carry a + # concrete chunking bound (the stream_consumer default of 4096). + d = CapabilityDescriptor.from_platform_entry(_entry(max_message_length=0)) + assert d.max_message_length == 4096 diff --git a/tests/gateway/relay/test_relay_media.py b/tests/gateway/relay/test_relay_media.py index 22ff3cfd2f..05b8fdd4db 100644 --- a/tests/gateway/relay/test_relay_media.py +++ b/tests/gateway/relay/test_relay_media.py @@ -229,3 +229,60 @@ async def test_download_sends_a_user_agent_on_every_request(): # The re-host request must still carry its bearer (no regression). rehost_headers = seen[1] assert (rehost_headers.get("Authorization") or "").startswith("Bearer ") + + +def test_is_relay_media_url_distinguishes_rehost_from_public(): + """Public compat helper: connector re-host refs need our bearer; ordinary + public URLs (CDN pass-throughs) do not. None/empty never raise.""" + c = RelayMediaClient("https://conn.example", "gw1", "sec") + assert c.is_relay_media_url("https://conn.example/relay/media/aa11") is True + assert c.is_relay_media_url("http://other.host:8080/relay/media/x") is True + assert c.is_relay_media_url("https://cdn.discordapp.com/attachments/1/2/v.ogg") is False + assert c.is_relay_media_url("https://conn.example/relay/mediafile") is False + assert c.is_relay_media_url("") is False + assert c.is_relay_media_url(None) is False # type: ignore[arg-type] + + +@pytest.mark.asyncio +async def test_download_routes_auth_decision_through_is_relay_media_url(monkeypatch): + """download() must consult is_relay_media_url (so subclasses/plugins that + override the classifier change auth behaviour): a disabled client refuses + re-host refs but still fetches public URLs without a bearer.""" + asked: list[str] = [] + c = RelayMediaClient("https://conn.example", None, None) # disabled: no creds + assert c.enabled is False + orig = c.is_relay_media_url + + def _spy(url): + asked.append(url) + return orig(url) + + monkeypatch.setattr(c, "is_relay_media_url", _spy) + # Re-host ref + disabled client → None before any network call. + assert await c.download("https://conn.example/relay/media/deadbeef") is None + assert asked == ["https://conn.example/relay/media/deadbeef"] + + seen: list[dict] = [] + + class _Resp: + headers = {"Content-Type": "image/png", "Content-Length": "4"} + + def read(self, *_a): + return b"\x89PNG" + + def __enter__(self): + return self + + def __exit__(self, *_a): + return False + + def _fake_urlopen(req, timeout=None): # noqa: ARG001 + seen.append(dict(req.headers)) + return _Resp() + + import urllib.request as _ur + + monkeypatch.setattr(_ur, "urlopen", _fake_urlopen) + assert await c.download("https://cdn.discordapp.com/attachments/1/2/i.png") + assert asked[-1] == "https://cdn.discordapp.com/attachments/1/2/i.png" + assert len(seen) == 1 and "Authorization" not in seen[0] diff --git a/tests/gateway/test_turn_lease.py b/tests/gateway/test_turn_lease.py index 1295720c1f..aa6bc6a183 100644 --- a/tests/gateway/test_turn_lease.py +++ b/tests/gateway/test_turn_lease.py @@ -492,3 +492,21 @@ def test_runner_release_turn_lease_is_token_scoped_and_bare_safe(): assert runner._release_turn_lease("", 1) is False _run(scenario()) + + +def test_registry_len_reports_tracked_sessions(): + """``len(registry)`` is public API (plugins/diagnostics size the registry + with it): 0 when empty, one per tracked session_id, and it follows eviction.""" + async def scenario(): + registry = SessionTurnLeaseRegistry(max_entries=8) + assert len(registry) == 0 + a = await registry.acquire("s1", owner_key="k1", generation=1, timeout=1) + assert len(registry) == 1 + b = await registry.acquire("s2", owner_key="k2", generation=1, timeout=1) + assert len(registry) == 2 + registry.release(a) + registry.release(b) + # Released (idle) entries stay tracked until eviction, matching _leases. + assert len(registry) == len(registry._leases) == 2 + + _run(scenario())