security: bind bearer sessions to endpoint provenance

This commit is contained in:
RaresKeY
2026-08-30 20:03:06 +00:00
parent a09b0b5722
commit c473240131
12 changed files with 1043 additions and 45 deletions
+26 -4
View File
@@ -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 = {
+42
View File
@@ -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()
+2
View File
@@ -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
+50
View File
@@ -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()
+89 -26
View File
@@ -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.
+54 -12
View File
@@ -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. "
+21 -3
View File
@@ -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:
+24
View File
@@ -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))
+33
View File
@@ -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
+30
View File
@@ -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]
+38
View File
@@ -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)
+634
View File
@@ -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)
)