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:
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
)
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
@@ -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]
|
||||
|
||||
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user