fix(security): close bearer auxiliary boundaries

This commit is contained in:
RaresKeY
2026-08-29 12:27:14 +00:00
parent 31c7249ef3
commit e1da1264dc
9 changed files with 804 additions and 88 deletions
+5
View File
@@ -2821,6 +2821,7 @@ def setup_chat_routes(
It just asks the LLM to rewrite the given text. It just asks the LLM to rewrite the given text.
""" """
require_chat_scope(request) require_chat_scope(request)
capability = build_request_capability(request)
try: try:
body = await request.json() body = await request.json()
except Exception: except Exception:
@@ -2855,6 +2856,9 @@ def setup_chat_routes(
async def stream_rewrite() -> AsyncGenerator[str, None]: async def stream_rewrite() -> AsyncGenerator[str, None]:
full_response = "" full_response = ""
try: try:
stream_kwargs = {}
if not capability.allow_live_probes:
stream_kwargs["allow_live_probes"] = False
async for chunk in stream_llm( async for chunk in stream_llm(
sess.endpoint_url, sess.endpoint_url,
sess.model, sess.model,
@@ -2867,6 +2871,7 @@ def setup_chat_routes(
# on "Rewriting...". Same fix as the chat max_tokens cap. # on "Rewriting...". Same fix as the chat max_tokens cap.
max_tokens=0, max_tokens=0,
tools=None, tools=None,
**stream_kwargs,
): ):
if chunk.startswith("data: ") and not chunk.startswith("data: [DONE]"): if chunk.startswith("data: ") and not chunk.startswith("data: [DONE]"):
try: try:
+30 -4
View File
@@ -10,10 +10,16 @@ from fastapi import APIRouter, Depends, Request, HTTPException
from core.models import ChatMessage from core.models import ChatMessage
from core.database import SessionLocal, ChatMessage as DbChatMessage, Session as DbSession from core.database import SessionLocal, ChatMessage as DbChatMessage, Session as DbSession
from src.auth_helpers import effective_user, is_bearer_principal, require_chat_scope from src.auth_helpers import (
effective_user,
is_bearer_principal,
request_capability,
require_chat_scope,
)
from src.message_metadata import ( from src.message_metadata import (
sanitize_client_message_metadata, sanitize_client_message_metadata,
sanitize_projected_message_metadata, sanitize_projected_message_metadata,
normalize_client_message_role,
) )
from src.topic_analyzer import analyze_topics from src.topic_analyzer import analyze_topics
from src.upload_handler import reserve_message_upload_references from src.upload_handler import reserve_message_upload_references
@@ -311,7 +317,7 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
_verify_session_owner(request, session_id) _verify_session_owner(request, session_id)
try: try:
body = await request.json() body = await request.json()
role = body.get("role", "assistant") role = normalize_client_message_role(body.get("role", "assistant"))
content = body.get("content", "") content = body.get("content", "")
if not content: if not content:
raise HTTPException(400, "content is required") raise HTTPException(400, "content is required")
@@ -720,6 +726,7 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
when the whole chat is approaching compaction. when the whole chat is approaching compaction.
""" """
require_chat_scope(request) require_chat_scope(request)
capability = request_capability(request)
_verify_session_owner(request, session_id) _verify_session_owner(request, session_id)
try: try:
session = session_manager.get_session(session_id) session = session_manager.get_session(session_id)
@@ -731,7 +738,14 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
messages = session.get_context_messages() messages = session.get_context_messages()
used = int(estimate_tokens(messages)) used = int(estimate_tokens(messages))
ctx_len = int(get_context_length(session.endpoint_url, session.model) or 0) context_kwargs = {}
if not capability.allow_live_probes:
context_kwargs["allow_live_probes"] = False
ctx_len = int(get_context_length(
session.endpoint_url,
session.model,
**context_kwargs,
) or 0)
pct = round((used / ctx_len) * 100, 1) if ctx_len else 0.0 pct = round((used / ctx_len) * 100, 1) if ctx_len else 0.0
pct = max(0.0, min(100.0, pct)) pct = max(0.0, min(100.0, pct))
visible_messages = sum( visible_messages = sum(
@@ -765,6 +779,7 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
async def compact_session(request: Request, session_id: str): async def compact_session(request: Request, session_id: str):
"""Manually trigger context compaction for a session.""" """Manually trigger context compaction for a session."""
require_chat_scope(request) require_chat_scope(request)
capability = request_capability(request)
_verify_session_owner(request, session_id) _verify_session_owner(request, session_id)
from src.auth_helpers import effective_user from src.auth_helpers import effective_user
owner = effective_user(request) owner = effective_user(request)
@@ -782,7 +797,14 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
if len(session.history) < 6: if len(session.history) < 6:
return {"status": "ok", "message": "Not enough messages to compact"} return {"status": "ok", "message": "Not enough messages to compact"}
ctx_len = get_context_length(session.endpoint_url, session.model) context_kwargs = {}
if not capability.allow_live_probes:
context_kwargs["allow_live_probes"] = False
ctx_len = get_context_length(
session.endpoint_url,
session.model,
**context_kwargs,
)
messages_before = session.get_context_messages() messages_before = session.get_context_messages()
used_before = estimate_tokens(messages_before) used_before = estimate_tokens(messages_before)
pct_before = round((used_before / ctx_len) * 100, 1) if ctx_len else 0 pct_before = round((used_before / ctx_len) * 100, 1) if ctx_len else 0
@@ -809,6 +831,9 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
from src.context_compactor import SELF_SUMMARY_SYSTEM_PROMPT, normalize_compaction_summary from src.context_compactor import SELF_SUMMARY_SYSTEM_PROMPT, normalize_compaction_summary
compaction_count = sum(1 for m in session.history if isinstance(m, ChatMessage) and "[Conversation summary" in (m.content or "")) compaction_count = sum(1 for m in session.history if isinstance(m, ChatMessage) and "[Conversation summary" in (m.content or ""))
sys_prompt = SELF_SUMMARY_SYSTEM_PROMPT.replace("{count}", str(len(older))).replace("{n}", str(compaction_count + 1)) sys_prompt = SELF_SUMMARY_SYSTEM_PROMPT.replace("{count}", str(len(older))).replace("{n}", str(compaction_count + 1))
compact_kwargs = {}
if not capability.allow_live_probes:
compact_kwargs["allow_live_probes"] = False
summary = await llm_call_async( summary = await llm_call_async(
compact_url, compact_model, compact_url, compact_model,
[ [
@@ -817,6 +842,7 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
], ],
temperature=0.2, max_tokens=1024, temperature=0.2, max_tokens=1024,
headers=compact_headers, timeout=30, headers=compact_headers, timeout=30,
**compact_kwargs,
) )
summary = normalize_compaction_summary(summary) summary = normalize_compaction_summary(summary)
+8 -6
View File
@@ -10,7 +10,7 @@ import time
from services.search import get_search_config, comprehensive_web_search, PROVIDER_INFO from services.search import get_search_config, comprehensive_web_search, PROVIDER_INFO
from services.search.core import _call_provider from services.search.core import _call_provider
from services.search.providers import _get_provider_key, _get_search_instance from services.search.providers import _get_provider_key, _get_search_instance
from src.auth_helpers import require_chat_scope from src.auth_helpers import require_interactive_request
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -40,11 +40,12 @@ async def _request_values(request: Request) -> Dict[str, Any]:
def setup_search_routes(config) -> APIRouter: def setup_search_routes(config) -> APIRouter:
router = APIRouter( router = APIRouter(
tags=["search"], tags=["search"],
dependencies=[Depends(require_chat_scope)], dependencies=[Depends(require_interactive_request)],
) )
@router.get("/api/search/config") @router.get("/api/search/config")
async def get_search_settings() -> Dict[str, Any]: async def get_search_settings(request: Request) -> Dict[str, Any]:
require_interactive_request(request)
return get_search_config() return get_search_config()
@router.post("/api/search") @router.post("/api/search")
@@ -53,7 +54,7 @@ def setup_search_routes(config) -> APIRouter:
Used by Compare mode to pre-search once and share results across panes. Used by Compare mode to pre-search once and share results across panes.
""" """
require_chat_scope(request) require_interactive_request(request)
values = await _request_values(request) values = await _request_values(request)
query = str(values.get("query") or values.get("q") or "").strip() query = str(values.get("query") or values.get("q") or "").strip()
if not query: if not query:
@@ -71,8 +72,9 @@ def setup_search_routes(config) -> APIRouter:
return {"context": "", "sources": [], "error": str(e)} return {"context": "", "sources": [], "error": str(e)}
@router.get("/api/search/providers") @router.get("/api/search/providers")
async def list_search_providers(): async def list_search_providers(request: Request):
"""Return available search providers with config status.""" """Return available search providers with config status."""
require_interactive_request(request)
providers = [] providers = []
for pid, (label, needs_key, needs_url) in PROVIDER_INFO.items(): for pid, (label, needs_key, needs_url) in PROVIDER_INFO.items():
if pid == "disabled": if pid == "disabled":
@@ -92,7 +94,7 @@ def setup_search_routes(config) -> APIRouter:
@router.post("/api/search/query") @router.post("/api/search/query")
async def search_with_provider(request: Request) -> Dict[str, Any]: async def search_with_provider(request: Request) -> Dict[str, Any]:
"""Search using a specific provider. Used by compare search mode.""" """Search using a specific provider. Used by compare search mode."""
require_chat_scope(request) require_interactive_request(request)
values = await _request_values(request) values = await _request_values(request)
query = str(values.get("query") or values.get("q") or "").strip() query = str(values.get("query") or values.get("q") or "").strip()
provider = str(values.get("provider") or "").strip() provider = str(values.get("provider") or "").strip()
+84 -51
View File
@@ -16,9 +16,14 @@ from src.auth_helpers import (
_auth_disabled, _auth_disabled,
is_bearer_principal, is_bearer_principal,
owner_filter, owner_filter,
request_capability,
require_chat_scope, require_chat_scope,
require_interactive_request,
)
from src.message_metadata import (
normalize_client_message_role,
sanitize_client_message_metadata,
) )
from src.message_metadata import sanitize_client_message_metadata
from src.session_image_cleanup import _generated_image_path_for_cleanup, session_image_refs 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.session_actions import is_session_recently_active
from src.upload_handler import reserve_message_upload_references from src.upload_handler import reserve_message_upload_references
@@ -246,32 +251,37 @@ def setup_session_routes(
# session is current and won't delete the live one — this server-side # session is current and won't delete the live one — this server-side
# purge exists only to catch ghosts the frontend missed (tab close, # purge exists only to catch ghosts the frontend missed (tab close,
# crash). Only clean up rows old enough to be definitely orphaned. # crash). Only clean up rows old enough to be definitely orphaned.
try: # Listing is an owner-scoped read for bearer integrations. The legacy
from datetime import timedelta as _td # incognito cleanup query has no owner predicate and would otherwise
_cutoff = utcnow_naive() - _td(minutes=10) # let a chat token mutate another user's stale sessions before the
_purge_db = SessionLocal() # owner-filtered result is assembled. Browser cleanup remains intact.
if not is_bearer_principal(request):
try: try:
from core.database import ChatMessage as _DbMsg from datetime import timedelta as _td
_ghosts = _purge_db.query(DbSession).filter( _cutoff = utcnow_naive() - _td(minutes=10)
DbSession.name.in_(("Nobody", "Incognito")), _purge_db = SessionLocal()
DbSession.created_at < _cutoff, try:
).all() from core.database import ChatMessage as _DbMsg
for _g in _ghosts: _ghosts = _purge_db.query(DbSession).filter(
if active_incognito_id and _g.id == active_incognito_id: DbSession.name.in_(("Nobody", "Incognito")),
continue DbSession.created_at < _cutoff,
_purge_db.query(_DbMsg).filter(_DbMsg.session_id == _g.id).delete() ).all()
_purge_db.delete(_g) for _g in _ghosts:
if hasattr(session_manager, "delete_session"): if active_incognito_id and _g.id == active_incognito_id:
try: continue
session_manager.delete_session(_g.id) _purge_db.query(_DbMsg).filter(_DbMsg.session_id == _g.id).delete()
except Exception: _purge_db.delete(_g)
pass if hasattr(session_manager, "delete_session"):
if _ghosts: try:
_purge_db.commit() session_manager.delete_session(_g.id)
finally: except Exception:
_purge_db.close() pass
except Exception: if _ghosts:
pass _purge_db.commit()
finally:
_purge_db.close()
except Exception:
pass
user_sessions = session_manager.get_sessions_for_user(user) user_sessions = session_manager.get_sessions_for_user(user)
# Fetch folder info from DB for each session # Fetch folder info from DB for each session
db = SessionLocal() db = SessionLocal()
@@ -354,6 +364,8 @@ def setup_session_routes(
endpoint_id: str = Form(""), endpoint_id: str = Form(""),
): ):
require_chat_scope(request) require_chat_scope(request)
capability = request_capability(request)
probe_kwargs = {} if capability.allow_live_probes else {"allow_live_probes": False}
skip_val = str(skip_validation).lower() == "true" skip_val = str(skip_validation).lower() == "true"
user = effective_user(request) user = effective_user(request)
endpoint_api_key = "" endpoint_api_key = ""
@@ -405,6 +417,7 @@ def setup_session_routes(
headers=validation_headers, headers=validation_headers,
owner=user, owner=user,
endpoint_id=endpoint_id.strip() if endpoint_id else None, endpoint_id=endpoint_id.strip() if endpoint_id else None,
**probe_kwargs,
) )
if not ids: if not ids:
raise HTTPException(400, "Cannot reach /v1/models") raise HTTPException(400, "Cannot reach /v1/models")
@@ -416,28 +429,35 @@ def setup_session_routes(
chat_ids = [m for m in ids if not any(p in m.lower() for p in _NON_CHAT)] chat_ids = [m for m in ids if not any(p in m.lower() for p in _NON_CHAT)]
model_to_use = (chat_ids or ids)[0] model_to_use = (chat_ids or ids)[0]
else: else:
from src.llm_core import list_model_ids # A bearer with an explicit model is already using an owner-scoped
import os as _os # registered endpoint (raw URLs are rejected above). Do not turn
req_base = _os.path.basename(model_to_use.rstrip("/")) # that synchronous session-creation request into a live catalog
avail = list_model_ids( # probe merely to validate a value the caller supplied. Interactive
endpoint_url, # requests retain the existing catalog-backed validation.
timeout=SESSION_MODEL_VALIDATION_TIMEOUT, if capability.allow_live_probes:
headers=validation_headers, from src.llm_core import list_model_ids
owner=user, import os as _os
endpoint_id=endpoint_id.strip() if endpoint_id else None, req_base = _os.path.basename(model_to_use.rstrip("/"))
) avail = list_model_ids(
if not avail: endpoint_url,
raise HTTPException(400, "Cannot reach /v1/models") timeout=SESSION_MODEL_VALIDATION_TIMEOUT,
if model_to_use not in avail: headers=validation_headers,
found = None owner=user,
for a in avail: endpoint_id=endpoint_id.strip() if endpoint_id else None,
if _os.path.basename(a.rstrip("/")) == req_base: **probe_kwargs,
found = a )
break if not avail:
if not found: raise HTTPException(400, "Cannot reach /v1/models")
raise HTTPException(400, if model_to_use not in avail:
f"Model not found at server. Available: {', '.join(avail)}") found = None
model_to_use = found for a in avail:
if _os.path.basename(a.rstrip("/")) == req_base:
found = a
break
if not found:
raise HTTPException(400,
f"Model not found at server. Available: {', '.join(avail)}")
model_to_use = found
sid = str(uuid.uuid4()) sid = str(uuid.uuid4())
user = effective_user(request) user = effective_user(request)
@@ -586,7 +606,7 @@ def setup_session_routes(
raise HTTPException(400, "Invalid message attachment metadata") from exc raise HTTPException(400, "Invalid message attachment metadata") from exc
for m in messages: for m in messages:
sess.add_message(ChatMessage( sess.add_message(ChatMessage(
m["role"], normalize_client_message_role(m.get("role", "user"), default="user"),
m["content"], m["content"],
metadata=sanitize_client_message_metadata(m.get("metadata")), metadata=sanitize_client_message_metadata(m.get("metadata")),
)) ))
@@ -1002,6 +1022,7 @@ def setup_session_routes(
async def compact_session(request: Request, session_id: str): async def compact_session(request: Request, session_id: str):
"""Summarize older messages into one compacted history entry.""" """Summarize older messages into one compacted history entry."""
require_chat_scope(request) require_chat_scope(request)
capability = request_capability(request)
_verify_session_owner(request, session_id) _verify_session_owner(request, session_id)
try: try:
session = session_manager.get_session(session_id) session = session_manager.get_session(session_id)
@@ -1046,6 +1067,9 @@ def setup_session_routes(
for m in older for m in older
) )
try: try:
compact_kwargs = {}
if not capability.allow_live_probes:
compact_kwargs["allow_live_probes"] = False
summary = await llm_call_async( summary = await llm_call_async(
url, url,
model, model,
@@ -1054,6 +1078,7 @@ def setup_session_routes(
max_tokens=1024, max_tokens=1024,
headers=headers, headers=headers,
timeout=60, timeout=60,
**compact_kwargs,
) )
except Exception as e: except Exception as e:
logger.error("Manual compaction failed: %s", e) logger.error("Manual compaction failed: %s", e)
@@ -1079,7 +1104,10 @@ def setup_session_routes(
"message_count": len(new_history), "message_count": len(new_history),
} }
@router.post("/sessions/auto-sort") @router.post(
"/sessions/auto-sort",
dependencies=[Depends(require_interactive_request)],
)
def auto_sort_sessions(request: Request, skip_llm: bool = False): def auto_sort_sessions(request: Request, skip_llm: bool = False):
"""Use AI to categorize all sessions into folders. """Use AI to categorize all sessions into folders.
@@ -1089,6 +1117,7 @@ def setup_session_routes(
users can clean junk without spending tokens. users can clean junk without spending tokens.
""" """
require_chat_scope(request) require_chat_scope(request)
require_interactive_request(request)
from src.llm_core import llm_call from src.llm_core import llm_call
user = effective_user(request) user = effective_user(request)
single_user_mode = not user and _auth_disabled() single_user_mode = not user and _auth_disabled()
@@ -1370,6 +1399,7 @@ def setup_session_routes(
async def get_context_info(request: Request, session_id: str): async def get_context_info(request: Request, session_id: str):
"""Get the real context length for a session's model from the endpoint.""" """Get the real context length for a session's model from the endpoint."""
require_chat_scope(request) require_chat_scope(request)
capability = request_capability(request)
_verify_session_owner(request, session_id) _verify_session_owner(request, session_id)
session = session_manager.get_session(session_id) session = session_manager.get_session(session_id)
if not session: if not session:
@@ -1378,7 +1408,10 @@ def setup_session_routes(
return {"context_length": None} return {"context_length": None}
try: try:
from src.model_context import get_context_length from src.model_context import get_context_length
ctx = get_context_length(session.endpoint_url, session.model) context_kwargs = {}
if not capability.allow_live_probes:
context_kwargs["allow_live_probes"] = False
ctx = get_context_length(session.endpoint_url, session.model, **context_kwargs)
return {"context_length": ctx, "model": session.model} return {"context_length": ctx, "model": session.model}
except Exception: except Exception:
return {"context_length": None} return {"context_length": None}
+14 -2
View File
@@ -21,7 +21,12 @@ from core.database import (
Note, Note,
Session as DbSession, Session as DbSession,
) )
from src.auth_helpers import effective_user, require_chat_scope, require_non_bearer_request from src.auth_helpers import (
effective_user,
is_bearer_principal,
require_chat_scope,
require_non_bearer_request,
)
from src.attachment_refs import attachment_refs_from_metadata from src.attachment_refs import attachment_refs_from_metadata
from src.constants import GENERATED_IMAGES_DIR from src.constants import GENERATED_IMAGES_DIR
from src.upload_handler import ( from src.upload_handler import (
@@ -379,7 +384,14 @@ def setup_upload_routes(upload_handler):
auth_configured = bool(auth_mgr and auth_mgr.is_configured) auth_configured = bool(auth_mgr and auth_mgr.is_configured)
current_user = effective_user(request) current_user = effective_user(request)
file_owner = info.get("owner") if info else None file_owner = info.get("owner") if info else None
if auth_configured: if is_bearer_principal(request):
# A token owner is an owner-bound data principal, even when that
# owner is an administrator. Do not reuse the browser admin
# fallback for bearer downloads or an admin token can read another
# user's upload by ID.
if not current_user or file_owner != current_user:
raise HTTPException(404, "File not found")
elif auth_configured:
if not current_user: if not current_user:
raise HTTPException(403, "Access denied") raise HTTPException(403, "Access denied")
if file_owner != current_user and not auth_mgr.is_admin(current_user): if file_owner != current_user and not auth_mgr.is_admin(current_user):
+48 -24
View File
@@ -2,6 +2,7 @@
import uuid import uuid
import logging import logging
import json
from typing import Optional from typing import Optional
import httpx import httpx
@@ -9,7 +10,12 @@ from fastapi import APIRouter, HTTPException, Request, Form
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
from core.database import SessionLocal, Webhook, ModelEndpoint from core.database import SessionLocal, Webhook, ModelEndpoint
from src.auth_helpers import is_bearer_principal, owner_filter, require_chat_scope from src.auth_helpers import (
is_bearer_principal,
owner_filter,
request_capability,
require_chat_scope,
)
from src.url_security import validate_public_http_url from src.url_security import validate_public_http_url
from src.webhook_manager import WebhookManager, validate_webhook_url, validate_events from src.webhook_manager import WebhookManager, validate_webhook_url, validate_events
@@ -62,6 +68,35 @@ def _caller_owns_session(sess_owner, caller) -> bool:
return sess_owner == caller return sess_owner == caller
def _cached_endpoint_model_ids(endpoint) -> list[str]:
"""Return model IDs already stored for a configured endpoint.
The synchronous bearer integration may use a cached model or the provider's
``auto`` alias, but it must not turn an ordinary chat request into a remote
catalog probe. Malformed/legacy cache shapes are treated as empty.
"""
raw = getattr(endpoint, "cached_models", None)
if not raw:
return []
try:
value = json.loads(raw) if isinstance(raw, str) else raw
except (TypeError, ValueError):
return []
if isinstance(value, dict):
value = value.get("data") or value.get("models") or []
if not isinstance(value, list):
return []
ids = []
for item in value:
if isinstance(item, str) and item.strip():
ids.append(item.strip())
elif isinstance(item, dict):
model_id = item.get("id") or item.get("name") or item.get("model")
if isinstance(model_id, str) and model_id.strip():
ids.append(model_id.strip())
return ids
def setup_webhook_routes( def setup_webhook_routes(
webhook_manager: WebhookManager, webhook_manager: WebhookManager,
auth_manager, auth_manager,
@@ -240,12 +275,13 @@ def setup_webhook_routes(
if getattr(request.state, "api_token", False) is not True: if getattr(request.state, "api_token", False) is not True:
raise HTTPException(403, "This endpoint requires an API token") raise HTTPException(403, "This endpoint requires an API token")
token_owner = require_chat_scope(request) token_owner = require_chat_scope(request)
capability = request_capability(request)
if not token_owner: if not token_owner:
raise HTTPException(403, "API token has no owner") raise HTTPException(403, "API token has no owner")
from core.models import ChatMessage from core.models import ChatMessage
from src.llm_core import llm_call_async from src.llm_core import llm_call_async
from src.endpoint_resolver import build_chat_url, build_headers, build_models_url, normalize_base from src.endpoint_resolver import build_chat_url, build_headers, normalize_base
message = body.message.strip() message = body.message.strip()
if not message: if not message:
@@ -337,28 +373,12 @@ def setup_webhook_routes(
raise HTTPException(500, "Could not resolve endpoint credentials") raise HTTPException(500, "Could not resolve endpoint credentials")
if model == "auto": if model == "auto":
try: # This route is bearer-only. Resolve auto from the endpoint's
async with httpx.AsyncClient(timeout=5) as client: # already persisted catalog and leave the provider alias in
models_url = build_models_url(base_url) # place when no cache exists; neither choice needs a new
hdrs = build_headers(api_key, base_url) # /models or /tags request during ordinary chat.
if models_url: ids = _cached_endpoint_model_ids(ep)
resp = await client.get(models_url, headers=hdrs) model = ids[0] if ids else "auto"
resp.raise_for_status()
data = resp.json()
items = data if isinstance(data, list) else (data.get("data") or [])
ids = [m.get("id") for m in items if isinstance(m, dict) and m.get("id")]
if not ids and isinstance(data, dict):
ids = [
m.get("name") or m.get("model")
for m in (data.get("models") or [])
if m.get("name") or m.get("model")
]
else:
import json as _json
ids = _json.loads(ep.cached_models or "[]")
model = ids[0] if ids else "auto"
except Exception:
raise HTTPException(500, "Could not discover models from endpoint")
if not session_manager: if not session_manager:
raise HTTPException(500, "Session manager not available") raise HTTPException(500, "Session manager not available")
@@ -378,9 +398,13 @@ def setup_webhook_routes(
messages = [{"role": m.role, "content": m.content} for m in sess.history] messages = [{"role": m.role, "content": m.content} for m in sess.history]
llm_kwargs = {}
if not capability.allow_live_probes:
llm_kwargs["allow_live_probes"] = False
reply = await llm_call_async( reply = await llm_call_async(
sess.endpoint_url, sess.model, messages, sess.endpoint_url, sess.model, messages,
headers=sess.headers, timeout=120, headers=sess.headers, timeout=120,
**llm_kwargs,
) )
sess.add_message(ChatMessage("assistant", reply)) sess.add_message(ChatMessage("assistant", reply))
session_manager.save_sessions() session_manager.save_sessions()
+20
View File
@@ -17,6 +17,26 @@ _APPROVAL_PROVENANCE_FIELDS = frozenset({
"session_id", "session_id",
}) })
_CLIENT_MESSAGE_ROLES = frozenset({"user", "assistant"})
def normalize_client_message_role(role: Any, *, default: str = "assistant") -> str:
"""Return a non-privileged role for a client-supplied message.
Durable ``system`` and ``tool`` records are still valid when created by
trusted server paths. Client ingress has no such provenance, so only the
ordinary conversation roles are accepted; every other value is demoted to
``user`` rather than becoming model-control metadata.
"""
if not isinstance(role, str):
return "user"
normalized = role.strip().casefold()
if normalized in _CLIENT_MESSAGE_ROLES:
return normalized
if normalized == "" and default in _CLIENT_MESSAGE_ROLES:
return default
return "user"
def _scrub_approval_metadata(value: Any, *, projection: bool, in_approval: bool = False): def _scrub_approval_metadata(value: Any, *, projection: bool, in_approval: bool = False):
"""Copy metadata while removing fields that can imply approval authority. """Copy metadata while removing fields that can imply approval authority.
+1 -1
View File
@@ -235,7 +235,7 @@ def _install_sync_chat_stubs(monkeypatch):
self.role = role self.role = role
self.content = content self.content = content
async def _llm_call_async(endpoint_url, model, messages, headers=None, timeout=None): async def _llm_call_async(endpoint_url, model, messages, headers=None, timeout=None, **kwargs):
return "mocked response" return "mocked response"
endpoint_resolver = types.ModuleType("src.endpoint_resolver") endpoint_resolver = types.ModuleType("src.endpoint_resolver")
+594
View File
@@ -0,0 +1,594 @@
"""Cycle-4 regressions for the API-token chat capability boundary.
These tests deliberately call route endpoints directly as well as exercising
the same request state that the auth middleware stamps. Router dependencies
are useful defense in depth, but they must not be the only authorization
check on a callable FastAPI endpoint.
"""
import asyncio
import json
from datetime import timedelta
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.database import ModelEndpoint, Session as DbSession
from core.models import ChatMessage, Session
class _Request:
def __init__(
self,
*,
bearer=True,
owner="alice",
scopes=("chat",),
current_user="api",
body=None,
auth_manager=None,
query_params=None,
):
self.state = SimpleNamespace(
api_token=bearer,
api_token_owner=owner if bearer else None,
api_token_scopes=list(scopes),
current_user=current_user,
)
self.app = SimpleNamespace(
state=SimpleNamespace(auth_manager=auth_manager)
)
self.headers = {}
self.query_params = query_params or {}
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}")
def _isolated_db(tmp_path):
engine = create_engine(
f"sqlite:///{tmp_path / 'cycle4.db'}",
connect_args={"check_same_thread": False},
poolclass=NullPool,
)
cdb.Base.metadata.create_all(engine)
return sessionmaker(bind=engine, autoflush=False, autocommit=False)
@pytest.mark.asyncio
async def test_search_alias_rejects_chat_bearer_on_every_standalone_entry_point(monkeypatch):
from routes import search_routes as alias_routes
from routes.search import search_routes
# The flat import is a sys.modules shim; use it as the exercised entry
# point so a future alias split cannot silently lose the gate.
router = alias_routes.setup_search_routes(None)
assert alias_routes is search_routes
monkeypatch.setattr(search_routes, "comprehensive_web_search", lambda *a, **k: ("hit", []))
monkeypatch.setattr(search_routes, "_call_provider", lambda *a, **k: [{"title": "hit"}])
request = _Request()
for path, kwargs in (
("/api/search/config", {}),
("/api/search/providers", {}),
("/api/search", {}),
("/api/search/query", {}),
):
endpoint = _endpoint(router, path, "GET" if path.endswith(("config", "providers")) else "POST")
with pytest.raises(HTTPException) as exc:
await endpoint(request=request, **kwargs)
assert exc.value.status_code == 403, path
def test_auto_sort_direct_handler_rejects_bearer_before_owner_side_effects(monkeypatch):
from routes import session_routes as sr
def unexpected(*args, **kwargs):
raise AssertionError("bearer reached auto-sort side effects")
manager = SimpleNamespace(
get_sessions_for_user=unexpected,
delete_session=unexpected,
)
router = sr.setup_session_routes(manager, {})
auto_sort = _endpoint(router, "/api/sessions/auto-sort", "POST")
with pytest.raises(HTTPException) as exc:
auto_sort(request=_Request(), skip_llm=True)
assert exc.value.status_code == 403
def test_admin_owned_bearer_cannot_use_browser_admin_upload_fallback(tmp_path, monkeypatch):
from routes import upload_routes as ur
file_id = "b" * 32 + ".png"
file_path = tmp_path / file_id
file_path.write_bytes(b"private upload")
class _AuthManager:
is_configured = True
def is_admin(self, user):
return user == "admin"
handler = SimpleNamespace(
upload_dir=str(tmp_path),
validate_upload_id=lambda value: value == file_id,
_load_upload_index=lambda: {
"bob:file": {
"id": file_id,
"name": "bob.png",
"mime": "image/png",
"owner": "bob",
}
},
)
router, _cleanup = ur.setup_upload_routes(handler)
download = _endpoint(router, "/api/upload/{file_id}", "GET")
request = _Request(
owner="admin",
current_user="api",
auth_manager=_AuthManager(),
)
with pytest.raises(HTTPException) as exc:
asyncio.run(download(request, file_id))
assert exc.value.status_code == 404
def test_bearer_session_listing_does_not_purge_other_users_incognito_rows(tmp_path, monkeypatch):
from routes import session_routes as sr
ts = _isolated_db(tmp_path)
monkeypatch.setattr(sr, "SessionLocal", ts)
db = ts()
try:
db.query(DbSession).delete()
old = cdb.utcnow_naive() - timedelta(hours=2)
ghost_id = "ghost-" + "a" * 8
owner_id = "owner-" + "b" * 8
db.add(DbSession(
id=ghost_id,
owner="bob",
name="Nobody",
endpoint_url="http://localhost",
model="model",
archived=False,
created_at=old,
updated_at=old,
))
db.add(DbSession(
id=owner_id,
owner="alice",
name="Alice chat",
endpoint_url="http://localhost",
model="model",
archived=False,
))
db.commit()
finally:
db.close()
visible = SimpleNamespace(
id=owner_id,
owner="alice",
name="Alice chat",
model="model",
endpoint_url="http://localhost",
rag=False,
archived=False,
)
manager = SimpleNamespace(
get_sessions_for_user=lambda owner: {owner_id: visible},
)
router = sr.setup_session_routes(manager, {})
list_sessions = _endpoint(router, "/api/sessions", "GET")
result = list_sessions(request=_Request(query_params={"active_incognito_id": ""}))
assert {item["id"] for item in result} == {owner_id}
db = ts()
try:
assert db.query(DbSession).filter(DbSession.id == ghost_id).first() is not None
finally:
db.close()
@pytest.mark.asyncio
@pytest.mark.parametrize("role", ["system", "tool"])
async def test_history_message_ingress_normalizes_privileged_client_roles(monkeypatch, role):
from routes.history import history_routes as hr
monkeypatch.setattr(hr, "_verify_session_owner", lambda *args, **kwargs: None)
stored = []
manager = SimpleNamespace(add_message=lambda sid, message: stored.append(message))
router = hr.setup_history_routes(manager)
add_message = _endpoint(router, "/api/session/{session_id}/message", "POST")
request = _Request(body={"role": role, "content": "client content"})
result = await add_message(request=request, session_id="sid")
assert result == {"status": "ok"}
assert stored[-1].role == "user"
@pytest.mark.asyncio
@pytest.mark.parametrize("role", ["system", "tool"])
async def test_bulk_message_ingress_normalizes_privileged_client_roles(monkeypatch, role):
from routes import session_routes as sr
monkeypatch.setattr(sr, "_verify_session_owner", lambda *args, **kwargs: None)
stored = []
session = SimpleNamespace(add_message=lambda message: stored.append(message))
manager = SimpleNamespace(
get_session=lambda sid: session,
save_sessions=lambda: None,
)
router = sr.setup_session_routes(manager, {})
inject = _endpoint(router, "/api/session/{sid}/inject_messages", "POST")
request = _Request(body={"messages": [{"role": role, "content": "client content"}]})
result = await inject(request=request, sid="sid")
assert result == {"ok": True, "count": 1}
assert stored[-1].role == "user"
def test_server_owned_system_and_tool_messages_remain_available_to_context():
session = Session(
id="sid",
name="chat",
endpoint_url="",
model="",
history=[
ChatMessage("system", "server policy"),
ChatMessage("tool", "server result"),
],
)
assert [message["role"] for message in session.get_context_messages()] == ["system", "tool"]
def test_session_creation_passes_bearer_no_live_capability_to_model_validation(monkeypatch):
from routes import session_routes as sr
from src import llm_core
monkeypatch.setattr(sr, "_reject_raw_endpoint_url_for_non_admin", lambda *args, **kwargs: None)
seen = {}
def list_model_ids(*args, **kwargs):
seen.update(kwargs)
return ["chosen"]
monkeypatch.setattr(llm_core, "list_model_ids", list_model_ids)
manager = SimpleNamespace(
create_session=lambda **kwargs: SimpleNamespace(
id=kwargs["session_id"],
name=kwargs["name"],
model=kwargs["model"],
endpoint_url=kwargs["endpoint_url"],
rag=kwargs["rag"],
headers={},
),
)
router = sr.setup_session_routes(manager, {})
create_session = _endpoint(router, "/api/session", "POST")
result = create_session(
request=_Request(),
name="chat",
endpoint_url="https://api.example.test/v1/chat/completions",
model="",
rag=None,
skip_validation=None,
api_key="",
endpoint_id="",
)
assert result.model == "chosen"
assert seen["allow_live_probes"] is False
def test_explicit_bearer_model_does_not_require_live_setup_probe(monkeypatch):
from routes import session_routes as sr
from src import llm_core
endpoint = SimpleNamespace(
id="ep",
is_enabled=True,
base_url="https://api.example.test/v1",
api_key=None,
)
monkeypatch.setattr(sr, "SessionLocal", lambda: _EndpointDb(endpoint))
def unexpected(*args, **kwargs):
raise AssertionError("explicit bearer model triggered setup catalog probe")
monkeypatch.setattr(llm_core, "list_model_ids", unexpected)
manager = SimpleNamespace(
create_session=lambda **kwargs: SimpleNamespace(
id=kwargs["session_id"],
name=kwargs["name"],
model=kwargs["model"],
endpoint_url=kwargs["endpoint_url"],
rag=kwargs["rag"],
headers={},
),
)
router = sr.setup_session_routes(manager, {})
create_session = _endpoint(router, "/api/session", "POST")
result = create_session(
request=_Request(),
name="chat",
endpoint_url="",
model="explicit-model",
rag=None,
skip_validation=None,
api_key="",
endpoint_id="ep",
)
assert result.model == "explicit-model"
@pytest.mark.asyncio
async def test_context_usage_and_context_info_pass_bearer_no_live_capability(monkeypatch):
from routes import session_routes as sr
from routes.history import history_routes as hr
from src import model_context
session = SimpleNamespace(
endpoint_url="http://127.0.0.1:8080/v1/chat/completions",
model="local-model",
history=[ChatMessage("user", "hello")],
get_context_messages=lambda: [{"role": "user", "content": "hello"}],
)
monkeypatch.setattr(hr, "_verify_session_owner", lambda *args, **kwargs: None)
monkeypatch.setattr(sr, "_verify_session_owner", lambda *args, **kwargs: None)
hr_seen = []
sr_seen = []
def history_context(*args, **kwargs):
hr_seen.append(kwargs)
return 4096
def session_context(*args, **kwargs):
sr_seen.append(kwargs)
return 4096
# Both route modules import this helper lazily, so the same patched helper
# proves each route family forwards the capability independently.
monkeypatch.setattr(model_context, "get_context_length", history_context)
history_manager = SimpleNamespace(get_session=lambda sid: session)
session_manager = SimpleNamespace(get_session=lambda sid: session)
history_router = hr.setup_history_routes(history_manager)
session_router = sr.setup_session_routes(session_manager, {})
# The first call records history's /context path; switch the shared patch
# after it so the second route's call is separately attributable.
history_context_endpoint = _endpoint(history_router, "/api/session/{session_id}/context", "GET")
await history_context_endpoint(request=_Request(), session_id="sid")
monkeypatch.setattr(model_context, "get_context_length", session_context)
info_endpoint = _endpoint(session_router, "/api/session/{session_id}/context_info", "GET")
await info_endpoint(request=_Request(), session_id="sid")
assert hr_seen == [{"allow_live_probes": False}]
assert sr_seen == [{"allow_live_probes": False}]
class _NoopDb:
def query(self, *args, **kwargs):
return self
def filter(self, *args, **kwargs):
return self
def order_by(self, *args, **kwargs):
return self
def all(self):
return []
def first(self):
return None
def add(self, *args, **kwargs):
return None
def commit(self):
return None
def rollback(self):
return None
def close(self):
return None
@pytest.mark.asyncio
async def test_bearer_compaction_routes_forward_no_live_capability(monkeypatch):
from routes import session_routes as sr
from routes.history import history_routes as hr
from src import endpoint_resolver, llm_core, model_context
history = [ChatMessage("user", f"message {i}") for i in range(6)]
session = SimpleNamespace(
id="sid",
owner="alice",
endpoint_url="https://api.example.test/v1/chat/completions",
model="model",
headers={},
history=list(history),
get_context_messages=lambda: [{"role": "user", "content": "message"}],
)
monkeypatch.setattr(hr, "_verify_session_owner", lambda *args, **kwargs: None)
monkeypatch.setattr(sr, "_verify_session_owner", lambda *args, **kwargs: None)
monkeypatch.setattr(hr, "_reject_compact_during_active_run", lambda *args, **kwargs: None)
monkeypatch.setattr(sr, "_reject_compact_during_active_run", lambda *args, **kwargs: None)
monkeypatch.setattr(endpoint_resolver, "resolve_endpoint", lambda *args, **kwargs: (None, None, None))
monkeypatch.setattr(hr, "SessionLocal", lambda: _NoopDb())
monkeypatch.setattr(sr, "SessionLocal", lambda: _NoopDb())
context_seen = []
llm_seen = []
def context_length(*args, **kwargs):
context_seen.append(kwargs)
return 4096
async def llm_call_async(*args, **kwargs):
llm_seen.append(kwargs)
return "summary"
monkeypatch.setattr(model_context, "get_context_length", context_length)
monkeypatch.setattr(llm_core, "llm_call_async", llm_call_async)
history_manager = SimpleNamespace(save_sessions=lambda: None)
history_manager.get_session = lambda sid: session
history_router = hr.setup_history_routes(history_manager)
history_compact = _endpoint(history_router, "/api/session/{session_id}/compact", "POST")
await history_compact(request=_Request(), session_id="sid")
session.history = list(history)
session_manager = SimpleNamespace(
get_session=lambda sid: session,
replace_messages=lambda *args: True,
)
session_router = sr.setup_session_routes(session_manager, {})
session_compact = _endpoint(session_router, "/api/session/{session_id}/compact", "POST")
await session_compact(request=_Request(), session_id="sid")
# The history compactor asks for context directly. The session-route
# compactor delegates context sizing to llm_call_async, so its explicit
# capability is asserted on the two LLM calls below.
assert context_seen == [{"allow_live_probes": False}]
assert len(llm_seen) == 2
assert all(call["allow_live_probes"] is False for call in llm_seen)
@pytest.mark.asyncio
async def test_rewrite_direct_handler_passes_bearer_no_live_capability(monkeypatch):
from routes import chat_routes as cr
monkeypatch.setattr(cr, "_verify_session_owner", lambda *args, **kwargs: None)
seen = {}
async def stream_llm(*args, **kwargs):
seen.update(kwargs)
yield 'data: {"delta":"rewritten"}\n\n'
yield "data: [DONE]\n\n"
monkeypatch.setattr(cr, "stream_llm", stream_llm)
monkeypatch.setattr(cr, "SessionLocal", lambda: _NoopDb())
session = SimpleNamespace(
endpoint_url="https://api.example.test/v1/chat/completions",
model="model",
headers={},
history=[ChatMessage("assistant", "old")],
)
manager = SimpleNamespace(
get_session=lambda sid: session,
save_sessions=lambda: None,
)
router = cr.setup_chat_routes(manager, None, None, None, None, None, webhook_manager=None)
rewrite = _endpoint(router, "/api/rewrite", "POST")
response = await rewrite(
request=_Request(body={
"session_id": "sid",
"original_text": "old",
"instruction": "shorter",
})
)
_chunks = [chunk async for chunk in response.body_iterator]
assert seen["allow_live_probes"] is False
class _EndpointDb:
def __init__(self, endpoint):
self.endpoint = endpoint
def query(self, *args, **kwargs):
return self
def filter(self, *args, **kwargs):
return self
def order_by(self, *args, **kwargs):
return self
def first(self):
return self.endpoint
def close(self):
return None
@pytest.mark.asyncio
async def test_sync_chat_fallback_uses_cached_models_without_provider_probe(monkeypatch):
from routes import webhook_routes as wr
from src import llm_core
endpoint = SimpleNamespace(
owner="alice",
is_enabled=True,
created_at=1,
base_url="http://127.0.0.1:11434/v1",
api_key="configured-key",
cached_models=json.dumps(["cached-model"]),
provider_auth_id=None,
)
monkeypatch.setattr(wr, "SessionLocal", lambda: _EndpointDb(endpoint))
monkeypatch.setattr(wr, "validate_public_http_url", lambda url: url)
class _ForbiddenHttpClient:
def __init__(self, *args, **kwargs):
raise AssertionError("bearer fallback attempted a model-list probe")
monkeypatch.setattr(wr.httpx, "AsyncClient", _ForbiddenHttpClient)
seen = {}
async def llm_call_async(*args, **kwargs):
seen.update(kwargs)
return "reply"
monkeypatch.setattr(llm_core, "llm_call_async", llm_call_async)
class _Session:
def __init__(self, **kwargs):
self.endpoint_url = kwargs["endpoint_url"]
self.model = kwargs["model"]
self.headers = {}
self.history = []
def add_message(self, message):
self.history.append(message)
manager = SimpleNamespace(
create_session=lambda **kwargs: _Session(**kwargs),
save_sessions=lambda: None,
)
webhook_manager = SimpleNamespace(fire_and_forget=lambda *args, **kwargs: None)
router = wr.setup_webhook_routes(webhook_manager, None, session_manager=manager)
sync_chat = _endpoint(router, "/api/v1/chat", "POST")
body = SimpleNamespace(
message="hello",
model=None,
session=None,
api_key=None,
base_url=None,
provider=None,
)
result = await sync_chat(request=_Request(), body=body)
assert result["model"] == "cached-model"
assert seen["allow_live_probes"] is False