Files
odysseus/tests/test_chatgpt_session_binding.py
T

163 lines
6.9 KiB
Python

"""Persistent account routing must survive reloads and fail closed on deletion."""
from types import SimpleNamespace
import pytest
from fastapi import HTTPException
from sqlalchemy import create_engine, text, inspect
from sqlalchemy.orm import sessionmaker
import core.database as cdb
import core.session_manager as sm
import routes.chat_helpers as helpers
import routes.chat_routes as chat
import routes.session_routes as sessions
from src import endpoint_resolver
BASE = "https://chatgpt.com/backend-api/codex"
@pytest.fixture
def state(monkeypatch):
engine = create_engine("sqlite:///:memory:")
cdb.Base.metadata.create_all(engine)
factory = sessionmaker(bind=engine)
for module in (sm, helpers, chat, sessions, endpoint_resolver):
monkeypatch.setattr(module, "SessionLocal", factory)
with factory() as db:
for key, owner in (("a", "alice"), ("b", "alice"), ("c", "bob")):
db.add(cdb.ModelEndpoint(id=key, owner=owner, name=key, base_url=BASE,
provider_auth_id="auth-" + key, is_enabled=True,
cached_models='["gpt-5.5"]', supports_tools=False))
db.commit()
manager = sm.SessionManager.__new__(sm.SessionManager)
manager.sessions = {}
manager.upload_handler = None
monkeypatch.setattr(endpoint_resolver, "resolve_endpoint_runtime",
lambda ep, owner=None: (BASE, "secret-" + ep.id))
yield factory, manager
engine.dispose()
def _session(manager, binding="b"):
return manager.create_session("s", "Chat", BASE + "/responses", "gpt-5.5",
owner="alice", endpoint_id=binding)
def test_binding_survives_metadata_full_reload_and_sync(state):
factory, manager = state
_session(manager)
with factory() as db:
db.add(cdb.ChatMessage(id="message", session_id="s", role="user", content="Hello"))
db.commit()
row = db.get(cdb.Session, "s")
assert manager._db_to_session_meta(row).endpoint_id == "b"
assert manager._db_to_session(row, db).endpoint_id == "b"
row.endpoint_id = "a"
db.commit()
manager.sync_session_metadata("s")
assert manager.sessions["s"].endpoint_id == "a"
@pytest.mark.parametrize("disabled", [True, False])
def test_missing_or_disabled_bound_account_cannot_borrow_only_remaining_sibling(state, disabled):
factory, manager = state
sess = _session(manager)
with factory() as db:
endpoint = db.get(cdb.ModelEndpoint, "b")
if disabled:
endpoint.is_enabled = False
else:
db.delete(endpoint)
db.get(cdb.Session, "s").headers = {"Authorization": "Bearer stale-secret"}
db.commit()
sess.headers = {"Authorization": "Bearer stale-secret"}
helpers.resolve_session_auth(sess, "s", "alice")
assert sess.headers == {}
assert sess.endpoint_id == "b"
assert chat._clear_orphaned_session_endpoint(sess, "alice") is True
assert chat._recover_empty_session_model(sess, "s", "alice") is False
with factory() as db:
assert db.get(cdb.Session, "s").headers == {}
assert db.get(cdb.Session, "s").endpoint_id == "b"
def test_legacy_binding_is_persisted_even_when_authentication_fails(state, monkeypatch):
factory, manager = state
sess = _session(manager, None)
def unavailable(*args, **kwargs):
raise RuntimeError("credentials unavailable")
monkeypatch.setattr(endpoint_resolver, "resolve_endpoint_runtime", unavailable)
helpers.resolve_session_auth(sess, "s", "alice")
assert sess.endpoint_id == "a"
with factory() as db:
assert db.get(cdb.Session, "s").endpoint_id == "a"
def test_explicit_selection_switches_same_url_and_model_binding(state):
factory, manager = state
sess = _session(manager)
assert chat._reconcile_selected_route_from_request(None, sess, "s", {
"selected_model": "gpt-5.5", "selected_endpoint_id": "a",
}, "alice")
assert sess.endpoint_id == "a"
with factory() as db:
assert db.get(cdb.Session, "s").endpoint_id == "a"
assert not chat._reconcile_selected_route_from_request(None, sess, "s", {
"selected_model": "gpt-5.5", "selected_endpoint_id": "c",
}, "alice")
assert sess.endpoint_id == "a"
def test_url_only_model_change_retains_exact_binding(state):
factory, manager = state
sess = _session(manager)
assert chat._reconcile_selected_route_from_request(None, sess, "s", {
"selected_model": "another-model", "selected_endpoint_url": BASE,
}, "alice")
assert sess.endpoint_id == "b"
def test_session_patch_and_unrelated_rename_preserve_binding(state, monkeypatch):
factory, manager = state
sess = _session(manager)
monkeypatch.setattr(sessions, "_verify_session_owner", lambda *args: None)
router = sessions.setup_session_routes(manager, {})
patch = [r.endpoint for r in router.routes if r.path == "/api/session/{sid}" and "PATCH" in r.methods][-1]
request = SimpleNamespace(state=SimpleNamespace(current_user="alice"))
kwargs = dict(request=request, sid="s", name=None, folder=None, cwd=None)
patch(**kwargs, model="gpt-5.5", endpoint_url=BASE, endpoint_id="a")
assert sess.endpoint_id == "a"
with factory() as db:
assert db.get(cdb.Session, "s").endpoint_id == "a"
with pytest.raises(HTTPException):
patch(**kwargs, model="gpt-5.5", endpoint_url=BASE, endpoint_id="c")
patch(**kwargs, model=None, endpoint_url=None, endpoint_id=None)
assert sess.endpoint_id == "a"
def test_session_list_reports_bound_account_for_duplicate_model_and_url(state, monkeypatch):
factory, manager = state
sess = _session(manager)
manager.get_sessions_for_user = lambda *args, **kwargs: {"s": sess}
router = sessions.setup_session_routes(manager, {})
listing = [r.endpoint for r in router.routes if r.path == "/api/sessions" and "GET" in r.methods][-1]
rows = listing(SimpleNamespace(state=SimpleNamespace(current_user="alice"), query_params={}))
assert rows[0]["endpoint_id"] == "b"
assert rows[0]["endpoint_name"] == "b"
def test_migration_is_additive_idempotent_and_repairs_missing_index(tmp_path, monkeypatch):
engine = create_engine(f"sqlite:///{tmp_path / 'legacy.db'}")
with engine.begin() as db:
db.execute(text("CREATE TABLE sessions (id VARCHAR PRIMARY KEY, model VARCHAR)"))
db.execute(text("INSERT INTO sessions VALUES ('legacy', 'gpt-5.5')"))
monkeypatch.setattr(cdb, "engine", engine)
cdb._migrate_add_session_endpoint_id_column()
cdb._migrate_add_session_endpoint_id_column()
with engine.begin() as db:
assert db.execute(text("SELECT model, endpoint_id FROM sessions")).one() == ("gpt-5.5", None)
db.execute(text("DROP INDEX ix_sessions_endpoint_id"))
cdb._migrate_add_session_endpoint_id_column()
assert [idx["name"] for idx in inspect(engine).get_indexes("sessions")] == ["ix_sessions_endpoint_id"]
engine.dispose()