unsloth/studio/backend/auth/authentication.py
Roland Tannous 92994e8b83 Studio: add RAG with hybrid search, reranker, chat integration
Backend (studio/backend/):
- core/rag/: parsers (PDF/TXT/MD/DOCX/HTML via pypdf/python-docx/bs4),
  recursive token-aware chunker, embeddings singleton via
  FastSentenceTransformer.from_pretrained(for_inference=True), Qdrant
  local vector store, bm25s lexical index, RRF hybrid retrieval,
  spawn-subprocess ingestion job with SSE progress, optional
  CrossEncoder reranker (off-by-default).
- routes/rag.py: KB CRUD, doc upload (KB + per-thread), doc list/delete,
  ingestion SSE, hybrid+rerank search, thread-index list/clear.
- routes/chat_history.py: purge thread RAG artifacts on thread delete
  and clear-all (rag_documents has no FK cascade to chat_threads so
  uploads work on un-persisted threads).
- studio.db gains 4 RAG tables; storage_roots gains rag_*() helpers.
- auth/authentication.py: get_current_subject_sse accepts ?token=... so
  EventSource can stream ingestion progress.

Frontend (studio/frontend/):
- features/rag/: api client, Zustand store, hooks, dropzone, KB list,
  doc rows, ingestion-progress, thread-index list components.
- Settings dialog gains a Knowledge Bases tab (master/detail + thread
  documents list); /knowledge-bases deep-links to it.
- features/chat/: per-thread ragSource/enableRerank/ragTopK state in
  chat-runtime-store; Retrieval section in chat-settings-sheet with KB
  DropdownMenu (active highlight + per-row trash), thread doc list with
  Clear-thread-index button, RAG Top K slider, reranker toggle;
  chat-adapter retrieves before /v1/chat/completions and injects hits
  as a system block; shared-composer + button routes documents into
  pendingDocs (auto-uploads, send blocked while indexing).
2026-05-23 18:46:15 +04:00

254 lines
7.6 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import secrets
from datetime import datetime, timedelta, timezone
from typing import Optional, Tuple
from fastapi import Depends, HTTPException, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
import jwt
from .storage import (
API_KEY_PREFIX,
get_jwt_secret,
get_user_and_secret,
load_jwt_secret,
save_refresh_token,
validate_api_key,
verify_refresh_token,
)
ALGORITHM = "HS256"
ACCESS_TOKEN_EXPIRE_MINUTES = 60
REFRESH_TOKEN_EXPIRE_DAYS = 7
security = HTTPBearer() # Reads Authorization: Bearer <token>
def _get_secret_for_subject(subject: str) -> str:
secret = get_jwt_secret(subject)
if secret is None:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Invalid or expired token",
)
return secret
def _decode_subject_without_verification(token: str) -> Optional[str]:
try:
payload = jwt.decode(
token,
options = {"verify_signature": False, "verify_exp": False},
)
except jwt.InvalidTokenError:
return None
subject = payload.get("sub")
return subject if isinstance(subject, str) else None
def create_access_token(
subject: str,
expires_delta: Optional[timedelta] = None,
*,
desktop: bool = False,
) -> str:
"""
Create a signed JWT for the given subject (e.g. username).
Tokens are valid across restarts because the signing secret is stored in SQLite.
"""
to_encode = {"sub": subject}
if desktop:
to_encode["desktop"] = True
expire = datetime.now(timezone.utc) + (
expires_delta or timedelta(minutes = ACCESS_TOKEN_EXPIRE_MINUTES)
)
to_encode.update({"exp": expire})
return jwt.encode(
to_encode,
_get_secret_for_subject(subject),
algorithm = ALGORITHM,
)
def is_desktop_access_token(token: str) -> bool:
"""Return true only for a valid desktop-issued JWT access token."""
if token.startswith(API_KEY_PREFIX):
return False
subject = _decode_subject_without_verification(token)
if subject is None:
return False
record = get_user_and_secret(subject)
if record is None:
return False
_salt, _pwd_hash, jwt_secret, _must_change_password = record
try:
payload = jwt.decode(token, jwt_secret, algorithms = [ALGORITHM])
except jwt.InvalidTokenError:
return False
return payload.get("sub") == subject and payload.get("desktop") is True
def create_refresh_token(subject: str, *, desktop: bool = False) -> str:
"""
Create a random refresh token, store its hash in SQLite, and return it.
Refresh tokens are opaque (not JWTs) and expire after REFRESH_TOKEN_EXPIRE_DAYS.
"""
token = secrets.token_urlsafe(48)
expires_at = datetime.now(timezone.utc) + timedelta(days = REFRESH_TOKEN_EXPIRE_DAYS)
save_refresh_token(token, subject, expires_at.isoformat(), is_desktop = desktop)
return token
def refresh_access_token(
refresh_token: str,
) -> Tuple[Optional[str], Optional[str], bool]:
"""
Validate a refresh token and issue a new access token.
The refresh token itself is NOT consumed — it stays valid until expiry.
Returns a new access_token or None if the refresh token is invalid/expired.
"""
verified = verify_refresh_token(refresh_token)
if verified is None:
return None, None, False
username, is_desktop = verified
return (
create_access_token(subject = username, desktop = is_desktop),
username,
is_desktop,
)
def reload_secret() -> None:
"""
Keep legacy API compatibility for callers expecting auth storage init.
Auth now resolves the current signing secret directly from SQLite.
"""
load_jwt_secret()
async def get_current_subject(
credentials: HTTPAuthorizationCredentials = Depends(security),
) -> str:
"""Validate JWT and require the password-change flow to be completed."""
return await _get_current_subject(
credentials,
allow_password_change = False,
)
async def get_current_subject_sse(
token: Optional[str] = None,
authorization: Optional[str] = None,
) -> str:
"""Auth dep for SSE endpoints.
EventSource cannot send custom headers, so callers pass the bearer
as a ``?token=…`` query param. Falls back to the Authorization
header so curl / API clients keep working.
Wire with ``Query(None)`` and ``Header(None)`` at the route layer:
async def stream(
current_subject: str = Depends(
lambda token = Query(None), authorization = Header(None):
get_current_subject_sse(token, authorization)
),
): ...
"""
raw = token
if not raw and authorization and authorization.lower().startswith("bearer "):
raw = authorization[len("bearer "):].strip()
if not raw:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Missing token",
)
credentials = HTTPAuthorizationCredentials(scheme = "Bearer", credentials = raw)
return await _get_current_subject(
credentials,
allow_password_change = False,
)
async def get_current_subject_allow_password_change(
credentials: HTTPAuthorizationCredentials = Depends(security),
) -> str:
"""Validate JWT but allow access to the password-change endpoint."""
return await _get_current_subject(
credentials,
allow_password_change = True,
)
async def _get_current_subject(
credentials: HTTPAuthorizationCredentials,
*,
allow_password_change: bool,
) -> str:
"""
FastAPI dependency to validate the JWT and return the subject.
Use this as a dependency on routes that should be protected, e.g.:
@router.get("/secure")
async def secure_endpoint(current_subject: str = Depends(get_current_subject)):
...
"""
token = credentials.credentials
# --- API key path (sk-unsloth-...) ---
if token.startswith(API_KEY_PREFIX):
username = validate_api_key(token)
if username is None:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Invalid or expired API key",
)
return username
# --- JWT path ---
subject = _decode_subject_without_verification(token)
if subject is None:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Invalid token payload",
)
record = get_user_and_secret(subject)
if record is None:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Invalid or expired token",
)
_salt, _pwd_hash, jwt_secret, must_change_password = record
try:
payload = jwt.decode(token, jwt_secret, algorithms = [ALGORITHM])
if payload.get("sub") != subject:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Invalid token payload",
)
is_desktop = payload.get("desktop") is True
if must_change_password and not allow_password_change and not is_desktop:
raise HTTPException(
status_code = status.HTTP_403_FORBIDDEN,
detail = "Password change required",
)
return subject
except jwt.InvalidTokenError:
raise HTTPException(
status_code = status.HTTP_401_UNAUTHORIZED,
detail = "Invalid or expired token",
)