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 63279301bc: /tmp/rf/rev/ab_publicapi.py identical output on both trees.
This commit is contained in:
Teknium
2026-09-03 09:32:40 -07:00
parent 58dca1be51
commit 96c104c903
8 changed files with 273 additions and 1 deletions
+51
View File
@@ -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,
+41
View File
@@ -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),
)
+5 -1
View File
@@ -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}
+3
View File
@@ -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()
+45
View File
@@ -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)
@@ -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
+57
View File
@@ -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]
+18
View File
@@ -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())