diff --git a/app.py b/app.py index bb4f51ffb..82c92d5d7 100644 --- a/app.py +++ b/app.py @@ -516,23 +516,45 @@ async def serve_generated_image(filename: str, request: Request): # SECURITY: filename is the only key, so anyone who knows / guesses a # 12-hex content hash could pull another user's image bytes. Require # auth and verify ownership via the gallery row (when one exists). + _is_bearer = False try: - from src.auth_helpers import get_current_user + from src.auth_helpers import ( + effective_user, + get_current_user, + is_bearer_principal, + require_chat_scope, + ) from core.database import SessionLocal as _SL, GalleryImage as _GI - _user = get_current_user(request) + _is_bearer = is_bearer_principal(request) + if _is_bearer: + # Gallery JSON attributes rows to the token owner. Reuse the same + # owner/scope gate for the binary follow-up so the returned URL is + # actually readable by that bearer principal. + require_chat_scope(request) + _user = effective_user(request) + else: + _user = get_current_user(request) if _user: _db = _SL() try: _row = _db.query(_GI).filter(_GI.filename == filename).first() # Generated-but-not-yet-imported images have no row → allow. - # Row exists with a different owner → 404 (don't confirm existence). - if _row is not None and _row.owner and _row.owner != _user: + # A bearer gallery row must have the exact token owner; cookie + # callers retain the legacy null-owner compatibility below. + if _row is not None and ( + (_is_bearer and _row.owner != _user) + or (not _is_bearer and _row.owner and _row.owner != _user) + ): raise HTTPException(status_code=404, detail="Image not found") finally: _db.close() except HTTPException: raise except Exception as _e: + if _is_bearer: + # An authenticated bearer request must not become a public file + # read because ownership lookup degraded or the DB was unavailable. + raise HTTPException(status_code=404, detail="Image not found") from _e logger.warning("Image ownership verification failed for %r", filename, exc_info=_e) ext = filename.rsplit('.', 1)[-1].lower() mime = { diff --git a/core/database.py b/core/database.py index dc2298e35..f085120bc 100644 --- a/core/database.py +++ b/core/database.py @@ -187,6 +187,13 @@ class Session(TimestampMixin, Base): endpoint_url = Column(String, nullable=False) model = Column(String, nullable=False) owner = Column(String, nullable=True, index=True) # username; null = legacy/shared + + # Bearer-chat sessions must retain the exact server-owned endpoint they + # were created from. Keep this reference non-cascading so endpoint + # disable/delete/owner changes remain observable as an orphan and fail + # closed at the next bearer LLM boundary. + model_endpoint_id = Column(String, nullable=True, index=True) + endpoint_provenance = Column(String, nullable=True) # Configuration flags rag = Column(Boolean, default=False) @@ -999,6 +1006,40 @@ def _migrate_add_owner_column(): except Exception: pass + +def _migrate_add_session_endpoint_provenance_columns(): + """Add the durable endpoint identity used by bearer session validation.""" + import sqlite3 + + db_path = DATABASE_URL.replace("sqlite:///", "") + if not os.path.exists(db_path): + return + conn = None + try: + conn = sqlite3.connect(db_path) + columns = {row[1] for row in conn.execute("PRAGMA table_info(sessions)")} + if "model_endpoint_id" not in columns: + conn.execute("ALTER TABLE sessions ADD COLUMN model_endpoint_id TEXT") + if "endpoint_provenance" not in columns: + conn.execute("ALTER TABLE sessions ADD COLUMN endpoint_provenance TEXT") + conn.execute( + "CREATE INDEX IF NOT EXISTS ix_sessions_model_endpoint_id " + "ON sessions(model_endpoint_id)" + ) + conn.commit() + logging.getLogger(__name__).info( + "Migrated: added session endpoint identity/provenance columns" + ) + except Exception as e: + logging.getLogger(__name__).warning( + "Session endpoint provenance migration failed: %s", e + ) + finally: + try: + conn.close() + except Exception: + pass + def _migrate_model_endpoints(): """Recreate model_endpoints table if schema changed (url->base_url).""" import sqlite3 @@ -2152,6 +2193,7 @@ def init_db(): _migrate_add_supports_tools_column() _migrate_add_task_run_model_column() _migrate_add_owner_column() + _migrate_add_session_endpoint_provenance_columns() _migrate_add_document_archived_column() _migrate_add_last_message_at_column() _migrate_add_folder_column() diff --git a/core/models.py b/core/models.py index b7a6f8c60..cfc8b53da 100644 --- a/core/models.py +++ b/core/models.py @@ -90,6 +90,8 @@ class Session: headers: Optional[Dict[str, str]] = None history: List[ChatMessage] = None owner: Optional[str] = None + model_endpoint_id: Optional[str] = None + endpoint_provenance: Optional[str] = None is_important: bool = False message_count: int = 0 diff --git a/core/session_manager.py b/core/session_manager.py index de8f674a8..7f0937312 100644 --- a/core/session_manager.py +++ b/core/session_manager.py @@ -165,6 +165,8 @@ class SessionManager: headers=headers, history=[], owner=getattr(db_session, "owner", None), + model_endpoint_id=getattr(db_session, "model_endpoint_id", None), + endpoint_provenance=getattr(db_session, "endpoint_provenance", None), is_important=getattr(db_session, "is_important", False) or False, ) session.message_count = getattr(db_session, "message_count", 0) or 0 @@ -221,6 +223,8 @@ class SessionManager: headers=headers, history=history, owner=getattr(db_session, 'owner', None), + model_endpoint_id=getattr(db_session, 'model_endpoint_id', None), + endpoint_provenance=getattr(db_session, 'endpoint_provenance', None), is_important=getattr(db_session, 'is_important', False) or False, ) @@ -502,6 +506,8 @@ class SessionManager: session.rag = db_session.rag session.archived = db_session.archived session.owner = getattr(db_session, "owner", None) + session.model_endpoint_id = getattr(db_session, "model_endpoint_id", None) + session.endpoint_provenance = getattr(db_session, "endpoint_provenance", None) session.is_important = getattr(db_session, "is_important", False) or False session.message_count = ( db.query(DbChatMessage) @@ -602,6 +608,50 @@ class SessionManager: finally: db.close() + def set_session_endpoint_provenance( + self, + session_id: str, + *, + model_endpoint_id: Optional[str], + endpoint_provenance: str, + ) -> bool: + """Persist the server-owned endpoint provenance for a session. + + ``registered`` rows carry an exact ModelEndpoint id. ``direct`` rows + deliberately carry no endpoint id and retain direct API-key + compatibility. The values are assigned only after the durable write + succeeds so an in-memory session cannot claim provenance the database + did not accept. + """ + provenance = str(endpoint_provenance or "").strip().lower() + endpoint_id = str(model_endpoint_id or "").strip() or None + if provenance == "registered" and not endpoint_id: + raise ValueError("registered session provenance requires an endpoint id") + if provenance == "direct": + endpoint_id = None + if provenance not in {"registered", "direct"}: + raise ValueError("unsupported session endpoint provenance") + + db = SessionLocal() + try: + db_session = db.query(DbSession).filter(DbSession.id == session_id).first() + if db_session is None: + raise KeyError(f"Session {session_id} not found") + db_session.model_endpoint_id = endpoint_id + db_session.endpoint_provenance = provenance + db.commit() + except Exception: + db.rollback() + raise + finally: + db.close() + + session = self.sessions.get(session_id) + if session is not None: + session.model_endpoint_id = endpoint_id + session.endpoint_provenance = provenance + return True + def delete_session(self, session_id: str) -> bool: """Permanently delete a session and all its messages.""" db = SessionLocal() diff --git a/routes/chat_helpers.py b/routes/chat_helpers.py index 256da4fb7..ffc5f582b 100644 --- a/routes/chat_helpers.py +++ b/routes/chat_helpers.py @@ -531,9 +531,20 @@ def resolve_session_auth( is_chatgpt_subscription = is_chatgpt_subscription_base(getattr(sess, "endpoint_url", "") or "") except Exception: is_chatgpt_subscription = False + provenance = (getattr(sess, "endpoint_provenance", None) or "").strip().lower() + endpoint_id = (getattr(sess, "model_endpoint_id", None) or "").strip() has_auth = _has_auth_keys(sess.headers) - if has_auth and not is_chatgpt_subscription: + if has_auth and not is_chatgpt_subscription and provenance != "registered": return + if provenance == "direct": + # A direct API-key session owns its request headers; a same-URL + # registered endpoint must never supply another user's credentials by + # coincidence. + return + if provenance == "registered": + # Do not carry a previously persisted key through endpoint rotation or + # an unavailable endpoint while attempting exact re-resolution below. + sess.headers = {} try: from src.endpoint_resolver import build_headers, resolve_endpoint_runtime @@ -549,6 +560,10 @@ def resolve_session_auth( # with similar endpoint URLs can borrow each other's API key. from src.auth_helpers import owner_filter q = owner_filter(q, ModelEndpoint, owner) + if provenance == "registered": + if not endpoint_id: + return + q = q.filter(ModelEndpoint.id == endpoint_id) for ep in q.all(): if not _session_url_matches_endpoint(target_url, ep.base_url or ""): continue @@ -617,6 +632,12 @@ def _normalize_model_id_from_cache(sess) -> Optional[str]: if not session_base: return None + provenance = getattr(sess, "endpoint_provenance", None) + endpoint_id = (getattr(sess, "model_endpoint_id", None) or "").strip() + if provenance == "direct": + # Direct API-key sessions are intentionally outside the registered + # endpoint inventory. Never borrow a same-URL endpoint's model list. + return None db = SessionLocal() try: q = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) @@ -624,6 +645,10 @@ def _normalize_model_id_from_cache(sess) -> Optional[str]: if owner: from src.auth_helpers import owner_filter q = owner_filter(q, ModelEndpoint, owner) + if provenance == "registered": + if not endpoint_id: + return None + q = q.filter(ModelEndpoint.id == endpoint_id) endpoints = q.all() for ep in endpoints: try: @@ -660,41 +685,79 @@ def _validate_bearer_session_model(sess, owner: str | None = None) -> Optional[s sessions, including provider-auth-backed rows, must use the visible server-owned inventory and never trigger a provider lookup here. """ + # Lightweight in-memory test doubles from older route tests do not carry + # durable provenance fields. They cannot represent a persisted bearer + # session; retain their historical seam while every SessionManager-loaded + # object (which always has both fields) takes the fail-closed path below. + if not hasattr(sess, "endpoint_provenance") and not hasattr(sess, "model_endpoint_id"): + return None + + provenance = (getattr(sess, "endpoint_provenance", None) or "").strip().lower() + endpoint_id = (getattr(sess, "model_endpoint_id", None) or "").strip() + if provenance == "direct": + if endpoint_id: + raise HTTPException(400, "Direct API-key sessions cannot carry a registered endpoint") + # No registered ModelEndpoint row is consulted for this documented + # compatibility path. + return None + if provenance != "registered": + raise HTTPException(400, "Session endpoint provenance is unavailable") + if not owner: + raise HTTPException(403, "A bearer session owner is required") + if not endpoint_id: + raise HTTPException(400, "Registered session endpoint identity is unavailable") + endpoint_url = (getattr(sess, "endpoint_url", "") or "").strip() requested = (getattr(sess, "model", "") or "").strip() - if not endpoint_url or not requested: - return None - try: - session_base = normalize_base(endpoint_url) - except Exception: - session_base = endpoint_url.rstrip("/") - if not session_base: - return None + if not endpoint_url: + raise HTTPException(400, "Registered session endpoint is not configured") + if not requested: + raise HTTPException(400, "Registered session model is not configured") db = SessionLocal() try: - q = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) - if owner: - from src.auth_helpers import owner_filter + from src.auth_helpers import owner_filter - q = owner_filter(q, ModelEndpoint, owner) - for ep in q.all(): - try: - if normalize_base(getattr(ep, "base_url", "") or "") != session_base: - continue - except Exception: - continue - from routes.model_routes import _validate_bearer_model_selection + q = db.query(ModelEndpoint).filter( + ModelEndpoint.id == endpoint_id, + ModelEndpoint.is_enabled == True, + ) + q = owner_filter(q, ModelEndpoint, owner) + endpoints = q.all() + if len(endpoints) != 1: + # This covers disabled/deleted/owner-mismatched rows as well as + # malformed duplicate results. Do not fall back to URL matching. + raise HTTPException(400, "Registered model endpoint is no longer available") + ep = endpoints[0] + if not _session_url_matches_endpoint(endpoint_url, getattr(ep, "base_url", "") or ""): + raise HTTPException(400, "Session endpoint provenance is stale") - sess.model = _validate_bearer_model_selection(ep, requested) - return sess.model + from routes.model_routes import _validate_bearer_model_selection + + validated = _validate_bearer_model_selection(ep, requested) + + # A session may outlive an endpoint-key rotation. For bearer calls, + # use the current static key for this exact endpoint and never trust a + # stale persisted Authorization header. Provider-auth rows remain + # request-local and are intentionally empty in cache-only mode. + try: + from src.endpoint_resolver import build_headers, resolve_endpoint_runtime + + base, api_key = resolve_endpoint_runtime( + ep, + owner=owner, + allow_live_probes=False, + ) + sess.headers = build_headers(api_key, base) + except Exception as exc: + logger.warning("Could not refresh bearer session endpoint auth: %s", exc) + sess.headers = {} + + sess.model = validated + return validated finally: db.close() - # No registered endpoint means this is the documented direct API-key - # compatibility path, not an endpoint-picker selection. - return None - def _session_is_research_spinoff(sess) -> bool: """True if this session was created via research "Discuss" spin-off. diff --git a/routes/chat_routes.py b/routes/chat_routes.py index 8019be436..3f3386996 100644 --- a/routes/chat_routes.py +++ b/routes/chat_routes.py @@ -434,11 +434,21 @@ def _clear_orphaned_session_endpoint( return False db = SessionLocal() try: + provenance = (getattr(sess, "endpoint_provenance", None) or "").strip().lower() + endpoint_id = (getattr(sess, "model_endpoint_id", None) or "").strip() + if provenance == "direct": + return False q = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) if owner: from src.auth_helpers import owner_filter q = owner_filter(q, ModelEndpoint, owner) - endpoints = q.all() + if provenance == "registered": + if not endpoint_id: + endpoints = [] + else: + endpoints = q.filter(ModelEndpoint.id == endpoint_id).all() + else: + endpoints = q.all() for ep in endpoints: if _session_url_matches_endpoint(sess.endpoint_url or "", ep.base_url or ""): return False @@ -568,15 +578,24 @@ def _recover_empty_session_model(sess, session_id: str, owner: str | None = None return False db = SessionLocal() try: - # Prefer the endpoint whose base URL matches the session — we know the - # user already pointed this session at that endpoint, so its first - # cached model is the most defensible default. + provenance = (getattr(sess, "endpoint_provenance", None) or "").strip().lower() + endpoint_id = (getattr(sess, "model_endpoint_id", None) or "").strip() + has_provenance_fields = hasattr(sess, "endpoint_provenance") or hasattr(sess, "model_endpoint_id") + if not allow_live_probes and has_provenance_fields and provenance not in {"registered", "direct"}: + return False + # Registered sessions use their immutable endpoint identity. URL-only + # fallback remains for legacy browser sessions, but is never used to + # recover a bearer session with a missing/invalid identity. ep = None if getattr(sess, "endpoint_url", ""): q = db.query(ModelEndpoint).filter(ModelEndpoint.is_enabled == True) if owner: from src.auth_helpers import owner_filter q = owner_filter(q, ModelEndpoint, owner) + if provenance == "registered": + if not endpoint_id: + return False + q = q.filter(ModelEndpoint.id == endpoint_id) endpoints = q.all() for cand in endpoints: if _session_url_matches_endpoint(sess.endpoint_url or "", cand.base_url or ""): @@ -707,6 +726,7 @@ def _reconcile_selected_route_from_request( endpoint_url = "" headers = None + resolved_endpoint_id = None if selected_endpoint_id or selected_endpoint_url: try: from src.auth_helpers import owner_filter @@ -718,16 +738,27 @@ def _reconcile_selected_route_from_request( q = q.filter(ModelEndpoint.id == selected_endpoint_id) if owner: q = owner_filter(q, ModelEndpoint, owner) - candidates = q.all() if selected_endpoint_url and not selected_endpoint_id else [q.first()] - ep = None - for cand in candidates: - if not cand: - continue - if selected_endpoint_id or _session_url_matches_endpoint(selected_endpoint_url, cand.base_url or ""): - ep = cand - break + if selected_endpoint_id: + candidates = [q.first()] + else: + candidates = [ + cand for cand in q.all() + if cand and _session_url_matches_endpoint( + selected_endpoint_url, + cand.base_url or "", + ) + ] + # A URL is not stable identity. Refuse to choose between + # duplicate visible endpoints instead of binding the + # session to whichever row happens to come first. + if len(candidates) != 1: + return False + ep = candidates[0] if candidates else None if not ep: return False + resolved_endpoint_id = str(getattr(ep, "id", "") or "").strip() or None + if not resolved_endpoint_id: + return False endpoint_url = build_chat_url(normalize_base(ep.base_url or "")) headers = build_headers(ep.api_key or "", ep.base_url or "") if ep.api_key else {} finally: @@ -742,12 +773,16 @@ def _reconcile_selected_route_from_request( if ( selected_model == (getattr(sess, "model", "") or "") and endpoint_url == (getattr(sess, "endpoint_url", "") or "") + and resolved_endpoint_id == (getattr(sess, "model_endpoint_id", "") or "") + and getattr(sess, "endpoint_provenance", None) == "registered" ): return False sess.model = selected_model sess.endpoint_url = endpoint_url sess.headers = headers or {} + sess.model_endpoint_id = resolved_endpoint_id + sess.endpoint_provenance = "registered" db = SessionLocal() try: db_session = db.query(DBSession).filter(DBSession.id == session_id).first() @@ -755,6 +790,8 @@ def _reconcile_selected_route_from_request( db_session.model = selected_model db_session.endpoint_url = endpoint_url db_session.headers = sess.headers or {} + db_session.model_endpoint_id = resolved_endpoint_id + db_session.endpoint_provenance = "registered" db_session.updated_at = datetime.utcnow() db.commit() finally: @@ -2930,6 +2967,11 @@ def setup_chat_routes( except (KeyError, SessionNotFoundError): raise HTTPException(404, "Session not found") + if capability.is_bearer: + # Rewrite is a direct streaming LLM consumer, so it must enforce + # the same server-owned session model/endpoint invariant as chat. + _validate_bearer_session_model(sess, owner=effective_user(request)) + messages = [ {"role": "system", "content": ( "You are rewriting a previous response. Follow the instruction exactly. " diff --git a/routes/compare/compare_routes.py b/routes/compare/compare_routes.py index 961e6dd29..0fedbd9af 100644 --- a/routes/compare/compare_routes.py +++ b/routes/compare/compare_routes.py @@ -12,6 +12,7 @@ import logging from core.database import Comparison, SessionLocal from core.session_manager import SessionManager from src.auth_helpers import effective_user, is_bearer_principal, require_chat_scope +from src.session_provenance import persist_session_endpoint_provenance from routes.session_routes import _reject_raw_endpoint_url_for_non_admin logger = logging.getLogger(__name__) @@ -220,15 +221,24 @@ def setup_compare_routes(session_manager: SessionManager): # `ep` is None (raw admin URL or no match), so a comparison can # never inherit another user's key/headers. headers = build_headers(ep.api_key, ep.base_url) if (ep and ep.api_key) else None - resolved.append((sid, selected_model, session_endpoint_url, headers)) + resolved.append( + ( + sid, + selected_model, + session_endpoint_url, + headers, + str(ep.id) if ep is not None else None, + "registered" if ep is not None else None, + ) + ) finally: db.close() # Both endpoints validated — only now create the ephemeral [CMP] # sessions and copy any resolved headers. - for sid, model, session_endpoint_url, headers in resolved: + for sid, model, session_endpoint_url, headers, endpoint_id, provenance in resolved: name = f"[CMP] {slot_name[sid]}" if blind else f"[CMP] {model.split('/')[-1]}" - session_manager.create_session( + comparison_session = session_manager.create_session( session_id=sid, name=name, endpoint_url=session_endpoint_url, @@ -236,6 +246,14 @@ def setup_compare_routes(session_manager: SessionManager): rag=False, owner=user, ) + if provenance in {"registered", "direct"}: + persist_session_endpoint_provenance( + session_manager, + sid, + comparison_session, + model_endpoint_id=endpoint_id, + endpoint_provenance=provenance, + ) if headers: s = session_manager.sessions.get(sid) if s: diff --git a/routes/history/history_routes.py b/routes/history/history_routes.py index 88dc53010..903fc1566 100644 --- a/routes/history/history_routes.py +++ b/routes/history/history_routes.py @@ -23,12 +23,14 @@ from src.message_metadata import ( ) from src.topic_analyzer import analyze_topics from src.upload_handler import reserve_message_upload_references +from src.session_provenance import persist_session_endpoint_provenance from routes.session_routes import ( _message_role, _message_text, _reject_compact_during_active_run, _verify_session_owner, ) +from routes.chat_helpers import _validate_bearer_session_model logger = logging.getLogger(__name__) @@ -664,6 +666,15 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter: if not source: raise HTTPException(404, "Session not found") + source_provenance = getattr(source, "endpoint_provenance", None) + source_endpoint_id = getattr(source, "model_endpoint_id", None) + if ( + is_bearer_principal(request) + and (hasattr(source, "endpoint_provenance") or hasattr(source, "model_endpoint_id")) + and source_provenance not in {"registered", "direct"} + ): + raise HTTPException(400, "Session endpoint provenance is unavailable") + # Create new session new_id = str(uuid.uuid4()) fork_name = f"\u2ADD {source.name}" @@ -675,6 +686,14 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter: rag=False, owner=getattr(source, 'owner', None), ) + if source_provenance in {"registered", "direct"}: + persist_session_endpoint_provenance( + session_manager, + new_id, + new_session, + model_endpoint_id=source_endpoint_id, + endpoint_provenance=source_provenance, + ) # Copy messages up to keep_count msgs_to_copy = source.history[:keep_count] @@ -796,6 +815,9 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter: if len(session.history) < 6: return {"status": "ok", "message": "Not enough messages to compact"} + if capability.is_bearer: + _validate_bearer_session_model(session, owner=owner) + context_kwargs = {} if not capability.allow_live_probes: context_kwargs["allow_live_probes"] = False @@ -927,6 +949,8 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter: "after": pct_after, } + except HTTPException: + raise except Exception as e: logger.error(f"Manual compact error {session_id}: {e}") raise HTTPException(500, str(e)) diff --git a/routes/session_routes.py b/routes/session_routes.py index bd84c182c..a31747f3f 100644 --- a/routes/session_routes.py +++ b/routes/session_routes.py @@ -27,6 +27,7 @@ from src.message_metadata import ( from src.session_image_cleanup import _generated_image_path_for_cleanup, session_image_refs from src.session_actions import is_session_recently_active from src.upload_handler import reserve_message_upload_references +from src.session_provenance import persist_session_endpoint_provenance def _sanitize_export_filename(name: str) -> str: @@ -477,6 +478,20 @@ def setup_session_routes( rag=str(rag).lower() == "true" if rag else False, owner=user, ) + if endpoint_row is not None or request_api_key: + try: + persist_session_endpoint_provenance( + session_manager, + sid, + session, + model_endpoint_id=getattr(endpoint_row, "id", None), + endpoint_provenance=( + "registered" if endpoint_row is not None else "direct" + ), + ) + except Exception as exc: + logger.error("Failed to persist session endpoint provenance for %s: %s", sid, exc) + raise HTTPException(500, "Failed to persist session endpoint provenance") from exc # Set auth headers for custom API-key endpoints resolved_key = request_api_key resolved_base = endpoint_url @@ -540,6 +555,7 @@ def setup_session_routes( _reject_raw_endpoint_url_for_non_admin(request, user, endpoint_id, endpoint_url) endpoint_api_key = "" endpoint_base_url = "" + endpoint_row = None if endpoint_id: from core.database import ModelEndpoint from src.auth_helpers import owner_filter @@ -555,13 +571,23 @@ def setup_session_routes( ep = q.first() if not ep: raise HTTPException(400, "Model endpoint no longer exists") + endpoint_row = ep endpoint_base_url = ep.base_url or "" endpoint_api_key = ep.api_key or "" endpoint_url = build_chat_url(normalize_base(endpoint_base_url)) finally: _db.close() + if is_bearer_principal(request) and endpoint_row is not None: + from routes.model_routes import _validate_bearer_model_selection + + # Validate before mutating either the in-memory or durable + # session. The same server-owned inventory is enforced again + # immediately before each bearer LLM consumer. + model = _validate_bearer_model_selection(endpoint_row, model) session.model = model session.endpoint_url = endpoint_url + session.model_endpoint_id = getattr(endpoint_row, "id", None) + session.endpoint_provenance = "registered" if endpoint_row is not None else None # Update auth headers from the endpoint's stored API key if endpoint_api_key: from src.endpoint_resolver import build_headers @@ -576,6 +602,8 @@ def setup_session_routes( db_session.model = model db_session.endpoint_url = endpoint_url db_session.headers = session.headers or {} + db_session.model_endpoint_id = getattr(endpoint_row, "id", None) + db_session.endpoint_provenance = "registered" if endpoint_row is not None else None db_session.updated_at = utcnow_naive() db.commit() finally: @@ -1050,6 +1078,11 @@ def setup_session_routes( if not older: raise HTTPException(400, "Nothing old enough to compact") + if capability.is_bearer: + from routes.chat_helpers import _validate_bearer_session_model + + _validate_bearer_session_model(session, owner=effective_user(request)) + from src.context_compactor import SELF_SUMMARY_SYSTEM_PROMPT from src.llm_core import llm_call_async diff --git a/routes/webhook/webhook_routes.py b/routes/webhook/webhook_routes.py index 8a6876435..2007c884e 100644 --- a/routes/webhook/webhook_routes.py +++ b/routes/webhook/webhook_routes.py @@ -18,6 +18,7 @@ from src.auth_helpers import ( ) from src.url_security import validate_public_http_url from src.webhook_manager import WebhookManager, validate_webhook_url, validate_events +from src.session_provenance import persist_session_endpoint_provenance logger = logging.getLogger(__name__) @@ -365,6 +366,12 @@ def setup_webhook_routes( _sess_owner = getattr(sess, "owner", None) if not _caller_owns_session(_sess_owner, _tok_user): raise HTTPException(404, "Session not found") + if is_bearer_principal(request): + from routes.chat_helpers import _validate_bearer_session_model + + # Existing-session resume is an LLM boundary too; ownership + # alone must not authorize the persisted endpoint/model. + _validate_bearer_session_model(sess, owner=token_owner) # --- Case 2: Direct API key + model (no pre-configured endpoint needed) --- if not sess and body.api_key: @@ -397,6 +404,12 @@ def setup_webhook_routes( session_id=sid, name="API Chat", endpoint_url=endpoint_url, model=model, owner=token_owner, ) + persist_session_endpoint_provenance( + session_manager, + sid, + sess, + endpoint_provenance="direct", + ) sess.headers = build_headers(api_key, base_url) session_manager.save_sessions() session_id = sid @@ -446,12 +459,29 @@ def setup_webhook_routes( session_id=sid, name="API Chat", endpoint_url=endpoint_url, model=model, owner=token_owner, ) + endpoint_id = getattr(ep, "id", None) + if endpoint_id: + persist_session_endpoint_provenance( + session_manager, + sid, + sess, + model_endpoint_id=endpoint_id, + endpoint_provenance="registered", + ) if api_key: sess.headers = build_headers(api_key, base_url) session_manager.save_sessions() session_id = sid # --- Send message and get response --- + if is_bearer_principal(request): + from routes.chat_helpers import _validate_bearer_session_model + + # The fallback branch has just created the session, so it did not + # pass through the existing-session gate above. Recheck the + # durable endpoint identity immediately before the LLM boundary + # for every bearer path, including malformed endpoint rows. + _validate_bearer_session_model(sess, owner=token_owner) sess.add_message(ChatMessage("user", message)) messages = [{"role": m.role, "content": m.content} for m in sess.history] diff --git a/src/session_provenance.py b/src/session_provenance.py new file mode 100644 index 000000000..5c775a002 --- /dev/null +++ b/src/session_provenance.py @@ -0,0 +1,38 @@ +"""Small shared seam for recording session endpoint provenance.""" + + +def persist_session_endpoint_provenance( + session_manager, + session_id: str, + session, + *, + model_endpoint_id=None, + endpoint_provenance: str, +) -> None: + """Record trusted endpoint provenance on a session and its durable row. + + The production ``SessionManager`` owns the database write. Lightweight + route test doubles may not implement that method, so they still receive + the same in-memory fields without changing the production contract. + """ + provenance = str(endpoint_provenance or "").strip().lower() + endpoint_id = str(model_endpoint_id or "").strip() or None + if provenance == "registered" and not endpoint_id: + raise ValueError("registered session provenance requires an endpoint id") + if provenance == "direct": + endpoint_id = None + if provenance not in {"registered", "direct"}: + raise ValueError("unsupported session endpoint provenance") + + setter = getattr(session_manager, "set_session_endpoint_provenance", None) + if callable(setter): + setter( + session_id, + model_endpoint_id=endpoint_id, + endpoint_provenance=provenance, + ) + if session is None: + session = getattr(session_manager, "sessions", {}).get(session_id) + if session is not None: + setattr(session, "model_endpoint_id", endpoint_id) + setattr(session, "endpoint_provenance", provenance) diff --git a/tests/test_api_token_repair_cycle7.py b/tests/test_api_token_repair_cycle7.py new file mode 100644 index 000000000..af1c31363 --- /dev/null +++ b/tests/test_api_token_repair_cycle7.py @@ -0,0 +1,634 @@ +"""Forward probes and regressions for the cycle-7 API-token repair. + +The first run of this file is intentionally against the vulnerable candidate: +the security assertions below should fail before the repair is applied. The +same tests remain as focused regressions after the fix. +""" + +import sys +from types import SimpleNamespace + +import pytest +from fastapi import HTTPException +from sqlalchemy import create_engine +from sqlalchemy.orm import sessionmaker +from sqlalchemy.pool import NullPool + +import core.database as cdb +from core.models import ChatMessage + + +class _Request: + def __init__(self, *, owner="alice", body=None, bearer=True): + self.state = SimpleNamespace( + api_token=bearer, + api_token_owner=owner if bearer else None, + api_token_scopes=["chat"] if bearer else [], + current_user="api" if bearer else owner, + ) + self.app = SimpleNamespace(state=SimpleNamespace(auth_manager=None)) + self.headers = {} + self.query_params = {} + self.client = SimpleNamespace(host="127.0.0.1") + self._body = body + + async def json(self): + return self._body + + +def _endpoint(router, path, method): + for route in reversed(router.routes): + if route.path == path and method in route.methods: + return route.endpoint + raise AssertionError(f"route not found: {method} {path}") + + +class _Query: + def __init__(self, rows): + self.rows = list(rows) + + def filter(self, *args, **kwargs): + return self + + def order_by(self, *args, **kwargs): + return self + + def all(self): + return list(self.rows) + + def first(self): + return self.rows[0] if self.rows else None + + +class _Db: + def __init__(self, rows_by_model): + self.rows_by_model = rows_by_model + + def query(self, model): + return _Query(self.rows_by_model.get(model, self.rows_by_model.get(None, []))) + + def close(self): + return None + + def commit(self): + return None + + def rollback(self): + return None + + def add(self, value): + return None + + def delete(self, value): + return None + + +def _registered_endpoint(*, endpoint_id="ep-1", models='["safe-model"]', owner="alice"): + return SimpleNamespace( + id=endpoint_id, + owner=owner, + base_url="https://api.example.test/v1", + is_enabled=True, + endpoint_kind="api", + cached_models=models, + pinned_models=None, + hidden_models=None, + api_key="", + provider_auth_id=None, + ) + + +def _registered_session(model="unsafe-model", endpoint_id="ep-1"): + return SimpleNamespace( + id="sid", + name="chat", + owner="alice", + endpoint_url="https://api.example.test/v1/chat/completions", + model=model, + headers={}, + history=[], + model_endpoint_id=endpoint_id, + endpoint_provenance="registered", + ) + + +def _patch_validator_db(monkeypatch, endpoint_rows): + from routes import chat_helpers + + monkeypatch.setattr( + chat_helpers, + "SessionLocal", + lambda: _Db({cdb.ModelEndpoint: endpoint_rows}), + ) + + +def _isolated_db(tmp_path): + engine = create_engine( + f"sqlite:///{tmp_path / 'repair-cycle7.db'}", + connect_args={"check_same_thread": False}, + poolclass=NullPool, + ) + cdb.Base.metadata.create_all(engine) + return sessionmaker(bind=engine, autoflush=False, autocommit=False) + + +def test_probe_registered_bearer_session_rejects_ambiguous_or_missing_provenance(monkeypatch): + from routes.chat_helpers import _validate_bearer_session_model + + endpoint_rows = [ + _registered_endpoint(endpoint_id="ep-a"), + _registered_endpoint(endpoint_id="ep-b"), + ] + _patch_validator_db(monkeypatch, endpoint_rows) + session = SimpleNamespace( + endpoint_url="https://api.example.test/v1/chat/completions", + model="safe-model", + model_endpoint_id=None, + endpoint_provenance="registered", + ) + + with pytest.raises(HTTPException): + _validate_bearer_session_model(session, owner="alice") + + +@pytest.mark.parametrize( + ("label", "endpoint_rows"), + [ + ("disabled-or-deleted", []), + ("owner-mismatch", []), + ("url-changed", [_registered_endpoint()]), + ("empty-inventory", [_registered_endpoint(models='[]')]), + ("malformed-inventory", [_registered_endpoint(models="not-json")]), + ("hidden-model", [_registered_endpoint()]), + ], +) +def test_registered_bearer_session_rejects_endpoint_boundary_cases( + monkeypatch, label, endpoint_rows +): + from routes.chat_helpers import _validate_bearer_session_model + + if label == "url-changed": + endpoint_rows[0].base_url = "https://other.example.test/v1" + elif label == "hidden-model": + endpoint_rows[0].hidden_models = '["unsafe-model"]' + endpoint_rows[0].cached_models = '["unsafe-model"]' + elif label == "empty-inventory": + endpoint_rows[0].pinned_models = "[]" + elif label == "owner-mismatch": + endpoint_rows = [] # the owner-scoped query has no visible row + + _patch_validator_db(monkeypatch, endpoint_rows) + with pytest.raises(HTTPException): + _validate_bearer_session_model(_registered_session(), owner="alice") + + +@pytest.mark.parametrize( + "case", + [ + "disabled", + "deleted", + "url-changed", + "empty-inventory", + "malformed-inventory", + "hidden-model", + "owner-mismatch", + ], +) +def test_registered_bearer_session_rejects_durable_endpoint_boundary_cases( + monkeypatch, tmp_path, case +): + from routes.chat_helpers import _validate_bearer_session_model + + session_factory = _isolated_db(tmp_path) + endpoint = cdb.ModelEndpoint( + id="ep-1", + name="Endpoint", + base_url="https://api.example.test/v1", + api_key="", + is_enabled=True, + owner="alice", + endpoint_kind="api", + cached_models='["safe-model"]', + pinned_models=None, + hidden_models=None, + ) + if case == "disabled": + endpoint.is_enabled = False + elif case == "deleted": + endpoint = None + elif case == "url-changed": + endpoint.base_url = "https://other.example.test/v1" + elif case == "empty-inventory": + endpoint.cached_models = "[]" + elif case == "malformed-inventory": + endpoint.cached_models = "not-json" + elif case == "hidden-model": + endpoint.cached_models = '["unsafe-model"]' + endpoint.hidden_models = '["unsafe-model"]' + elif case == "owner-mismatch": + endpoint.owner = "bob" + + if endpoint is not None: + db = session_factory() + try: + db.add(endpoint) + db.commit() + finally: + db.close() + monkeypatch.setattr( + __import__("routes.chat_helpers", fromlist=["SessionLocal"]), + "SessionLocal", + session_factory, + ) + with pytest.raises(HTTPException): + _validate_bearer_session_model(_registered_session(), owner="alice") + + +def test_registered_bearer_session_uses_exact_id_when_base_urls_are_duplicated( + monkeypatch, tmp_path +): + from routes.chat_helpers import _validate_bearer_session_model + + session_factory = _isolated_db(tmp_path) + db = session_factory() + try: + db.add_all( + [ + cdb.ModelEndpoint( + id="ep-wrong", + name="Wrong duplicate", + base_url="https://api.example.test/v1", + api_key="", + is_enabled=True, + owner="alice", + endpoint_kind="api", + cached_models='["wrong-model"]', + ), + cdb.ModelEndpoint( + id="ep-1", + name="Exact duplicate", + base_url="https://api.example.test/v1", + api_key="", + is_enabled=True, + owner="alice", + endpoint_kind="api", + cached_models='["safe-model"]', + ), + ] + ) + db.commit() + finally: + db.close() + monkeypatch.setattr( + __import__("routes.chat_helpers", fromlist=["SessionLocal"]), + "SessionLocal", + session_factory, + ) + session = _registered_session(model="safe-model", endpoint_id="ep-1") + assert _validate_bearer_session_model(session, owner="alice") == "safe-model" + + +def test_registered_bearer_session_refreshes_static_endpoint_headers(monkeypatch): + from routes.chat_helpers import _validate_bearer_session_model + + endpoint = _registered_endpoint() + endpoint.api_key = "current-key" + _patch_validator_db(monkeypatch, [endpoint]) + session = _registered_session(model="safe-model") + session.headers = {"Authorization": "Bearer stale-key"} + + assert _validate_bearer_session_model(session, owner="alice") == "safe-model" + assert session.headers == {"Authorization": "Bearer current-key"} + + +def test_direct_api_key_session_preserves_compatibility_without_inventory_lookup(monkeypatch): + from routes import chat_helpers + + def unexpected_db(): + raise AssertionError("direct API-key sessions must not consult endpoint inventory") + + monkeypatch.setattr(chat_helpers, "SessionLocal", unexpected_db) + session = SimpleNamespace( + endpoint_url="https://direct.example.test/v1/chat/completions", + model="unlisted-direct-model", + model_endpoint_id=None, + endpoint_provenance="direct", + ) + assert chat_helpers._validate_bearer_session_model(session, owner="alice") is None + + +def test_registered_local_endpoint_keeps_explicit_model_without_catalog(monkeypatch): + from routes.chat_helpers import _validate_bearer_session_model + + endpoint = _registered_endpoint(models=None) + endpoint.base_url = "http://localhost:8000/v1" + endpoint.endpoint_kind = "local" + _patch_validator_db(monkeypatch, [endpoint]) + session = _registered_session(model="operator-model") + session.endpoint_url = "http://localhost:8000/v1/chat/completions" + assert _validate_bearer_session_model(session, owner="alice") == "operator-model" + + +def test_unclassified_persisted_session_fails_closed_for_bearer_validation(monkeypatch): + from routes.chat_helpers import _validate_bearer_session_model + + session = SimpleNamespace( + endpoint_url="https://api.example.test/v1/chat/completions", + model="safe-model", + model_endpoint_id=None, + endpoint_provenance=None, + ) + with pytest.raises(HTTPException): + _validate_bearer_session_model(session, owner="alice") + + +def test_session_manager_round_trips_endpoint_provenance(monkeypatch, tmp_path): + import core.session_manager as session_manager_module + from core.session_manager import SessionManager + + session_factory = _isolated_db(tmp_path) + monkeypatch.setattr(session_manager_module, "SessionLocal", session_factory) + manager = SessionManager() + session = manager.create_session( + session_id="durable-sid", + name="durable", + endpoint_url="https://api.example.test/v1/chat/completions", + model="safe-model", + owner="alice", + ) + manager.set_session_endpoint_provenance( + "durable-sid", + model_endpoint_id="ep-1", + endpoint_provenance="registered", + ) + + db = session_factory() + try: + row = db.query(cdb.Session).filter(cdb.Session.id == "durable-sid").first() + assert row.model_endpoint_id == "ep-1" + assert row.endpoint_provenance == "registered" + finally: + db.close() + assert session.model_endpoint_id == "ep-1" + assert session.endpoint_provenance == "registered" + manager.sessions.clear() + reloaded = manager.get_session("durable-sid") + assert reloaded.model_endpoint_id == "ep-1" + assert reloaded.endpoint_provenance == "registered" + + +def test_probe_bearer_patch_rejects_unlisted_model_before_persisting(monkeypatch): + from routes import session_routes as sr + + endpoint = _registered_endpoint() + db_session = SimpleNamespace( + id="sid", + owner="alice", + model="safe-model", + endpoint_url="https://api.example.test/v1/chat/completions", + headers={}, + updated_at=None, + folder=None, + ) + _db = _Db({cdb.Session: [db_session], cdb.ModelEndpoint: [endpoint], None: [db_session]}) + monkeypatch.setattr(sr, "SessionLocal", lambda: _db) + session = _registered_session(model="safe-model") + manager = SimpleNamespace( + get_session=lambda sid: session, + update_session_name=lambda *args, **kwargs: None, + ) + router = sr.setup_session_routes(manager, {}) + patch_session = _endpoint(router, "/api/session/{sid}", "PATCH") + + with pytest.raises(HTTPException) as exc: + patch_session( + request=_Request(), + sid="sid", + model="unsafe-model", + endpoint_url="https://api.example.test/v1/chat/completions", + endpoint_id="ep-1", + ) + assert "permitted" in str(exc.value.detail).lower() + assert session.model == "safe-model" + assert db_session.model == "safe-model" + + +def test_bearer_patch_binds_exact_endpoint_provenance(monkeypatch): + from routes import session_routes as sr + + endpoint = _registered_endpoint() + db_session = SimpleNamespace( + id="sid", + owner="alice", + model="safe-model", + endpoint_url="https://api.example.test/v1/chat/completions", + headers={}, + updated_at=None, + folder=None, + ) + db = _Db({cdb.Session: [db_session], cdb.ModelEndpoint: [endpoint], None: [db_session]}) + monkeypatch.setattr(sr, "SessionLocal", lambda: db) + session = _registered_session(model="safe-model") + manager = SimpleNamespace( + get_session=lambda sid: session, + update_session_name=lambda *args, **kwargs: None, + ) + router = sr.setup_session_routes(manager, {}) + patch_session = _endpoint(router, "/api/session/{sid}", "PATCH") + + patch_session( + request=_Request(), + sid="sid", + model="safe-model", + endpoint_url="https://api.example.test/v1/chat/completions", + endpoint_id="ep-1", + ) + assert session.model_endpoint_id == "ep-1" + assert session.endpoint_provenance == "registered" + assert db_session.model_endpoint_id == "ep-1" + assert db_session.endpoint_provenance == "registered" + + +@pytest.mark.asyncio +async def test_bearer_patch_then_sync_resume_uses_validated_model(monkeypatch): + from routes import session_routes as sr + from routes.webhook import webhook_routes as wr + from src import llm_core + + endpoint = _registered_endpoint() + db_session = SimpleNamespace( + id="sid", + owner="alice", + model="safe-model", + endpoint_url="https://api.example.test/v1/chat/completions", + headers={}, + updated_at=None, + folder=None, + ) + db = _Db({cdb.Session: [db_session], cdb.ModelEndpoint: [endpoint], None: [db_session]}) + monkeypatch.setattr(sr, "SessionLocal", lambda: db) + session = _registered_session(model="safe-model") + session.add_message = lambda message: session.history.append(message) + manager = SimpleNamespace( + get_session=lambda sid: session, + update_session_name=lambda *args, **kwargs: None, + save_sessions=lambda: None, + ) + patch_session = _endpoint(sr.setup_session_routes(manager, {}), "/api/session/{sid}", "PATCH") + patch_session( + request=_Request(), + sid="sid", + model="safe-model", + endpoint_url="https://api.example.test/v1/chat/completions", + endpoint_id="ep-1", + ) + _patch_validator_db(monkeypatch, [endpoint]) + + async def fake_llm(*args, **kwargs): + return "reply" + + monkeypatch.setattr(llm_core, "llm_call_async", fake_llm) + sync_chat = _endpoint( + wr.setup_webhook_routes(SimpleNamespace(), None, session_manager=manager), + "/api/v1/chat", + "POST", + ) + body = SimpleNamespace( + message="hello", + model=None, + session="sid", + api_key=None, + base_url=None, + provider=None, + ) + result = await sync_chat(request=_Request(), body=body) + assert result["model"] == "safe-model" + + +@pytest.mark.asyncio +async def test_probe_bearer_sync_resume_revalidates_persisted_model(monkeypatch): + from routes.webhook import webhook_routes as wr + from src import llm_core + + session = _registered_session() + session.history = [] + session.add_message = lambda message: session.history.append(message) + manager = SimpleNamespace(get_session=lambda sid: session, save_sessions=lambda: None) + _patch_validator_db(monkeypatch, [_registered_endpoint()]) + + async def unexpected_llm(*args, **kwargs): + raise AssertionError("unlisted persisted model reached the LLM") + + monkeypatch.setattr(llm_core, "llm_call_async", unexpected_llm) + router = wr.setup_webhook_routes( + webhook_manager=SimpleNamespace(), + auth_manager=None, + session_manager=manager, + ) + sync_chat = _endpoint(router, "/api/v1/chat", "POST") + body = SimpleNamespace( + message="hello", + model=None, + session="sid", + api_key=None, + base_url=None, + provider=None, + ) + + with pytest.raises(HTTPException): + await sync_chat(request=_Request(), body=body) + + +@pytest.mark.asyncio +async def test_probe_bearer_rewrite_revalidates_before_streaming(monkeypatch): + from routes import chat_routes as cr + + session = _registered_session() + session.history = [] + manager = SimpleNamespace(get_session=lambda sid: session, save_sessions=lambda: None) + _patch_validator_db(monkeypatch, [_registered_endpoint()]) + monkeypatch.setattr(cr, "_verify_session_owner", lambda *args, **kwargs: None) + router = cr.setup_chat_routes(manager, None, None, None, None, None, webhook_manager=None) + rewrite = _endpoint(router, "/api/rewrite", "POST") + + with pytest.raises(HTTPException): + await rewrite( + request=_Request( + body={ + "session_id": "sid", + "original_text": "old", + "instruction": "shorter", + } + ) + ) + + +@pytest.mark.asyncio +async def test_probe_bearer_compaction_aliases_revalidate_before_llm(monkeypatch): + from routes import session_routes as sr + from routes.history import history_routes as hr + from core.models import ChatMessage + from src import llm_core, model_context + + session = _registered_session() + session.history = [ChatMessage("user", f"message {i}") for i in range(6)] + session.get_context_messages = lambda: [{"role": "user", "content": "message"}] + manager = SimpleNamespace( + get_session=lambda sid: session, + replace_messages=lambda *args: True, + save_sessions=lambda: None, + ) + _patch_validator_db(monkeypatch, [_registered_endpoint()]) + monkeypatch.setattr(sr, "_verify_session_owner", lambda *args, **kwargs: None) + monkeypatch.setattr(hr, "_verify_session_owner", lambda *args, **kwargs: None) + monkeypatch.setattr(sr, "_reject_compact_during_active_run", lambda *args, **kwargs: None) + monkeypatch.setattr(hr, "_reject_compact_during_active_run", lambda *args, **kwargs: None) + monkeypatch.setattr(sr, "SessionLocal", lambda: _Db({})) + monkeypatch.setattr(hr, "SessionLocal", lambda: _Db({})) + monkeypatch.setattr(model_context, "get_context_length", lambda *args, **kwargs: 4096) + + async def compact_llm(*args, **kwargs): + return "summary" + + monkeypatch.setattr(llm_core, "llm_call_async", compact_llm) + + session_router = sr.setup_session_routes(manager, {}) + history_router = hr.setup_history_routes(manager) + session_compact = _endpoint(session_router, "/api/session/{session_id}/compact", "POST") + history_compact = _endpoint(history_router, "/api/session/{session_id}/compact", "POST") + + with pytest.raises(HTTPException): + await session_compact(request=_Request(), session_id="sid") + with pytest.raises(HTTPException): + await history_compact(request=_Request(), session_id="sid") + + +@pytest.mark.asyncio +async def test_probe_bearer_gallery_json_reference_serves_owned_binary(monkeypatch, tmp_path): + # app.py normally calls load_dotenv at import time. Replace that call in + # this isolated probe so the probe never reads any .env* file. + import dotenv + + monkeypatch.setattr(dotenv, "load_dotenv", lambda *args, **kwargs: None) + if "app" in sys.modules: + app = sys.modules["app"] + else: + import app # noqa: PLC0415 + + image_path = tmp_path / "image.png" + image_path.write_bytes(b"owned image") + row = SimpleNamespace(filename="image.png", owner="alice") + monkeypatch.setattr(app, "resolve_generated_image_path", lambda filename: image_path) + monkeypatch.setattr(cdb, "SessionLocal", lambda: _Db({cdb.GalleryImage: [row]})) + + response = await app.serve_generated_image("image.png", _Request()) + assert response.path == str(image_path) + + cookie_response = await app.serve_generated_image( + "image.png", _Request(owner="alice", bearer=False) + ) + assert cookie_response.path == str(image_path) + with pytest.raises(HTTPException): + await app.serve_generated_image( + "image.png", _Request(owner="bob", bearer=False) + )