mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-09-10 18:22:20 +02:00
security: bind bearer sessions to endpoint provenance
This commit is contained in:
@@ -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 = {
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
@@ -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. "
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
)
|
||||
Reference in New Issue
Block a user