fix(scopes): scope observer sessions in SQL instead of a fetched name list (#1065)

`get_observation_context` resolved scope by fetching every session name the
observer has a membership record in, then expanding that list into
`session_name IN (...)` twice in one statement — once in the CTE and once in
the outer select. That puts psycopg's 65535-bind-parameter ceiling at roughly
32,765 sessions, and the count only ever grows: the loose membership
definition (`active_only=False`) counts sessions the peer has since left, so
leaving a session does not shrink the scope. A workspace with tens of
thousands of sessions for one peer produced a statement the driver could not
serialize at all.

Two new helpers in `crud.message` express the observer half as a correlated
EXISTS over `session_peers`. Scope now costs two bind parameters regardless of
membership size, and the membership query disappears (two round trips become
one). The `session_peers` primary key is `(workspace_name, session_name,
peer_name)`, so the correlated probe is an exact-match index hit.

The caller-supplied allowlist stays an IN clause — it is route-capped at 1000
entries and carries none of the unbounded-growth risk. `resolve_session_scope`
is left in place: three other callers still need the materialized list,
including `_search_messages_external`, which sends session names to the vector
store as a filter payload and cannot take SQL.

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Eugene Eisenstein 2026-08-25 13:24:31 -04:00 committed by GitHub
parent 2dbc25093d
commit 2f7658577e
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 273 additions and 7 deletions

View File

@ -6,6 +6,7 @@ from typing import Any
from nanoid import generate as generate_nanoid
from sqlalchemy import ColumnElement, Select, and_, func, or_, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import InstrumentedAttribute
from src import models, schemas
from src.config import settings
@ -159,6 +160,91 @@ async def resolve_session_scope(
return (allowed, False) if allowed else (None, True)
def observer_scope_clause(
workspace_name: str,
observer: str,
session_column: InstrumentedAttribute[str],
) -> ColumnElement[bool]:
"""Correlated EXISTS restricting ``session_column`` to the observer's sessions.
The in-database equivalent of filtering on :func:`get_peer_session_names`.
Prefer it whenever the scope feeds a single SQL statement: a peer's
membership count is unbounded, and materializing the names turns each one
into its own bind parameter. The PostgreSQL wire protocol caps parameters
at 65535 per statement, so a peer in enough sessions produces a query the
driver cannot serialize at all and the resulting error carries every
parameter in its text.
Matches the loose membership definition ``get_peer_session_names`` uses by
default: any membership record grants visibility, whether or not the peer
has since left the session.
"""
session_peers = models.session_peers_table
return (
select(1)
.where(session_peers.c.workspace_name == workspace_name)
.where(session_peers.c.peer_name == observer)
.where(session_peers.c.session_name == session_column)
.exists()
)
def resolve_session_scope_clauses(
workspace_name: str,
session_name: str | None,
session_allowlist: list[str] | None,
observer: str | None,
session_column: InstrumentedAttribute[str],
) -> tuple[list[ColumnElement[bool]], bool]:
"""SQL-side counterpart to :func:`resolve_session_scope`.
Returns ``(clauses, deny)``, where ``clauses`` are ANDed onto the caller's
statement and ``deny=True`` means return an empty result without querying.
Unlike :func:`resolve_session_scope` this touches no database and grows no
bind parameters with the observer's session count — the observer half
becomes a correlated EXISTS instead of an ``IN`` over fetched names.
Scoping matches :func:`resolve_session_scope` case for case, with one
deliberate difference: where that function returns ``deny=True`` because an
observer's membership (or its intersection with the allowlist) is empty,
this returns clauses that simply match no rows. Callers reach the same empty
result, at the cost of running one indexed query that returns nothing.
``session_allowlist`` stays an ``IN`` clause: it is caller-supplied and
therefore bounded, so it carries none of the unbounded-growth risk.
Args:
workspace_name: Name of the workspace
session_name: A single pinned session, if the caller named one. The
caller applies its own equality filter; this function only checks
the allowlist permits it.
session_allowlist: Optional session allowlist. ``None`` is
unrestricted; an empty list fails closed.
observer: When set, scope is limited to this peer's sessions
session_column: The session-name column to scope, e.g.
``models.Message.session_name``
"""
if session_name:
# Fail closed when the allowlist forbids the pinned session, matching
# `resolve_session_scope` — routes guard this too, but the dialectic
# tools reach CRUD directly, so enforce it at the boundary.
if session_allowlist is not None and session_name not in session_allowlist:
return [], True
return [], False
clauses: list[ColumnElement[bool]] = []
if observer is not None:
clauses.append(observer_scope_clause(workspace_name, observer, session_column))
if session_allowlist is not None:
if not session_allowlist:
return [], True
clauses.append(session_column.in_(session_allowlist))
return clauses, False
def _apply_token_limit(
base_conditions: list[ColumnElement[Any]], token_limit: int
) -> Select[tuple[models.Message]]:

View File

@ -1273,10 +1273,19 @@ async def get_observation_context(
if not message_ids:
return []
from src.crud.message import resolve_session_scope
from src.crud.message import resolve_session_scope_clauses
allowed_session_names, deny = await resolve_session_scope(
db, workspace_name, session_name, session_allowlist, observer
# Scope as SQL rather than as a fetched name list. The scope is applied to
# both the CTE and the outer select, so a materialized list would spend two
# bind parameters per session the observer belongs to — enough sessions and
# the statement exceeds the driver's 65535-parameter ceiling and cannot be
# sent at all.
scope_clauses, deny = resolve_session_scope_clauses(
workspace_name,
session_name,
session_allowlist,
observer,
models.Message.session_name,
)
if deny:
return []
@ -1290,8 +1299,8 @@ async def get_observation_context(
if session_name:
stmt = stmt.where(models.Message.session_name == session_name)
elif allowed_session_names is not None:
stmt = stmt.where(models.Message.session_name.in_(allowed_session_names))
for clause in scope_clauses:
stmt = stmt.where(clause)
target_seqs_cte = stmt.cte("target_seqs")
@ -1314,8 +1323,8 @@ async def get_observation_context(
if session_name:
stmt = stmt.where(models.Message.session_name == session_name)
elif allowed_session_names is not None:
stmt = stmt.where(models.Message.session_name.in_(allowed_session_names))
for clause in scope_clauses:
stmt = stmt.where(clause)
result = await db.execute(stmt)
messages = list(result.scalars().all())

View File

@ -86,6 +86,8 @@ _RUNTIME_MOCK_TEST_BLOCKLIST_PREFIXES = (
# LLM transport tests mock providers directly and don't need database/runtime setup.
"tests/utils/test_length_finish_reason.py",
"tests/utils/test_clients.py",
# Session-scope SQL shape — asserts on compiled statements, never executes one.
"tests/crud/test_session_scope_clauses.py",
# Pure JWT scope tests — operate on src.security directly, no DB needed.
"tests/test_security.py",
"tests/test_generate_jwt_script.py",

View File

@ -0,0 +1,169 @@
"""Observer session scope is enforced in SQL, not as a fetched name list.
A peer's session-membership count is unbounded. Materializing it turns every
session name into its own bind parameter, and the PostgreSQL wire protocol caps
parameters at 65535 per statement, so a peer in enough sessions yields a
statement the driver refuses to serialize. `get_observation_context` applies the
scope twice in one statement, which halves that ceiling to ~32.7k sessions.
These tests assert on compiled SQL and never execute a statement.
"""
from typing import Any
import pytest
from sqlalchemy import Select, select
from sqlalchemy.dialects import postgresql
from src import models
from src.crud.message import resolve_session_scope_clauses
from src.utils.agent_tools import get_observation_context
def _compile(stmt: Select[Any]) -> tuple[str, dict[str, Any]]:
compiled = stmt.compile(
dialect=postgresql.dialect(),
compile_kwargs={"render_postcompile": True},
)
return str(compiled), dict(compiled.params)
class _FakeResult:
def scalars(self) -> "_FakeResult":
return self
def all(self) -> list[Any]:
return []
class _CapturingDB:
"""Captures statements instead of executing them."""
def __init__(self) -> None:
self.statements: list[Any] = []
async def execute(self, stmt: Any) -> _FakeResult:
self.statements.append(stmt)
return _FakeResult()
@pytest.mark.parametrize(
("session_allowlist", "observer", "expect_exists", "expected_params"),
[
# Observer only: membership becomes a correlated EXISTS, so only the
# workspace and peer are bound — never the session names, which is what
# keeps the parameter count from growing with membership.
(None, "observer-peer", True, ["observer-peer", "workspace"]),
# Allowlist only: caller-supplied and therefore bounded, so IN is fine.
(["s1", "s2"], None, False, ["s1", "s2"]),
# Both: the EXISTS is intersected with the bounded IN.
(["s1"], "observer-peer", True, ["observer-peer", "s1", "workspace"]),
# Neither: unrestricted, nothing filtered and nothing bound.
(None, None, False, []),
],
ids=["observer-only", "allowlist-only", "observer-and-allowlist", "unrestricted"],
)
def test_scope_clause_shape(
session_allowlist: list[str] | None,
observer: str | None,
expect_exists: bool,
expected_params: list[str],
) -> None:
clauses, deny = resolve_session_scope_clauses(
"workspace",
None,
session_allowlist,
observer,
models.Message.session_name,
)
assert not deny
sql, params = _compile(select(models.Message.public_id).where(*clauses))
assert ("EXISTS" in sql.upper()) is expect_exists
assert ("session_peers" in sql) is expect_exists
assert sorted(params.values()) == expected_params
@pytest.mark.parametrize(
("session_name", "session_allowlist", "observer"),
[
# An empty allowlist fails closed rather than matching everything.
(None, [], "observer-peer"),
(None, [], None),
# A pinned session the allowlist forbids fails closed.
("s9", ["s1", "s2"], "observer-peer"),
],
ids=[
"empty-allowlist-with-observer",
"empty-allowlist-without-observer",
"pinned-session-not-in-allowlist",
],
)
def test_scope_fails_closed(
session_name: str | None,
session_allowlist: list[str] | None,
observer: str | None,
) -> None:
clauses, deny = resolve_session_scope_clauses(
"workspace",
session_name,
session_allowlist,
observer,
models.Message.session_name,
)
assert deny
assert clauses == []
@pytest.mark.asyncio
async def test_get_observation_context_scope_costs_no_per_session_parameters() -> None:
"""The statement's parameter count depends on message_ids, not membership."""
db = _CapturingDB()
message_ids = [f"msg-{i}" for i in range(5)]
await get_observation_context(
db, # pyright: ignore[reportArgumentType]
"workspace",
None,
message_ids,
observer="observer-peer",
)
assert len(db.statements) == 1
sql, params = _compile(db.statements[0])
# The scope must be applied to *both* the CTE and the outer select — a
# materialized list would have cost two parameters per session there.
# Count the subquery's FROM rather than EXISTS: the adjacency check is also
# an EXISTS, so counting those would pass with the CTE's scope missing.
assert sql.count("FROM public.session_peers") == 2
# Both must correlate to the enclosing `messages` row. An uncorrelated
# subquery compiles just as happily and would silently drop scoping
# instead of enforcing it.
assert sql.count("session_peers.session_name = public.messages.session_name") == 2
# Everything bound is either a message id, the workspace, the peer, or the
# ±1 adjacency window — nothing that scales with the peer's session count.
expected = set(message_ids) | {"workspace", "observer-peer", -1, 1}
assert set(params.values()) <= expected
@pytest.mark.asyncio
async def test_get_observation_context_denies_without_querying() -> None:
"""Fail-closed scopes must not reach the database at all."""
db = _CapturingDB()
result = await get_observation_context(
db, # pyright: ignore[reportArgumentType]
"workspace",
None,
["msg-1"],
observer="observer-peer",
session_allowlist=[],
)
assert result == []
assert db.statements == []