mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-09-10 18:22:20 +02:00
fix(security): close bearer auxiliary boundaries
This commit is contained in:
@@ -2821,6 +2821,7 @@ def setup_chat_routes(
|
||||
It just asks the LLM to rewrite the given text.
|
||||
"""
|
||||
require_chat_scope(request)
|
||||
capability = build_request_capability(request)
|
||||
try:
|
||||
body = await request.json()
|
||||
except Exception:
|
||||
@@ -2855,6 +2856,9 @@ def setup_chat_routes(
|
||||
async def stream_rewrite() -> AsyncGenerator[str, None]:
|
||||
full_response = ""
|
||||
try:
|
||||
stream_kwargs = {}
|
||||
if not capability.allow_live_probes:
|
||||
stream_kwargs["allow_live_probes"] = False
|
||||
async for chunk in stream_llm(
|
||||
sess.endpoint_url,
|
||||
sess.model,
|
||||
@@ -2867,6 +2871,7 @@ def setup_chat_routes(
|
||||
# on "Rewriting...". Same fix as the chat max_tokens cap.
|
||||
max_tokens=0,
|
||||
tools=None,
|
||||
**stream_kwargs,
|
||||
):
|
||||
if chunk.startswith("data: ") and not chunk.startswith("data: [DONE]"):
|
||||
try:
|
||||
|
||||
@@ -10,10 +10,16 @@ from fastapi import APIRouter, Depends, Request, HTTPException
|
||||
|
||||
from core.models import ChatMessage
|
||||
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 (
|
||||
sanitize_client_message_metadata,
|
||||
sanitize_projected_message_metadata,
|
||||
normalize_client_message_role,
|
||||
)
|
||||
from src.topic_analyzer import analyze_topics
|
||||
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)
|
||||
try:
|
||||
body = await request.json()
|
||||
role = body.get("role", "assistant")
|
||||
role = normalize_client_message_role(body.get("role", "assistant"))
|
||||
content = body.get("content", "")
|
||||
if not content:
|
||||
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.
|
||||
"""
|
||||
require_chat_scope(request)
|
||||
capability = request_capability(request)
|
||||
_verify_session_owner(request, session_id)
|
||||
try:
|
||||
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()
|
||||
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 = max(0.0, min(100.0, pct))
|
||||
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):
|
||||
"""Manually trigger context compaction for a session."""
|
||||
require_chat_scope(request)
|
||||
capability = request_capability(request)
|
||||
_verify_session_owner(request, session_id)
|
||||
from src.auth_helpers import effective_user
|
||||
owner = effective_user(request)
|
||||
@@ -782,7 +797,14 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
if len(session.history) < 6:
|
||||
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()
|
||||
used_before = estimate_tokens(messages_before)
|
||||
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
|
||||
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))
|
||||
compact_kwargs = {}
|
||||
if not capability.allow_live_probes:
|
||||
compact_kwargs["allow_live_probes"] = False
|
||||
summary = await llm_call_async(
|
||||
compact_url, compact_model,
|
||||
[
|
||||
@@ -817,6 +842,7 @@ def setup_history_routes(session_manager, upload_handler=None) -> APIRouter:
|
||||
],
|
||||
temperature=0.2, max_tokens=1024,
|
||||
headers=compact_headers, timeout=30,
|
||||
**compact_kwargs,
|
||||
)
|
||||
summary = normalize_compaction_summary(summary)
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import time
|
||||
from services.search import get_search_config, comprehensive_web_search, PROVIDER_INFO
|
||||
from services.search.core import _call_provider
|
||||
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__)
|
||||
|
||||
@@ -40,11 +40,12 @@ async def _request_values(request: Request) -> Dict[str, Any]:
|
||||
def setup_search_routes(config) -> APIRouter:
|
||||
router = APIRouter(
|
||||
tags=["search"],
|
||||
dependencies=[Depends(require_chat_scope)],
|
||||
dependencies=[Depends(require_interactive_request)],
|
||||
)
|
||||
|
||||
@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()
|
||||
|
||||
@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.
|
||||
"""
|
||||
require_chat_scope(request)
|
||||
require_interactive_request(request)
|
||||
values = await _request_values(request)
|
||||
query = str(values.get("query") or values.get("q") or "").strip()
|
||||
if not query:
|
||||
@@ -71,8 +72,9 @@ def setup_search_routes(config) -> APIRouter:
|
||||
return {"context": "", "sources": [], "error": str(e)}
|
||||
|
||||
@router.get("/api/search/providers")
|
||||
async def list_search_providers():
|
||||
async def list_search_providers(request: Request):
|
||||
"""Return available search providers with config status."""
|
||||
require_interactive_request(request)
|
||||
providers = []
|
||||
for pid, (label, needs_key, needs_url) in PROVIDER_INFO.items():
|
||||
if pid == "disabled":
|
||||
@@ -92,7 +94,7 @@ def setup_search_routes(config) -> APIRouter:
|
||||
@router.post("/api/search/query")
|
||||
async def search_with_provider(request: Request) -> Dict[str, Any]:
|
||||
"""Search using a specific provider. Used by compare search mode."""
|
||||
require_chat_scope(request)
|
||||
require_interactive_request(request)
|
||||
values = await _request_values(request)
|
||||
query = str(values.get("query") or values.get("q") or "").strip()
|
||||
provider = str(values.get("provider") or "").strip()
|
||||
|
||||
+84
-51
@@ -16,9 +16,14 @@ from src.auth_helpers import (
|
||||
_auth_disabled,
|
||||
is_bearer_principal,
|
||||
owner_filter,
|
||||
request_capability,
|
||||
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_actions import is_session_recently_active
|
||||
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
|
||||
# purge exists only to catch ghosts the frontend missed (tab close,
|
||||
# crash). Only clean up rows old enough to be definitely orphaned.
|
||||
try:
|
||||
from datetime import timedelta as _td
|
||||
_cutoff = utcnow_naive() - _td(minutes=10)
|
||||
_purge_db = SessionLocal()
|
||||
# Listing is an owner-scoped read for bearer integrations. The legacy
|
||||
# incognito cleanup query has no owner predicate and would otherwise
|
||||
# let a chat token mutate another user's stale sessions before the
|
||||
# owner-filtered result is assembled. Browser cleanup remains intact.
|
||||
if not is_bearer_principal(request):
|
||||
try:
|
||||
from core.database import ChatMessage as _DbMsg
|
||||
_ghosts = _purge_db.query(DbSession).filter(
|
||||
DbSession.name.in_(("Nobody", "Incognito")),
|
||||
DbSession.created_at < _cutoff,
|
||||
).all()
|
||||
for _g in _ghosts:
|
||||
if active_incognito_id and _g.id == active_incognito_id:
|
||||
continue
|
||||
_purge_db.query(_DbMsg).filter(_DbMsg.session_id == _g.id).delete()
|
||||
_purge_db.delete(_g)
|
||||
if hasattr(session_manager, "delete_session"):
|
||||
try:
|
||||
session_manager.delete_session(_g.id)
|
||||
except Exception:
|
||||
pass
|
||||
if _ghosts:
|
||||
_purge_db.commit()
|
||||
finally:
|
||||
_purge_db.close()
|
||||
except Exception:
|
||||
pass
|
||||
from datetime import timedelta as _td
|
||||
_cutoff = utcnow_naive() - _td(minutes=10)
|
||||
_purge_db = SessionLocal()
|
||||
try:
|
||||
from core.database import ChatMessage as _DbMsg
|
||||
_ghosts = _purge_db.query(DbSession).filter(
|
||||
DbSession.name.in_(("Nobody", "Incognito")),
|
||||
DbSession.created_at < _cutoff,
|
||||
).all()
|
||||
for _g in _ghosts:
|
||||
if active_incognito_id and _g.id == active_incognito_id:
|
||||
continue
|
||||
_purge_db.query(_DbMsg).filter(_DbMsg.session_id == _g.id).delete()
|
||||
_purge_db.delete(_g)
|
||||
if hasattr(session_manager, "delete_session"):
|
||||
try:
|
||||
session_manager.delete_session(_g.id)
|
||||
except Exception:
|
||||
pass
|
||||
if _ghosts:
|
||||
_purge_db.commit()
|
||||
finally:
|
||||
_purge_db.close()
|
||||
except Exception:
|
||||
pass
|
||||
user_sessions = session_manager.get_sessions_for_user(user)
|
||||
# Fetch folder info from DB for each session
|
||||
db = SessionLocal()
|
||||
@@ -354,6 +364,8 @@ def setup_session_routes(
|
||||
endpoint_id: str = Form(""),
|
||||
):
|
||||
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"
|
||||
user = effective_user(request)
|
||||
endpoint_api_key = ""
|
||||
@@ -405,6 +417,7 @@ def setup_session_routes(
|
||||
headers=validation_headers,
|
||||
owner=user,
|
||||
endpoint_id=endpoint_id.strip() if endpoint_id else None,
|
||||
**probe_kwargs,
|
||||
)
|
||||
if not ids:
|
||||
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)]
|
||||
model_to_use = (chat_ids or ids)[0]
|
||||
else:
|
||||
from src.llm_core import list_model_ids
|
||||
import os as _os
|
||||
req_base = _os.path.basename(model_to_use.rstrip("/"))
|
||||
avail = list_model_ids(
|
||||
endpoint_url,
|
||||
timeout=SESSION_MODEL_VALIDATION_TIMEOUT,
|
||||
headers=validation_headers,
|
||||
owner=user,
|
||||
endpoint_id=endpoint_id.strip() if endpoint_id else None,
|
||||
)
|
||||
if not avail:
|
||||
raise HTTPException(400, "Cannot reach /v1/models")
|
||||
if model_to_use not in avail:
|
||||
found = None
|
||||
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
|
||||
# A bearer with an explicit model is already using an owner-scoped
|
||||
# registered endpoint (raw URLs are rejected above). Do not turn
|
||||
# that synchronous session-creation request into a live catalog
|
||||
# probe merely to validate a value the caller supplied. Interactive
|
||||
# requests retain the existing catalog-backed validation.
|
||||
if capability.allow_live_probes:
|
||||
from src.llm_core import list_model_ids
|
||||
import os as _os
|
||||
req_base = _os.path.basename(model_to_use.rstrip("/"))
|
||||
avail = list_model_ids(
|
||||
endpoint_url,
|
||||
timeout=SESSION_MODEL_VALIDATION_TIMEOUT,
|
||||
headers=validation_headers,
|
||||
owner=user,
|
||||
endpoint_id=endpoint_id.strip() if endpoint_id else None,
|
||||
**probe_kwargs,
|
||||
)
|
||||
if not avail:
|
||||
raise HTTPException(400, "Cannot reach /v1/models")
|
||||
if model_to_use not in avail:
|
||||
found = None
|
||||
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())
|
||||
user = effective_user(request)
|
||||
@@ -586,7 +606,7 @@ def setup_session_routes(
|
||||
raise HTTPException(400, "Invalid message attachment metadata") from exc
|
||||
for m in messages:
|
||||
sess.add_message(ChatMessage(
|
||||
m["role"],
|
||||
normalize_client_message_role(m.get("role", "user"), default="user"),
|
||||
m["content"],
|
||||
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):
|
||||
"""Summarize older messages into one compacted history entry."""
|
||||
require_chat_scope(request)
|
||||
capability = request_capability(request)
|
||||
_verify_session_owner(request, session_id)
|
||||
try:
|
||||
session = session_manager.get_session(session_id)
|
||||
@@ -1046,6 +1067,9 @@ def setup_session_routes(
|
||||
for m in older
|
||||
)
|
||||
try:
|
||||
compact_kwargs = {}
|
||||
if not capability.allow_live_probes:
|
||||
compact_kwargs["allow_live_probes"] = False
|
||||
summary = await llm_call_async(
|
||||
url,
|
||||
model,
|
||||
@@ -1054,6 +1078,7 @@ def setup_session_routes(
|
||||
max_tokens=1024,
|
||||
headers=headers,
|
||||
timeout=60,
|
||||
**compact_kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error("Manual compaction failed: %s", e)
|
||||
@@ -1079,7 +1104,10 @@ def setup_session_routes(
|
||||
"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):
|
||||
"""Use AI to categorize all sessions into folders.
|
||||
|
||||
@@ -1089,6 +1117,7 @@ def setup_session_routes(
|
||||
users can clean junk without spending tokens.
|
||||
"""
|
||||
require_chat_scope(request)
|
||||
require_interactive_request(request)
|
||||
from src.llm_core import llm_call
|
||||
user = effective_user(request)
|
||||
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):
|
||||
"""Get the real context length for a session's model from the endpoint."""
|
||||
require_chat_scope(request)
|
||||
capability = request_capability(request)
|
||||
_verify_session_owner(request, session_id)
|
||||
session = session_manager.get_session(session_id)
|
||||
if not session:
|
||||
@@ -1378,7 +1408,10 @@ def setup_session_routes(
|
||||
return {"context_length": None}
|
||||
try:
|
||||
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}
|
||||
except Exception:
|
||||
return {"context_length": None}
|
||||
|
||||
+14
-2
@@ -21,7 +21,12 @@ from core.database import (
|
||||
Note,
|
||||
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.constants import GENERATED_IMAGES_DIR
|
||||
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)
|
||||
current_user = effective_user(request)
|
||||
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:
|
||||
raise HTTPException(403, "Access denied")
|
||||
if file_owner != current_user and not auth_mgr.is_admin(current_user):
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
import uuid
|
||||
import logging
|
||||
import json
|
||||
from typing import Optional
|
||||
|
||||
import httpx
|
||||
@@ -9,7 +10,12 @@ from fastapi import APIRouter, HTTPException, Request, Form
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
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.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
|
||||
|
||||
|
||||
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(
|
||||
webhook_manager: WebhookManager,
|
||||
auth_manager,
|
||||
@@ -240,12 +275,13 @@ def setup_webhook_routes(
|
||||
if getattr(request.state, "api_token", False) is not True:
|
||||
raise HTTPException(403, "This endpoint requires an API token")
|
||||
token_owner = require_chat_scope(request)
|
||||
capability = request_capability(request)
|
||||
if not token_owner:
|
||||
raise HTTPException(403, "API token has no owner")
|
||||
|
||||
from core.models import ChatMessage
|
||||
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()
|
||||
if not message:
|
||||
@@ -337,28 +373,12 @@ def setup_webhook_routes(
|
||||
raise HTTPException(500, "Could not resolve endpoint credentials")
|
||||
|
||||
if model == "auto":
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=5) as client:
|
||||
models_url = build_models_url(base_url)
|
||||
hdrs = build_headers(api_key, base_url)
|
||||
if models_url:
|
||||
resp = await client.get(models_url, headers=hdrs)
|
||||
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")
|
||||
# This route is bearer-only. Resolve auto from the endpoint's
|
||||
# already persisted catalog and leave the provider alias in
|
||||
# place when no cache exists; neither choice needs a new
|
||||
# /models or /tags request during ordinary chat.
|
||||
ids = _cached_endpoint_model_ids(ep)
|
||||
model = ids[0] if ids else "auto"
|
||||
|
||||
if not session_manager:
|
||||
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]
|
||||
|
||||
llm_kwargs = {}
|
||||
if not capability.allow_live_probes:
|
||||
llm_kwargs["allow_live_probes"] = False
|
||||
reply = await llm_call_async(
|
||||
sess.endpoint_url, sess.model, messages,
|
||||
headers=sess.headers, timeout=120,
|
||||
**llm_kwargs,
|
||||
)
|
||||
sess.add_message(ChatMessage("assistant", reply))
|
||||
session_manager.save_sessions()
|
||||
|
||||
@@ -17,6 +17,26 @@ _APPROVAL_PROVENANCE_FIELDS = frozenset({
|
||||
"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):
|
||||
"""Copy metadata while removing fields that can imply approval authority.
|
||||
|
||||
@@ -235,7 +235,7 @@ def _install_sync_chat_stubs(monkeypatch):
|
||||
self.role = role
|
||||
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"
|
||||
|
||||
endpoint_resolver = types.ModuleType("src.endpoint_resolver")
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user