mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-08-09 10:39:11 +02:00
refactor(routes): move auth domain into routes/auth/ subpackage
Slice 2n of the route-domain reorganization (#4082/#4071, per
specs/architecture-runtime-inventory.md §6.3). Moves auth_routes.py
(836 lines), api_token_routes.py (209 lines), and device_flow.py (193 lines)
into routes/auth/, leaving backward-compat sys.modules shims at the old
paths. Pure file reorganization, no behavior change.
All three shims use sys.modules replacement so:
- sys.modules.pop("routes.auth_routes") + re-import (test_auth_policy.py)
- monkeypatch.delitem(sys.modules, "routes.api_token_routes") + re-import
(test_api_token_routes.py)
- import ... as ar + monkeypatch.setattr(ar, "DEEP_RESEARCH_DIR", ...)
(test_rename_user_owner_sync.py)
- monkeypatch.setattr(device_flow, "require_admin", ...)
(test_device_flow_routes.py)
all reach the canonical modules.
The 3 auth files do NOT import each other — they are independent modules
that share the auth domain grouping. device_flow is also imported by
copilot_routes and chatgpt_subscription_routes (external); the shim keeps
those working unchanged.
Zero source-introspection landmines. SESSION_COOKIE constant imported by
app.py resolves automatically through the sys.modules shim.
Adds tests/test_auth_routes_shim.py to pin all 3 sys.modules shim contracts.
Verified: compileall clean; full suite 4806 passed, 3 skipped.
This commit is contained in:
parent
20e7fc0164
commit
2d933c8dff
9 changed files with 1324 additions and 1228 deletions
4
app.py
4
app.py
|
|
@ -244,7 +244,7 @@ app.add_middleware(_InteractiveActivityMiddleware)
|
|||
app.add_middleware(_SlowRequestLogMiddleware)
|
||||
|
||||
# ========= AUTH =========
|
||||
from routes.auth_routes import setup_auth_routes, SESSION_COOKIE
|
||||
from routes.auth.auth_routes import setup_auth_routes, SESSION_COOKIE
|
||||
|
||||
auth_manager = AuthManager()
|
||||
app.state.auth_manager = auth_manager
|
||||
|
|
@ -824,7 +824,7 @@ from routes.webhook.webhook_routes import setup_webhook_routes
|
|||
app.include_router(setup_webhook_routes(webhook_manager, auth_manager, session_manager, api_key_manager))
|
||||
|
||||
# API Tokens
|
||||
from routes.api_token_routes import setup_api_token_routes
|
||||
from routes.auth.api_token_routes import setup_api_token_routes
|
||||
app.include_router(setup_api_token_routes())
|
||||
|
||||
logger.info("Webhook & API token routes initialized")
|
||||
|
|
|
|||
|
|
@ -1,209 +1,15 @@
|
|||
"""API Token management routes — /api/tokens/*."""
|
||||
"""Backward-compat shim — canonical location is routes/auth/api_token_routes.py.
|
||||
|
||||
import secrets
|
||||
import uuid
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.api_token_routes``, ``from routes.api_token_routes
|
||||
import X``, and the ``monkeypatch.delitem(sys.modules,
|
||||
"routes.api_token_routes")`` + re-import pattern in test_api_token_routes.py
|
||||
all operate on the *same* object. Keeps existing import paths working after
|
||||
slice 2n (#4082/#4071).
|
||||
"""
|
||||
|
||||
import bcrypt
|
||||
from fastapi import APIRouter, HTTPException, Request, Form
|
||||
import sys as _sys
|
||||
|
||||
from core.database import get_db_session, ApiToken
|
||||
from core.middleware import require_admin
|
||||
from src.auth_helpers import get_current_user
|
||||
from routes.auth import api_token_routes as _canonical # noqa: F401
|
||||
|
||||
MAX_NAME_LEN = 100
|
||||
DEFAULT_SCOPES = "chat"
|
||||
ALLOWED_SCOPES = {
|
||||
"chat",
|
||||
"todos:read",
|
||||
"todos:write",
|
||||
"documents:read",
|
||||
"documents:write",
|
||||
"email:read",
|
||||
"email:draft",
|
||||
"email:send",
|
||||
"calendar:read",
|
||||
"calendar:write",
|
||||
"memory:read",
|
||||
"memory:write",
|
||||
"cookbook:read",
|
||||
"cookbook:launch",
|
||||
}
|
||||
TOKEN_PROFILES = {
|
||||
"chat": ["chat"],
|
||||
"codex_todos": ["todos:read", "todos:write"],
|
||||
"codex_documents": ["documents:read", "documents:write"],
|
||||
"codex_email_drafts": ["email:read", "email:draft", "documents:read", "documents:write"],
|
||||
}
|
||||
|
||||
|
||||
def _normalize_scopes(scopes: str | list[str] | None = None, profile: str | None = None) -> list[str]:
|
||||
profile = profile if isinstance(profile, str) else None
|
||||
profile_key = (profile or "").strip()
|
||||
if profile_key:
|
||||
if profile_key not in TOKEN_PROFILES:
|
||||
raise HTTPException(400, "Unknown token profile")
|
||||
requested = list(TOKEN_PROFILES[profile_key])
|
||||
elif isinstance(scopes, list):
|
||||
requested = [str(s).strip() for s in scopes if str(s).strip()]
|
||||
elif isinstance(scopes, str) and scopes:
|
||||
requested = [s.strip() for s in scopes.replace(" ", ",").split(",") if s.strip()]
|
||||
else:
|
||||
requested = [DEFAULT_SCOPES]
|
||||
|
||||
normalized = []
|
||||
for scope in requested:
|
||||
if scope not in ALLOWED_SCOPES:
|
||||
raise HTTPException(400, f"Unknown token scope: {scope}")
|
||||
if scope not in normalized:
|
||||
normalized.append(scope)
|
||||
|
||||
def ensure_before(write_scope: str, read_scope: str):
|
||||
if write_scope not in normalized or read_scope in normalized:
|
||||
return
|
||||
idx = normalized.index(write_scope)
|
||||
normalized.insert(idx, read_scope)
|
||||
|
||||
ensure_before("todos:write", "todos:read")
|
||||
ensure_before("documents:write", "documents:read")
|
||||
ensure_before("calendar:write", "calendar:read")
|
||||
ensure_before("memory:write", "memory:read")
|
||||
ensure_before("email:draft", "email:read")
|
||||
ensure_before("cookbook:launch", "cookbook:read")
|
||||
|
||||
return normalized or [DEFAULT_SCOPES]
|
||||
|
||||
|
||||
def setup_api_token_routes() -> APIRouter:
|
||||
router = APIRouter(prefix="/api", tags=["api_tokens"])
|
||||
|
||||
@router.get("/tokens")
|
||||
def list_tokens(request: Request):
|
||||
require_admin(request)
|
||||
with get_db_session() as db:
|
||||
tokens = db.query(ApiToken).all()
|
||||
return [
|
||||
{
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"owner": getattr(t, "owner", None),
|
||||
"token_prefix": t.token_prefix,
|
||||
"scopes": [s.strip() for s in (getattr(t, "scopes", "") or DEFAULT_SCOPES).split(",") if s.strip()],
|
||||
"is_active": t.is_active,
|
||||
"last_used_at": t.last_used_at.isoformat() if t.last_used_at else None,
|
||||
"created_at": t.created_at.isoformat() if t.created_at else None,
|
||||
}
|
||||
for t in tokens
|
||||
]
|
||||
|
||||
def _invalidate_cache(request: Request):
|
||||
"""Tell the auth middleware its cached token map is stale."""
|
||||
try:
|
||||
invalidator = getattr(request.app.state, "invalidate_token_cache", None)
|
||||
if invalidator:
|
||||
invalidator()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@router.get("/tokens/profiles")
|
||||
def token_profiles(request: Request):
|
||||
require_admin(request)
|
||||
return {
|
||||
"profiles": TOKEN_PROFILES,
|
||||
"allowed_scopes": sorted(ALLOWED_SCOPES),
|
||||
}
|
||||
|
||||
@router.post("/tokens")
|
||||
def create_token(
|
||||
request: Request,
|
||||
name: str = Form(""),
|
||||
scopes: str = Form(None),
|
||||
profile: str = Form(None),
|
||||
):
|
||||
require_admin(request)
|
||||
name = name.strip()[:MAX_NAME_LEN]
|
||||
if not name:
|
||||
raise HTTPException(400, "Token name is required")
|
||||
owner = get_current_user(request)
|
||||
scope_list = _normalize_scopes(scopes, profile)
|
||||
scopes_value = ",".join(scope_list)
|
||||
|
||||
raw_token = "ody_" + secrets.token_urlsafe(32)
|
||||
token_hash = bcrypt.hashpw(raw_token.encode(), bcrypt.gensalt()).decode()
|
||||
token_id = str(uuid.uuid4())[:8]
|
||||
|
||||
with get_db_session() as db:
|
||||
db.add(ApiToken(
|
||||
id=token_id,
|
||||
owner=owner,
|
||||
name=name,
|
||||
token_hash=token_hash,
|
||||
token_prefix=raw_token[:8],
|
||||
scopes=scopes_value,
|
||||
is_active=True,
|
||||
))
|
||||
_invalidate_cache(request)
|
||||
|
||||
return {
|
||||
"id": token_id,
|
||||
"name": name,
|
||||
"owner": owner,
|
||||
"token": raw_token,
|
||||
"token_prefix": raw_token[:8],
|
||||
"scopes": scope_list,
|
||||
}
|
||||
|
||||
@router.patch("/tokens/{token_id}")
|
||||
async def update_token(request: Request, token_id: str):
|
||||
require_admin(request)
|
||||
current_user = get_current_user(request)
|
||||
try:
|
||||
payload = await request.json()
|
||||
except Exception:
|
||||
payload = {}
|
||||
if not isinstance(payload, dict):
|
||||
payload = {}
|
||||
with get_db_session() as db:
|
||||
token = db.query(ApiToken).filter(ApiToken.id == token_id).first()
|
||||
if not token:
|
||||
raise HTTPException(404, "Token not found")
|
||||
if current_user and token.owner != current_user:
|
||||
raise HTTPException(403, "Not your token")
|
||||
if isinstance(payload.get("name"), str) and payload["name"].strip():
|
||||
token.name = payload["name"].strip()[:MAX_NAME_LEN]
|
||||
# Only touch scopes when the caller actually sent them. A partial
|
||||
# update such as a rename ({"name": ...} with no "scopes" key) must
|
||||
# not silently reset the token to the default scope — that dropped
|
||||
# every previously granted scope.
|
||||
if "scopes" in payload:
|
||||
token.scopes = ",".join(_normalize_scopes(payload.get("scopes")))
|
||||
db.add(token)
|
||||
current_scopes = [
|
||||
s.strip()
|
||||
for s in (getattr(token, "scopes", "") or DEFAULT_SCOPES).split(",")
|
||||
if s.strip()
|
||||
]
|
||||
response = {
|
||||
"id": token_id,
|
||||
"name": getattr(token, "name", ""),
|
||||
"owner": getattr(token, "owner", None),
|
||||
"token_prefix": getattr(token, "token_prefix", ""),
|
||||
"scopes": current_scopes,
|
||||
}
|
||||
_invalidate_cache(request)
|
||||
return response
|
||||
|
||||
@router.delete("/tokens/{token_id}")
|
||||
def delete_token(request: Request, token_id: str):
|
||||
require_admin(request)
|
||||
current_user = get_current_user(request)
|
||||
with get_db_session() as db:
|
||||
token = db.query(ApiToken).filter(ApiToken.id == token_id).first()
|
||||
if not token:
|
||||
raise HTTPException(404, "Token not found")
|
||||
if current_user and token.owner != current_user:
|
||||
raise HTTPException(403, "Not your token")
|
||||
db.delete(token)
|
||||
_invalidate_cache(request)
|
||||
return {"status": "deleted"}
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
|
|
|||
6
routes/auth/__init__.py
Normal file
6
routes/auth/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""Auth route domain package (slice 2n, #4082/#4071).
|
||||
|
||||
Contains auth_routes.py, api_token_routes.py, and device_flow.py, migrated
|
||||
from the flat routes/ directory. Backward-compat shims at the old paths
|
||||
re-export from here.
|
||||
"""
|
||||
209
routes/auth/api_token_routes.py
Normal file
209
routes/auth/api_token_routes.py
Normal file
|
|
@ -0,0 +1,209 @@
|
|||
"""API Token management routes — /api/tokens/*."""
|
||||
|
||||
import secrets
|
||||
import uuid
|
||||
|
||||
import bcrypt
|
||||
from fastapi import APIRouter, HTTPException, Request, Form
|
||||
|
||||
from core.database import get_db_session, ApiToken
|
||||
from core.middleware import require_admin
|
||||
from src.auth_helpers import get_current_user
|
||||
|
||||
MAX_NAME_LEN = 100
|
||||
DEFAULT_SCOPES = "chat"
|
||||
ALLOWED_SCOPES = {
|
||||
"chat",
|
||||
"todos:read",
|
||||
"todos:write",
|
||||
"documents:read",
|
||||
"documents:write",
|
||||
"email:read",
|
||||
"email:draft",
|
||||
"email:send",
|
||||
"calendar:read",
|
||||
"calendar:write",
|
||||
"memory:read",
|
||||
"memory:write",
|
||||
"cookbook:read",
|
||||
"cookbook:launch",
|
||||
}
|
||||
TOKEN_PROFILES = {
|
||||
"chat": ["chat"],
|
||||
"codex_todos": ["todos:read", "todos:write"],
|
||||
"codex_documents": ["documents:read", "documents:write"],
|
||||
"codex_email_drafts": ["email:read", "email:draft", "documents:read", "documents:write"],
|
||||
}
|
||||
|
||||
|
||||
def _normalize_scopes(scopes: str | list[str] | None = None, profile: str | None = None) -> list[str]:
|
||||
profile = profile if isinstance(profile, str) else None
|
||||
profile_key = (profile or "").strip()
|
||||
if profile_key:
|
||||
if profile_key not in TOKEN_PROFILES:
|
||||
raise HTTPException(400, "Unknown token profile")
|
||||
requested = list(TOKEN_PROFILES[profile_key])
|
||||
elif isinstance(scopes, list):
|
||||
requested = [str(s).strip() for s in scopes if str(s).strip()]
|
||||
elif isinstance(scopes, str) and scopes:
|
||||
requested = [s.strip() for s in scopes.replace(" ", ",").split(",") if s.strip()]
|
||||
else:
|
||||
requested = [DEFAULT_SCOPES]
|
||||
|
||||
normalized = []
|
||||
for scope in requested:
|
||||
if scope not in ALLOWED_SCOPES:
|
||||
raise HTTPException(400, f"Unknown token scope: {scope}")
|
||||
if scope not in normalized:
|
||||
normalized.append(scope)
|
||||
|
||||
def ensure_before(write_scope: str, read_scope: str):
|
||||
if write_scope not in normalized or read_scope in normalized:
|
||||
return
|
||||
idx = normalized.index(write_scope)
|
||||
normalized.insert(idx, read_scope)
|
||||
|
||||
ensure_before("todos:write", "todos:read")
|
||||
ensure_before("documents:write", "documents:read")
|
||||
ensure_before("calendar:write", "calendar:read")
|
||||
ensure_before("memory:write", "memory:read")
|
||||
ensure_before("email:draft", "email:read")
|
||||
ensure_before("cookbook:launch", "cookbook:read")
|
||||
|
||||
return normalized or [DEFAULT_SCOPES]
|
||||
|
||||
|
||||
def setup_api_token_routes() -> APIRouter:
|
||||
router = APIRouter(prefix="/api", tags=["api_tokens"])
|
||||
|
||||
@router.get("/tokens")
|
||||
def list_tokens(request: Request):
|
||||
require_admin(request)
|
||||
with get_db_session() as db:
|
||||
tokens = db.query(ApiToken).all()
|
||||
return [
|
||||
{
|
||||
"id": t.id,
|
||||
"name": t.name,
|
||||
"owner": getattr(t, "owner", None),
|
||||
"token_prefix": t.token_prefix,
|
||||
"scopes": [s.strip() for s in (getattr(t, "scopes", "") or DEFAULT_SCOPES).split(",") if s.strip()],
|
||||
"is_active": t.is_active,
|
||||
"last_used_at": t.last_used_at.isoformat() if t.last_used_at else None,
|
||||
"created_at": t.created_at.isoformat() if t.created_at else None,
|
||||
}
|
||||
for t in tokens
|
||||
]
|
||||
|
||||
def _invalidate_cache(request: Request):
|
||||
"""Tell the auth middleware its cached token map is stale."""
|
||||
try:
|
||||
invalidator = getattr(request.app.state, "invalidate_token_cache", None)
|
||||
if invalidator:
|
||||
invalidator()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@router.get("/tokens/profiles")
|
||||
def token_profiles(request: Request):
|
||||
require_admin(request)
|
||||
return {
|
||||
"profiles": TOKEN_PROFILES,
|
||||
"allowed_scopes": sorted(ALLOWED_SCOPES),
|
||||
}
|
||||
|
||||
@router.post("/tokens")
|
||||
def create_token(
|
||||
request: Request,
|
||||
name: str = Form(""),
|
||||
scopes: str = Form(None),
|
||||
profile: str = Form(None),
|
||||
):
|
||||
require_admin(request)
|
||||
name = name.strip()[:MAX_NAME_LEN]
|
||||
if not name:
|
||||
raise HTTPException(400, "Token name is required")
|
||||
owner = get_current_user(request)
|
||||
scope_list = _normalize_scopes(scopes, profile)
|
||||
scopes_value = ",".join(scope_list)
|
||||
|
||||
raw_token = "ody_" + secrets.token_urlsafe(32)
|
||||
token_hash = bcrypt.hashpw(raw_token.encode(), bcrypt.gensalt()).decode()
|
||||
token_id = str(uuid.uuid4())[:8]
|
||||
|
||||
with get_db_session() as db:
|
||||
db.add(ApiToken(
|
||||
id=token_id,
|
||||
owner=owner,
|
||||
name=name,
|
||||
token_hash=token_hash,
|
||||
token_prefix=raw_token[:8],
|
||||
scopes=scopes_value,
|
||||
is_active=True,
|
||||
))
|
||||
_invalidate_cache(request)
|
||||
|
||||
return {
|
||||
"id": token_id,
|
||||
"name": name,
|
||||
"owner": owner,
|
||||
"token": raw_token,
|
||||
"token_prefix": raw_token[:8],
|
||||
"scopes": scope_list,
|
||||
}
|
||||
|
||||
@router.patch("/tokens/{token_id}")
|
||||
async def update_token(request: Request, token_id: str):
|
||||
require_admin(request)
|
||||
current_user = get_current_user(request)
|
||||
try:
|
||||
payload = await request.json()
|
||||
except Exception:
|
||||
payload = {}
|
||||
if not isinstance(payload, dict):
|
||||
payload = {}
|
||||
with get_db_session() as db:
|
||||
token = db.query(ApiToken).filter(ApiToken.id == token_id).first()
|
||||
if not token:
|
||||
raise HTTPException(404, "Token not found")
|
||||
if current_user and token.owner != current_user:
|
||||
raise HTTPException(403, "Not your token")
|
||||
if isinstance(payload.get("name"), str) and payload["name"].strip():
|
||||
token.name = payload["name"].strip()[:MAX_NAME_LEN]
|
||||
# Only touch scopes when the caller actually sent them. A partial
|
||||
# update such as a rename ({"name": ...} with no "scopes" key) must
|
||||
# not silently reset the token to the default scope — that dropped
|
||||
# every previously granted scope.
|
||||
if "scopes" in payload:
|
||||
token.scopes = ",".join(_normalize_scopes(payload.get("scopes")))
|
||||
db.add(token)
|
||||
current_scopes = [
|
||||
s.strip()
|
||||
for s in (getattr(token, "scopes", "") or DEFAULT_SCOPES).split(",")
|
||||
if s.strip()
|
||||
]
|
||||
response = {
|
||||
"id": token_id,
|
||||
"name": getattr(token, "name", ""),
|
||||
"owner": getattr(token, "owner", None),
|
||||
"token_prefix": getattr(token, "token_prefix", ""),
|
||||
"scopes": current_scopes,
|
||||
}
|
||||
_invalidate_cache(request)
|
||||
return response
|
||||
|
||||
@router.delete("/tokens/{token_id}")
|
||||
def delete_token(request: Request, token_id: str):
|
||||
require_admin(request)
|
||||
current_user = get_current_user(request)
|
||||
with get_db_session() as db:
|
||||
token = db.query(ApiToken).filter(ApiToken.id == token_id).first()
|
||||
if not token:
|
||||
raise HTTPException(404, "Token not found")
|
||||
if current_user and token.owner != current_user:
|
||||
raise HTTPException(403, "Not your token")
|
||||
db.delete(token)
|
||||
_invalidate_cache(request)
|
||||
return {"status": "deleted"}
|
||||
|
||||
return router
|
||||
836
routes/auth/auth_routes.py
Normal file
836
routes/auth/auth_routes.py
Normal file
|
|
@ -0,0 +1,836 @@
|
|||
"""Authentication routes — login, logout, signup, status, user management."""
|
||||
|
||||
from fastapi import APIRouter, Request, Response, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from typing import Optional
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
|
||||
from core.atomic_io import atomic_write_json, atomic_write_text
|
||||
from core.auth import AuthManager, RESERVED_USERNAMES, SetAdminResult, TOKEN_TTL
|
||||
from src.constants import DEEP_RESEARCH_DIR, MEMORY_FILE, PASSWORD_MIN_LENGTH, SKILLS_DIR
|
||||
from src.rate_limiter import RateLimiter
|
||||
from src.settings_scrub import scrub_settings
|
||||
from src.settings import (
|
||||
load_settings as _load_settings,
|
||||
save_settings as _save_settings,
|
||||
load_features as _load_features,
|
||||
save_features as _save_features,
|
||||
DEFAULT_SETTINGS,
|
||||
)
|
||||
from src.integrations import (
|
||||
load_integrations,
|
||||
add_integration,
|
||||
update_integration,
|
||||
delete_integration,
|
||||
get_integration,
|
||||
mask_integration_secret,
|
||||
execute_api_call,
|
||||
INTEGRATION_PRESETS,
|
||||
migrate_from_settings,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
remember: bool = True
|
||||
totp_code: Optional[str] = None
|
||||
|
||||
|
||||
class SetupRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
class SignupRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
class ChangePasswordRequest(BaseModel):
|
||||
current_password: str
|
||||
new_password: str
|
||||
|
||||
|
||||
class CreateUserRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
is_admin: bool = False
|
||||
|
||||
|
||||
class DeleteUserRequest(BaseModel):
|
||||
username: str
|
||||
|
||||
|
||||
class RenameUserRequest(BaseModel):
|
||||
username: str
|
||||
|
||||
|
||||
class SetAdminRequest(BaseModel):
|
||||
is_admin: bool
|
||||
|
||||
|
||||
class SetOpenRegistrationRequest(BaseModel):
|
||||
enabled: bool
|
||||
|
||||
SESSION_COOKIE = "odysseus_session"
|
||||
|
||||
|
||||
def setup_auth_routes(auth_manager: AuthManager) -> APIRouter:
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
_login_limiter = RateLimiter(max_requests=15, window_seconds=60)
|
||||
_signup_limiter = RateLimiter(max_requests=3, window_seconds=300)
|
||||
_setup_limiter = RateLimiter(max_requests=3, window_seconds=300)
|
||||
|
||||
def _get_current_user(request: Request) -> Optional[str]:
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
return auth_manager.get_username_for_token(token)
|
||||
|
||||
@router.post("/setup")
|
||||
async def first_run_setup(body: SetupRequest, request: Request):
|
||||
"""Create initial admin account. Only works if no accounts exist."""
|
||||
if not _setup_limiter.check(request.client.host):
|
||||
raise HTTPException(429, "Too many requests — try again later")
|
||||
if auth_manager.is_configured:
|
||||
raise HTTPException(400, "Already configured")
|
||||
if len(body.password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
if len(body.username.strip()) < 1:
|
||||
raise HTTPException(400, "Username is required")
|
||||
if body.username.lower() in RESERVED_USERNAMES:
|
||||
raise HTTPException(403, "Username is reserved")
|
||||
ok = await asyncio.to_thread(auth_manager.setup, body.username, body.password)
|
||||
if not ok:
|
||||
raise HTTPException(500, "Setup failed")
|
||||
return {"ok": True, "message": "Admin account created"}
|
||||
|
||||
@router.post("/signup")
|
||||
async def signup(body: SignupRequest, request: Request):
|
||||
"""Create a new user account. Only works if signup is enabled by admin."""
|
||||
if not _signup_limiter.check(request.client.host):
|
||||
raise HTTPException(429, "Too many requests — try again later")
|
||||
if not auth_manager.is_configured:
|
||||
raise HTTPException(400, "Run setup first")
|
||||
if not auth_manager.signup_enabled:
|
||||
raise HTTPException(403, "Registration is disabled. Ask an admin for an account.")
|
||||
if len(body.password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
if len(body.username.strip()) < 1:
|
||||
raise HTTPException(400, "Username is required")
|
||||
if body.username.lower() in RESERVED_USERNAMES:
|
||||
raise HTTPException(403, "Username is reserved")
|
||||
ok = await asyncio.to_thread(auth_manager.create_user, body.username, body.password, is_admin=False)
|
||||
if not ok:
|
||||
raise HTTPException(409, "Username already taken")
|
||||
return {"ok": True, "message": "Account created"}
|
||||
|
||||
@router.post("/login")
|
||||
async def login(body: LoginRequest, request: Request, response: Response):
|
||||
if not _login_limiter.check(request.client.host):
|
||||
raise HTTPException(429, "Too many requests — try again later")
|
||||
# Verify password first
|
||||
username = body.username.strip().lower()
|
||||
if not await asyncio.to_thread(auth_manager.verify_password, username, body.password):
|
||||
raise HTTPException(401, "Invalid credentials")
|
||||
# Check 2FA if enabled
|
||||
if auth_manager.totp_enabled(username):
|
||||
if not body.totp_code:
|
||||
# Password OK but need TOTP — tell client to show code input
|
||||
return {"ok": False, "requires_totp": True, "username": username}
|
||||
if not auth_manager.totp_verify(username, body.totp_code):
|
||||
raise HTTPException(401, "Invalid 2FA code")
|
||||
# All checks passed — create session (password already verified above)
|
||||
token = await asyncio.to_thread(auth_manager.create_session_trusted, username)
|
||||
if not token:
|
||||
raise HTTPException(401, "Invalid credentials")
|
||||
cookie_kwargs = dict(
|
||||
key=SESSION_COOKIE,
|
||||
value=token,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
secure=os.getenv("SECURE_COOKIES", "false").lower() == "true",
|
||||
path="/",
|
||||
)
|
||||
if body.remember:
|
||||
cookie_kwargs["max_age"] = TOKEN_TTL
|
||||
response.set_cookie(**cookie_kwargs)
|
||||
return {"ok": True, "username": username}
|
||||
|
||||
@router.post("/logout")
|
||||
async def logout(request: Request, response: Response):
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
if token:
|
||||
auth_manager.revoke_token(token)
|
||||
response.delete_cookie(SESSION_COOKIE, path="/")
|
||||
return {"ok": True}
|
||||
|
||||
@router.get("/status")
|
||||
async def auth_status(request: Request):
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
result = auth_manager.status(token)
|
||||
result["signup_enabled"] = auth_manager.signup_enabled
|
||||
# Include the caller's effective privileges so the frontend can
|
||||
# hide / dim UI controls the user isn't allowed to use. Admins get
|
||||
# ADMIN_PRIVILEGES (everything on), regular users get their stored
|
||||
# set merged with DEFAULT_PRIVILEGES.
|
||||
try:
|
||||
u = result.get("username")
|
||||
if u:
|
||||
result["privileges"] = auth_manager.get_privileges(u)
|
||||
except Exception:
|
||||
pass
|
||||
return result
|
||||
|
||||
@router.get("/policy")
|
||||
async def auth_policy():
|
||||
"""Return public auth policy constants for the frontend."""
|
||||
return auth_manager.policy()
|
||||
|
||||
@router.post("/change-password")
|
||||
async def change_password(body: ChangePasswordRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if len(body.new_password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
current_token = request.cookies.get(SESSION_COOKIE)
|
||||
ok = await asyncio.to_thread(auth_manager.change_password, user, body.current_password, body.new_password)
|
||||
if not ok:
|
||||
raise HTTPException(400, "Current password is incorrect")
|
||||
await asyncio.to_thread(auth_manager.revoke_user_sessions, user, current_token)
|
||||
return {"ok": True}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Two-factor authentication
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@router.post("/2fa/setup")
|
||||
async def totp_setup(request: Request):
|
||||
"""Generate a TOTP secret and return the QR code URI."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if auth_manager.totp_enabled(user):
|
||||
raise HTTPException(400, "2FA is already enabled")
|
||||
secret = auth_manager.totp_generate_secret(user)
|
||||
if not secret:
|
||||
raise HTTPException(500, "Failed to generate secret")
|
||||
uri = auth_manager.totp_get_provisioning_uri(user, secret)
|
||||
# Generate QR code as base64 PNG
|
||||
import qrcode, io, base64
|
||||
qr = qrcode.make(uri, box_size=6, border=2)
|
||||
buf = io.BytesIO()
|
||||
qr.save(buf, format="PNG")
|
||||
qr_b64 = base64.b64encode(buf.getvalue()).decode("ascii")
|
||||
return {"secret": secret, "uri": uri, "qr_code": f"data:image/png;base64,{qr_b64}"}
|
||||
|
||||
class TotpVerifyRequest(BaseModel):
|
||||
code: str
|
||||
|
||||
@router.post("/2fa/confirm")
|
||||
async def totp_confirm(body: TotpVerifyRequest, request: Request):
|
||||
"""Verify a TOTP code to confirm 2FA setup. Returns backup codes."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if not auth_manager.totp_confirm_enable(user, body.code):
|
||||
raise HTTPException(400, "Invalid code — try again")
|
||||
backup = auth_manager.users.get(user, {}).get("totp_backup_codes", [])
|
||||
return {"ok": True, "backup_codes": backup}
|
||||
|
||||
class TotpDisableRequest(BaseModel):
|
||||
password: str
|
||||
|
||||
@router.post("/2fa/disable")
|
||||
async def totp_disable(body: TotpDisableRequest, request: Request):
|
||||
"""Disable 2FA. Requires password confirmation."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if not auth_manager.totp_disable(user, body.password):
|
||||
raise HTTPException(400, "Invalid password")
|
||||
return {"ok": True}
|
||||
|
||||
@router.get("/2fa/status")
|
||||
async def totp_status(request: Request):
|
||||
"""Check if 2FA is enabled for the current user."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
return {"enabled": auth_manager.totp_enabled(user)}
|
||||
|
||||
# Admin-only routes
|
||||
@router.get("/users")
|
||||
async def list_users(request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
return {"users": auth_manager.list_users()}
|
||||
|
||||
@router.post("/users")
|
||||
async def admin_create_user(body: CreateUserRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
if len(body.password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
if len(body.username.strip()) < 1:
|
||||
raise HTTPException(400, "Username is required")
|
||||
if body.username.lower() in RESERVED_USERNAMES:
|
||||
raise HTTPException(403, "Username is reserved")
|
||||
ok = auth_manager.create_user(body.username, body.password, body.is_admin)
|
||||
if not ok:
|
||||
raise HTTPException(409, "Username already taken")
|
||||
return {"ok": True}
|
||||
|
||||
@router.put("/users/{username}/privileges")
|
||||
async def update_user_privileges(username: str, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
ok = auth_manager.set_privileges(username, body)
|
||||
if not ok:
|
||||
raise HTTPException(404, "User not found or is admin")
|
||||
return {"ok": True, "privileges": auth_manager.get_privileges(username)}
|
||||
|
||||
@router.put("/users/{username}/rename")
|
||||
async def rename_user(username: str, body: RenameUserRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
old_username = (username or "").strip().lower()
|
||||
new_username = (body.username or "").strip().lower()
|
||||
if not new_username:
|
||||
raise HTTPException(400, "Username required")
|
||||
if old_username == new_username:
|
||||
return {"ok": True, "username": new_username, "renamed_self": old_username == user}
|
||||
if old_username not in auth_manager.users:
|
||||
raise HTTPException(404, "User not found")
|
||||
if new_username in auth_manager.users:
|
||||
raise HTTPException(409, "Username already taken")
|
||||
|
||||
# Gate on auth first. Every mutation below is contingent on this
|
||||
# succeeding — doing it last meant a rejected rename (e.g. reserved
|
||||
# username) left file-backed owner fields already rewritten with no
|
||||
# way to roll them back.
|
||||
ok = auth_manager.rename_user(old_username, new_username, user)
|
||||
if not ok:
|
||||
raise HTTPException(400, "Cannot rename user")
|
||||
|
||||
def _rollback_auth_rename() -> bool:
|
||||
# On self-rename the admin session has already moved to the new
|
||||
# username, so the rollback must authenticate as the new user.
|
||||
rollback_user = new_username if user == old_username else user
|
||||
try:
|
||||
return bool(auth_manager.rename_user(new_username, old_username, rollback_user))
|
||||
except Exception as rollback_err:
|
||||
logger.error(
|
||||
"Failed to roll back auth rename %s -> %s after owner migration failure: %s",
|
||||
new_username, old_username, rollback_err,
|
||||
)
|
||||
return False
|
||||
|
||||
# Usernames are ownership keys for user data. Rename the common
|
||||
# owner-scoped DB rows so the account keeps access to its sessions,
|
||||
# docs, email accounts, tasks, etc.
|
||||
try:
|
||||
from sqlalchemy import func
|
||||
from core.database import Base, SessionLocal
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for mapper in Base.registry.mappers:
|
||||
model = mapper.class_
|
||||
if not hasattr(model, "owner"):
|
||||
continue
|
||||
(
|
||||
db.query(model)
|
||||
.filter(func.lower(model.owner) == old_username)
|
||||
.update({"owner": new_username}, synchronize_session=False)
|
||||
)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as e:
|
||||
logger.error("Failed to rename owner references %s -> %s: %s", old_username, new_username, e)
|
||||
if not _rollback_auth_rename():
|
||||
logger.error(
|
||||
"Auth rename %s -> %s could not be rolled back after owner migration failure",
|
||||
old_username, new_username,
|
||||
)
|
||||
raise HTTPException(500, "Failed to rename user data")
|
||||
|
||||
# Per-user prefs are JSON-backed, not SQL-backed.
|
||||
try:
|
||||
from routes.prefs_routes import _load as _load_prefs, _save as _save_prefs
|
||||
prefs = _load_prefs()
|
||||
users = prefs.get("_users") if isinstance(prefs, dict) else None
|
||||
if isinstance(users, dict):
|
||||
prefs_key = next(
|
||||
(k for k in users if str(k).strip().lower() == old_username),
|
||||
None,
|
||||
)
|
||||
new_taken = any(str(k).strip().lower() == new_username for k in users)
|
||||
if prefs_key is not None and not new_taken:
|
||||
users[new_username] = users.pop(prefs_key)
|
||||
_save_prefs(prefs)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename user prefs %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# In-flight deep-research tasks live in the process-local
|
||||
# ResearchHandler registry. They are not covered by the persisted JSON
|
||||
# migration above, but the research routes filter and cancel by this
|
||||
# owner field while the job is running. Do this before sweeping
|
||||
# completed JSON files so a job that finishes during the rename saves
|
||||
# with the new owner or is caught by the disk sweep below.
|
||||
try:
|
||||
rh = getattr(request.app.state, "research_handler", None)
|
||||
rename_owner = getattr(rh, "rename_owner", None)
|
||||
if callable(rename_owner):
|
||||
rename_owner(old_username, new_username)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename active research tasks %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# deep_research: each completed report is a standalone JSON file with
|
||||
# an `owner` field. research_routes filters by d.get("owner") == user,
|
||||
# so a stale owner makes every report invisible to the renamed user.
|
||||
try:
|
||||
dr_dir = Path(DEEP_RESEARCH_DIR)
|
||||
if dr_dir.is_dir():
|
||||
for p in dr_dir.glob("*.json"):
|
||||
try:
|
||||
d = json.loads(p.read_text(encoding="utf-8"))
|
||||
if str(d.get("owner", "")).strip().lower() == old_username:
|
||||
d["owner"] = new_username
|
||||
atomic_write_json(str(p), d)
|
||||
except Exception as err:
|
||||
logger.warning("Failed to update research owner in %s: %s", p.name, err)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename research owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# memory.json: a flat JSON array where each entry carries an `owner`
|
||||
# field. memory_manager.load(owner=user) filters on it, so stale
|
||||
# entries disappear from the memory panel.
|
||||
try:
|
||||
if os.path.isfile(MEMORY_FILE):
|
||||
with open(MEMORY_FILE, encoding="utf-8") as fh:
|
||||
entries = json.loads(fh.read())
|
||||
if isinstance(entries, list):
|
||||
changed = False
|
||||
for entry in entries:
|
||||
if isinstance(entry, dict) and str(entry.get("owner", "")).strip().lower() == old_username:
|
||||
entry["owner"] = new_username
|
||||
changed = True
|
||||
if changed:
|
||||
atomic_write_json(MEMORY_FILE, entries)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename memory.json owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# uploads.json: upload rows use owner metadata for access checks and
|
||||
# owner-prefixed index keys for dedupe. Rename both so attachments keep
|
||||
# resolving after the account username changes.
|
||||
try:
|
||||
upload_handler = getattr(request.app.state, "upload_handler", None)
|
||||
rename_owner = getattr(upload_handler, "rename_owner", None)
|
||||
if callable(rename_owner):
|
||||
rename_owner(old_username, new_username)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename upload owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# direct personal RAG uploads live in per-owner directories and the
|
||||
# vector metadata also carries the username used for owner-filtered
|
||||
# search. Keep both in sync with the auth rename.
|
||||
try:
|
||||
from routes.personal_routes import rename_personal_upload_owner
|
||||
personal_docs_manager = getattr(request.app.state, "personal_docs_manager", None)
|
||||
if personal_docs_manager is not None:
|
||||
rag_manager = getattr(personal_docs_manager, "rag_manager", None)
|
||||
rename_personal_upload_owner(
|
||||
old_username,
|
||||
new_username,
|
||||
personal_docs_manager=personal_docs_manager,
|
||||
rag_manager=rag_manager,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename personal RAG upload owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# skills: SKILL.md frontmatter carries owner: <username>; the usage
|
||||
# sidecar (_usage.json) keys entries as owner::skill-name. Both must
|
||||
# be updated or the renamed user's Skills panel goes empty.
|
||||
try:
|
||||
skills_root = Path(SKILLS_DIR)
|
||||
if skills_root.is_dir():
|
||||
_owner_re = re.compile(
|
||||
r'(?m)^(owner:\s*)' + re.escape(old_username) + r'\s*$',
|
||||
re.IGNORECASE,
|
||||
)
|
||||
for p in skills_root.rglob("SKILL.md"):
|
||||
try:
|
||||
text = p.read_text(encoding="utf-8")
|
||||
new_text = _owner_re.sub(r'\g<1>' + new_username, text)
|
||||
if new_text != text:
|
||||
atomic_write_text(str(p), new_text)
|
||||
except Exception as err:
|
||||
logger.warning("Failed to update skill owner in %s: %s", p, err)
|
||||
usage_path = skills_root / "_usage.json"
|
||||
if usage_path.is_file():
|
||||
try:
|
||||
usage = json.loads(usage_path.read_text(encoding="utf-8"))
|
||||
if isinstance(usage, dict):
|
||||
new_usage = {}
|
||||
changed = False
|
||||
for k, v in usage.items():
|
||||
owner_part, sep, skill_part = k.partition("::")
|
||||
if sep and owner_part.lower() == old_username:
|
||||
new_usage[new_username + "::" + skill_part] = v
|
||||
changed = True
|
||||
else:
|
||||
new_usage[k] = v
|
||||
if changed:
|
||||
atomic_write_json(str(usage_path), new_usage)
|
||||
except Exception as err:
|
||||
logger.warning("Failed to update skills usage keys %s -> %s: %s", old_username, new_username, err)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename skills owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# The in-memory session cache (session_manager.sessions) stores each
|
||||
# session's owner at load time. Without this patch the renamed user's
|
||||
# sessions are invisible on the next /api/sessions call because
|
||||
# get_sessions_for_user does an exact `s.owner == username` comparison
|
||||
# against stale in-memory values.
|
||||
sm = getattr(request.app.state, "session_manager", None)
|
||||
if sm is not None:
|
||||
for sess in list(getattr(sm, "sessions", {}).values()):
|
||||
if str(getattr(sess, "owner", None) or "").strip().lower() == old_username:
|
||||
sess.owner = new_username
|
||||
|
||||
# The owner-rename loop above updated ApiToken.owner in the DB, but the
|
||||
# bearer-token cache still maps each token to the OLD owner. Without
|
||||
# refreshing it, the renamed user's API tokens resolve to the old (now
|
||||
# non-existent) owner and stop reaching their data until the cache next
|
||||
# goes dirty. Invalidate it now, like the token CRUD routes do.
|
||||
invalidator = getattr(request.app.state, "invalidate_token_cache", None)
|
||||
if callable(invalidator):
|
||||
invalidator()
|
||||
return {"ok": True, "username": new_username, "renamed_self": old_username == user}
|
||||
|
||||
@router.put("/users/{username}/admin")
|
||||
async def set_user_admin(username: str, body: SetAdminRequest, request: Request):
|
||||
"""Promote/demote a user to/from admin. Admin only.
|
||||
|
||||
The last remaining admin can't be demoted (no lockout). Self-demotion
|
||||
is allowed while another admin exists; the `self` flag tells the UI to
|
||||
reload the acting user into the normal-user view.
|
||||
"""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
result = auth_manager.set_admin(username, body.is_admin, user)
|
||||
if result is SetAdminResult.USER_NOT_FOUND:
|
||||
raise HTTPException(404, "User not found")
|
||||
if result is SetAdminResult.NOT_AUTHORIZED:
|
||||
raise HTTPException(403, "Admin only")
|
||||
if result is SetAdminResult.LAST_ADMIN:
|
||||
raise HTTPException(400, "Cannot demote the last admin")
|
||||
target = (username or "").strip().lower()
|
||||
return {
|
||||
"ok": True,
|
||||
"is_admin": body.is_admin,
|
||||
"self": target == (user or "").strip().lower(),
|
||||
}
|
||||
|
||||
@router.post("/signup-toggle", deprecated=True)
|
||||
async def toggle_signup(request: Request):
|
||||
"""
|
||||
Toggle open registration on/off. Admin only.
|
||||
|
||||
DEPRECATED: This endpoint uses toggle semantics which can lead to unsafe state changes.
|
||||
Use PUT /open-signup instead.
|
||||
|
||||
This endpoint is kept for backward compatibility and may be removed in future versions.
|
||||
"""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
auth_manager.signup_enabled = not auth_manager.signup_enabled
|
||||
return {"ok": True, "signup_enabled": auth_manager.signup_enabled}
|
||||
|
||||
@router.put("/open-signup")
|
||||
async def set_signup_enabled(body: SetOpenRegistrationRequest, request: Request):
|
||||
"""Set open signup enabled state. Admin only."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
auth_manager.signup_enabled = body.enabled
|
||||
return {"ok": True,"signup_enabled": auth_manager.signup_enabled}
|
||||
|
||||
@router.delete("/users")
|
||||
async def admin_delete_user(body: DeleteUserRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
|
||||
def _invalidate_api_token_cache():
|
||||
try:
|
||||
invalidator = getattr(request.app.state, "invalidate_token_cache", None)
|
||||
if invalidator:
|
||||
invalidator()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
ok = auth_manager.delete_user(body.username, user)
|
||||
except Exception:
|
||||
# delete_user can touch ApiToken rows before a later auth-store write
|
||||
# fails. Dirty the bearer cache anyway so a partial token purge does
|
||||
# not leave already-cached tokens authenticating until restart.
|
||||
_invalidate_api_token_cache()
|
||||
raise
|
||||
if not ok:
|
||||
raise HTTPException(400, "Cannot delete user")
|
||||
# delete_user removes the user's ApiToken rows, but the bearer-auth
|
||||
# middleware serves from an in-memory prefix->token cache that only
|
||||
# rebuilds when flagged dirty. Without this, a deleted user's already
|
||||
# cached token keeps authenticating until some other token op or a
|
||||
# restart clears the cache. Mirror what the token routes do.
|
||||
_invalidate_api_token_cache()
|
||||
return {"ok": True}
|
||||
|
||||
# ---- Feature visibility (admin-managed) ----
|
||||
|
||||
@router.get("/features")
|
||||
async def get_features():
|
||||
"""Public: returns which UI features are enabled."""
|
||||
return _load_features()
|
||||
|
||||
@router.post("/features")
|
||||
async def set_features(request: Request):
|
||||
"""Admin only: update feature toggles."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
current = _load_features()
|
||||
for key in current:
|
||||
if key in body and isinstance(body[key], bool):
|
||||
current[key] = body[key]
|
||||
_save_features(current)
|
||||
return current
|
||||
|
||||
# ---- App settings (admin-managed) ----
|
||||
|
||||
@router.get("/settings")
|
||||
async def get_settings(request: Request):
|
||||
"""Returns app settings. Admins get the full set; non-admins get
|
||||
a scrubbed copy with secret keys blanked. The frontend uses this
|
||||
for keybinds + TTS prefs, so it stays callable without admin."""
|
||||
user = _get_current_user(request)
|
||||
settings = _load_settings()
|
||||
if user and auth_manager.is_admin(user):
|
||||
return settings
|
||||
return scrub_settings(settings)
|
||||
|
||||
@router.post("/settings")
|
||||
async def set_settings(request: Request):
|
||||
"""Admin only: update app settings."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
current = _load_settings()
|
||||
# Per-key validation for numeric settings: coerce to int and clamp to a
|
||||
# sane range so a bad value can't disable the agent or let it run away.
|
||||
_INT_RANGES = {
|
||||
"agent_max_rounds": (1, 200),
|
||||
"agent_max_tool_calls": (0, 1000), # 0 = unlimited
|
||||
}
|
||||
for key in DEFAULT_SETTINGS:
|
||||
if key not in body:
|
||||
continue
|
||||
val = body[key]
|
||||
if key in _INT_RANGES:
|
||||
lo, hi = _INT_RANGES[key]
|
||||
try:
|
||||
val = int(val)
|
||||
except (TypeError, ValueError):
|
||||
raise HTTPException(400, f"{key} must be an integer")
|
||||
val = max(lo, min(val, hi))
|
||||
current[key] = val
|
||||
_save_settings(current)
|
||||
return current
|
||||
|
||||
# ---- Integrations CRUD ----
|
||||
|
||||
# Run migration on startup
|
||||
migrate_from_settings()
|
||||
|
||||
@router.get("/integrations")
|
||||
async def list_integrations_route(request: Request):
|
||||
"""List all integrations (admin only, keys masked)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
items = load_integrations()
|
||||
# Mask API keys for frontend display
|
||||
safe = [mask_integration_secret(item) for item in items]
|
||||
return {"integrations": safe}
|
||||
|
||||
@router.get("/integrations/presets")
|
||||
async def list_presets():
|
||||
"""List available integration presets."""
|
||||
return {"presets": {k: {kk: vv for kk, vv in v.items() if kk != "api_key"} for k, v in INTEGRATION_PRESETS.items()}}
|
||||
|
||||
@router.post("/integrations")
|
||||
async def create_integration(request: Request):
|
||||
"""Create a new integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
item = add_integration(body)
|
||||
return {"ok": True, "integration": mask_integration_secret(item)}
|
||||
|
||||
@router.put("/integrations/{integration_id}")
|
||||
async def update_integration_route(integration_id: str, request: Request):
|
||||
"""Update an existing integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
item = update_integration(integration_id, body)
|
||||
if not item:
|
||||
raise HTTPException(404, "Integration not found")
|
||||
return {"ok": True, "integration": mask_integration_secret(item)}
|
||||
|
||||
@router.delete("/integrations/{integration_id}")
|
||||
async def delete_integration_route(integration_id: str, request: Request):
|
||||
"""Delete an integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
ok = delete_integration(integration_id)
|
||||
if not ok:
|
||||
raise HTTPException(404, "Integration not found")
|
||||
return {"ok": True}
|
||||
|
||||
@router.post("/integrations/{integration_id}/test")
|
||||
async def test_integration_route(integration_id: str, request: Request):
|
||||
"""Test connectivity to an integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
integ = get_integration(integration_id)
|
||||
if not integ:
|
||||
raise HTTPException(404, "Integration not found")
|
||||
preset = (integ.get("preset") or integ.get("name", "")).lower()
|
||||
|
||||
# ntfy is special: a GET / proves the server is reachable but
|
||||
# publishes nothing, so the user has no way to know whether
|
||||
# subscribers will actually receive notifications. Instead, do
|
||||
# the real thing — POST a one-line "connectivity test" message
|
||||
# to the topic the Reminders panel is configured to use. If the
|
||||
# subscriber app is wired up correctly, this is what the green
|
||||
# checkmark + a phone ping confirms together.
|
||||
if preset == "ntfy":
|
||||
import httpx
|
||||
from urllib.parse import urlparse
|
||||
# Strip any path/query the user accidentally pasted in the
|
||||
# base URL (e.g. `http://host:8091/odysseus`) — otherwise
|
||||
# the topic gets appended after the path and we publish to
|
||||
# `/odysseus/odysseus` (which ntfy 404s on). ntfy itself
|
||||
# only ever serves from the root.
|
||||
raw_base = (integ.get("base_url") or "").strip()
|
||||
parsed = urlparse(raw_base)
|
||||
base = f"{parsed.scheme}://{parsed.netloc}" if parsed.scheme and parsed.netloc else raw_base.rstrip("/")
|
||||
settings = _load_settings()
|
||||
topic = (settings.get("reminder_ntfy_topic") or "reminders").strip() or "reminders"
|
||||
full_url = f"{base}/{topic}"
|
||||
api_key = integ.get("api_key", "")
|
||||
auth_type = (integ.get("auth_type") or "none").lower()
|
||||
headers = {
|
||||
"Title": "Odysseus connectivity test",
|
||||
"Tags": "white_check_mark",
|
||||
"Priority": "default",
|
||||
}
|
||||
if api_key:
|
||||
if auth_type == "bearer":
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
elif auth_type == "header":
|
||||
headers[integ.get("auth_header") or "Authorization"] = api_key
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=8.0) as client:
|
||||
r = await client.post(
|
||||
full_url,
|
||||
content="Connectivity test from Odysseus. If you see this on your phone, ntfy is wired up correctly.",
|
||||
headers=headers,
|
||||
)
|
||||
if r.is_success:
|
||||
# Tell the user EXACTLY where it went and what to
|
||||
# subscribe to on their phone, so they can match
|
||||
# without guesswork. The doubled-topic / wrong-host
|
||||
# mistakes are easier to spot when the actual URL
|
||||
# is right there in the success line.
|
||||
return {
|
||||
"ok": True,
|
||||
"message": (
|
||||
f"Sent to {full_url} — on your ntfy app, "
|
||||
f"subscribe to topic \"{topic}\" with server "
|
||||
f"\"{base}\" (or paste the full URL: {full_url})."
|
||||
),
|
||||
}
|
||||
return {"ok": False, "message": f"ntfy returned HTTP {r.status_code} from {full_url}: {r.text[:200]}"}
|
||||
except Exception as e:
|
||||
hint = ""
|
||||
if parsed.hostname not in ("127.0.0.1", "localhost"):
|
||||
hint = " If this is Docker Compose ntfy, set NTFY_BIND to that host/Tailscale IP and NTFY_BASE_URL to the same server URL in .env, then recreate ntfy."
|
||||
return {"ok": False, "message": f"ntfy publish to {full_url} failed: {e}.{hint}"[:500]}
|
||||
|
||||
if preset == "discord_webhook":
|
||||
import httpx
|
||||
webhook_url = (integ.get("base_url") or "").strip()
|
||||
if not webhook_url:
|
||||
return {"ok": False, "message": "No webhook URL set — paste the full Discord webhook URL into the Base URL field."}
|
||||
payload = {
|
||||
"embeds": [{
|
||||
"title": "Odysseus connectivity test",
|
||||
"description": "If you see this, your Discord Webhook integration is wired up correctly.",
|
||||
"color": 5793266,
|
||||
}]
|
||||
}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=8.0) as client:
|
||||
r = await client.post(webhook_url, json=payload)
|
||||
if r.is_success:
|
||||
return {"ok": True, "message": "Test embed sent — check your Discord channel to confirm it arrived."}
|
||||
return {"ok": False, "message": f"Discord returned HTTP {r.status_code}: {r.text[:200]}"}
|
||||
except Exception as e:
|
||||
return {"ok": False, "message": f"Request failed: {e}"[:400]}
|
||||
|
||||
# All other presets: GET against a known health endpoint.
|
||||
# Fall back to detecting from name if preset is missing.
|
||||
health_paths = {
|
||||
"miniflux": "/v1/me",
|
||||
"gitea": "/api/v1/version",
|
||||
"linkding": "/api/tags/",
|
||||
"homeassistant": "/api/",
|
||||
"home assistant": "/api/",
|
||||
}
|
||||
path = health_paths.get(preset, "/")
|
||||
result = await execute_api_call(integration_id, "GET", path)
|
||||
if result.get("exit_code", 1) == 0:
|
||||
return {"ok": True, "message": "Connection successful"}
|
||||
return {"ok": False, "message": (result.get("error") or "Connection failed")[:300]}
|
||||
|
||||
return router
|
||||
193
routes/auth/device_flow.py
Normal file
193
routes/auth/device_flow.py
Normal file
|
|
@ -0,0 +1,193 @@
|
|||
"""Shared OAuth/device-flow route scaffolding for provider setup."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Iterable, Mapping, Optional
|
||||
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
|
||||
from core.middleware import require_admin
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeviceFlowStart:
|
||||
"""Provider-specific start result consumed by the shared route wrapper."""
|
||||
|
||||
pending: Mapping[str, Any]
|
||||
response: Mapping[str, Any]
|
||||
interval: int = 5
|
||||
expires_in: int = 900
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeviceFlowPoll:
|
||||
"""Normalized provider poll outcome."""
|
||||
|
||||
status: str
|
||||
endpoint: Optional[Mapping[str, Any]] = None
|
||||
error: Optional[str] = None
|
||||
detail: Optional[str] = None
|
||||
interval: Optional[int] = None
|
||||
|
||||
@classmethod
|
||||
def pending(cls, detail: Optional[str] = None) -> "DeviceFlowPoll":
|
||||
return cls(status="pending", detail=detail)
|
||||
|
||||
@classmethod
|
||||
def slow_down(cls, interval: Optional[int] = None, detail: Optional[str] = None) -> "DeviceFlowPoll":
|
||||
return cls(status="slow_down", interval=interval, detail=detail)
|
||||
|
||||
@classmethod
|
||||
def authorized(cls, endpoint: Mapping[str, Any]) -> "DeviceFlowPoll":
|
||||
return cls(status="authorized", endpoint=endpoint)
|
||||
|
||||
@classmethod
|
||||
def failed(cls, error: str) -> "DeviceFlowPoll":
|
||||
return cls(status="failed", error=error)
|
||||
|
||||
|
||||
class PendingDeviceFlowStore:
|
||||
"""Thread-safe in-memory pending device-flow store.
|
||||
|
||||
Device codes and provider-side secrets stay inside this process. Each entry
|
||||
stores provider payload separately from poll metadata so provider callbacks
|
||||
only receive the fields they created.
|
||||
"""
|
||||
|
||||
def __init__(self, *, time_func: Callable[[], float] = time.time):
|
||||
self._pending: dict[str, dict[str, Any]] = {}
|
||||
self._lock = threading.Lock()
|
||||
self._time = time_func
|
||||
|
||||
def _now(self) -> float:
|
||||
return float(self._time())
|
||||
|
||||
def prune_expired(self) -> None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
for key in [k for k, v in self._pending.items() if v.get("expires_at", 0) < now]:
|
||||
self._pending.pop(key, None)
|
||||
|
||||
def add(self, payload: Mapping[str, Any], *, interval: int, expires_in: int) -> str:
|
||||
self.prune_expired()
|
||||
poll_id = uuid.uuid4().hex
|
||||
with self._lock:
|
||||
self._pending[poll_id] = {
|
||||
"payload": dict(payload),
|
||||
"interval": max(int(interval or 5), 1),
|
||||
"expires_at": self._now() + max(int(expires_in or 900), 1),
|
||||
"next_poll_at": 0.0,
|
||||
}
|
||||
return poll_id
|
||||
|
||||
def get_payload(self, poll_id: str) -> Optional[dict[str, Any]]:
|
||||
self.prune_expired()
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
if entry is None:
|
||||
return None
|
||||
return dict(entry.get("payload") or {})
|
||||
|
||||
def is_throttled(self, poll_id: str) -> bool:
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
return bool(entry and self._now() < float(entry.get("next_poll_at") or 0))
|
||||
|
||||
def schedule_next(self, poll_id: str) -> None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
if entry is not None:
|
||||
entry["next_poll_at"] = now + int(entry.get("interval") or 5)
|
||||
|
||||
def slow_down(self, poll_id: str, interval: Optional[int] = None) -> None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
if entry is not None:
|
||||
new_interval = int(interval or (int(entry.get("interval") or 5) + 5))
|
||||
entry["interval"] = max(new_interval, 1)
|
||||
entry["next_poll_at"] = now + entry["interval"]
|
||||
|
||||
def pop(self, poll_id: str) -> None:
|
||||
with self._lock:
|
||||
self._pending.pop(poll_id, None)
|
||||
|
||||
|
||||
async def _maybe_await(value: Any) -> Any:
|
||||
if inspect.isawaitable(value):
|
||||
return await value
|
||||
return value
|
||||
|
||||
|
||||
def _pending_response(detail: Optional[str] = None) -> dict[str, Any]:
|
||||
response: dict[str, Any] = {"status": "pending"}
|
||||
if detail:
|
||||
response["detail"] = detail
|
||||
return response
|
||||
|
||||
|
||||
def create_device_flow_router(
|
||||
*,
|
||||
prefix: str,
|
||||
tags: Iterable[str],
|
||||
store: PendingDeviceFlowStore,
|
||||
start_flow: Callable[[Request, Mapping[str, Any]], DeviceFlowStart],
|
||||
poll_flow: Callable[[Request, Mapping[str, Any]], DeviceFlowPoll],
|
||||
) -> APIRouter:
|
||||
"""Create standard `/device/start|poll|cancel` routes for a provider."""
|
||||
|
||||
router = APIRouter(prefix=prefix, tags=list(tags))
|
||||
|
||||
@router.post("/device/start")
|
||||
async def device_start(request: Request):
|
||||
require_admin(request)
|
||||
form = await request.form()
|
||||
start = await _maybe_await(start_flow(request, form))
|
||||
interval = int(start.interval or 5)
|
||||
expires_in = int(start.expires_in or 900)
|
||||
poll_id = store.add(start.pending, interval=interval, expires_in=expires_in)
|
||||
response = dict(start.response)
|
||||
response.update({"poll_id": poll_id, "interval": interval, "expires_in": expires_in})
|
||||
return response
|
||||
|
||||
@router.post("/device/poll")
|
||||
async def device_poll(request: Request, poll_id: str = Form(...)):
|
||||
require_admin(request)
|
||||
payload = store.get_payload(poll_id)
|
||||
if payload is None:
|
||||
raise HTTPException(404, "Unknown or expired login session")
|
||||
if store.is_throttled(poll_id):
|
||||
return {"status": "pending"}
|
||||
|
||||
try:
|
||||
outcome = await _maybe_await(poll_flow(request, payload))
|
||||
except Exception:
|
||||
store.pop(poll_id)
|
||||
raise
|
||||
|
||||
if outcome.status == "authorized":
|
||||
store.pop(poll_id)
|
||||
return {"status": "authorized", "endpoint": dict(outcome.endpoint or {})}
|
||||
if outcome.status == "failed":
|
||||
store.pop(poll_id)
|
||||
return {"status": "failed", "error": outcome.error or "denied"}
|
||||
if outcome.status == "slow_down":
|
||||
store.slow_down(poll_id, outcome.interval)
|
||||
return _pending_response(outcome.detail)
|
||||
|
||||
store.schedule_next(poll_id)
|
||||
return _pending_response(outcome.detail)
|
||||
|
||||
@router.post("/device/cancel")
|
||||
def device_cancel(request: Request, poll_id: str = Form(...)):
|
||||
require_admin(request)
|
||||
store.pop(poll_id)
|
||||
return {"status": "cancelled"}
|
||||
|
||||
return router
|
||||
|
|
@ -1,836 +1,18 @@
|
|||
"""Authentication routes — login, logout, signup, status, user management."""
|
||||
"""Backward-compat shim — canonical location is routes/auth/auth_routes.py.
|
||||
|
||||
from fastapi import APIRouter, Request, Response, HTTPException
|
||||
from pydantic import BaseModel
|
||||
from typing import Optional
|
||||
import asyncio
|
||||
import logging
|
||||
import os
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.auth_routes``, ``from routes.auth_routes import X``,
|
||||
``importlib.import_module("routes.auth_routes")``, the
|
||||
``sys.modules.pop("routes.auth_routes")`` + re-import pattern in
|
||||
test_auth_policy.py / test_auth_session_revocation.py, and the
|
||||
``import ... as ar`` + ``monkeypatch.setattr(ar, ...)`` pattern in
|
||||
test_rename_user_owner_sync.py / test_integrations_store_shape.py all operate
|
||||
on the *same* object the application actually uses. Keeps existing import
|
||||
paths working after slice 2n (#4082/#4071).
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
from pathlib import Path
|
||||
import sys as _sys
|
||||
|
||||
from core.atomic_io import atomic_write_json, atomic_write_text
|
||||
from core.auth import AuthManager, RESERVED_USERNAMES, SetAdminResult, TOKEN_TTL
|
||||
from src.constants import DEEP_RESEARCH_DIR, MEMORY_FILE, PASSWORD_MIN_LENGTH, SKILLS_DIR
|
||||
from src.rate_limiter import RateLimiter
|
||||
from src.settings_scrub import scrub_settings
|
||||
from src.settings import (
|
||||
load_settings as _load_settings,
|
||||
save_settings as _save_settings,
|
||||
load_features as _load_features,
|
||||
save_features as _save_features,
|
||||
DEFAULT_SETTINGS,
|
||||
)
|
||||
from src.integrations import (
|
||||
load_integrations,
|
||||
add_integration,
|
||||
update_integration,
|
||||
delete_integration,
|
||||
get_integration,
|
||||
mask_integration_secret,
|
||||
execute_api_call,
|
||||
INTEGRATION_PRESETS,
|
||||
migrate_from_settings,
|
||||
)
|
||||
from routes.auth import auth_routes as _canonical # noqa: F401
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
remember: bool = True
|
||||
totp_code: Optional[str] = None
|
||||
|
||||
|
||||
class SetupRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
class SignupRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
class ChangePasswordRequest(BaseModel):
|
||||
current_password: str
|
||||
new_password: str
|
||||
|
||||
|
||||
class CreateUserRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
is_admin: bool = False
|
||||
|
||||
|
||||
class DeleteUserRequest(BaseModel):
|
||||
username: str
|
||||
|
||||
|
||||
class RenameUserRequest(BaseModel):
|
||||
username: str
|
||||
|
||||
|
||||
class SetAdminRequest(BaseModel):
|
||||
is_admin: bool
|
||||
|
||||
|
||||
class SetOpenRegistrationRequest(BaseModel):
|
||||
enabled: bool
|
||||
|
||||
SESSION_COOKIE = "odysseus_session"
|
||||
|
||||
|
||||
def setup_auth_routes(auth_manager: AuthManager) -> APIRouter:
|
||||
router = APIRouter(prefix="/api/auth", tags=["auth"])
|
||||
|
||||
_login_limiter = RateLimiter(max_requests=15, window_seconds=60)
|
||||
_signup_limiter = RateLimiter(max_requests=3, window_seconds=300)
|
||||
_setup_limiter = RateLimiter(max_requests=3, window_seconds=300)
|
||||
|
||||
def _get_current_user(request: Request) -> Optional[str]:
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
return auth_manager.get_username_for_token(token)
|
||||
|
||||
@router.post("/setup")
|
||||
async def first_run_setup(body: SetupRequest, request: Request):
|
||||
"""Create initial admin account. Only works if no accounts exist."""
|
||||
if not _setup_limiter.check(request.client.host):
|
||||
raise HTTPException(429, "Too many requests — try again later")
|
||||
if auth_manager.is_configured:
|
||||
raise HTTPException(400, "Already configured")
|
||||
if len(body.password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
if len(body.username.strip()) < 1:
|
||||
raise HTTPException(400, "Username is required")
|
||||
if body.username.lower() in RESERVED_USERNAMES:
|
||||
raise HTTPException(403, "Username is reserved")
|
||||
ok = await asyncio.to_thread(auth_manager.setup, body.username, body.password)
|
||||
if not ok:
|
||||
raise HTTPException(500, "Setup failed")
|
||||
return {"ok": True, "message": "Admin account created"}
|
||||
|
||||
@router.post("/signup")
|
||||
async def signup(body: SignupRequest, request: Request):
|
||||
"""Create a new user account. Only works if signup is enabled by admin."""
|
||||
if not _signup_limiter.check(request.client.host):
|
||||
raise HTTPException(429, "Too many requests — try again later")
|
||||
if not auth_manager.is_configured:
|
||||
raise HTTPException(400, "Run setup first")
|
||||
if not auth_manager.signup_enabled:
|
||||
raise HTTPException(403, "Registration is disabled. Ask an admin for an account.")
|
||||
if len(body.password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
if len(body.username.strip()) < 1:
|
||||
raise HTTPException(400, "Username is required")
|
||||
if body.username.lower() in RESERVED_USERNAMES:
|
||||
raise HTTPException(403, "Username is reserved")
|
||||
ok = await asyncio.to_thread(auth_manager.create_user, body.username, body.password, is_admin=False)
|
||||
if not ok:
|
||||
raise HTTPException(409, "Username already taken")
|
||||
return {"ok": True, "message": "Account created"}
|
||||
|
||||
@router.post("/login")
|
||||
async def login(body: LoginRequest, request: Request, response: Response):
|
||||
if not _login_limiter.check(request.client.host):
|
||||
raise HTTPException(429, "Too many requests — try again later")
|
||||
# Verify password first
|
||||
username = body.username.strip().lower()
|
||||
if not await asyncio.to_thread(auth_manager.verify_password, username, body.password):
|
||||
raise HTTPException(401, "Invalid credentials")
|
||||
# Check 2FA if enabled
|
||||
if auth_manager.totp_enabled(username):
|
||||
if not body.totp_code:
|
||||
# Password OK but need TOTP — tell client to show code input
|
||||
return {"ok": False, "requires_totp": True, "username": username}
|
||||
if not auth_manager.totp_verify(username, body.totp_code):
|
||||
raise HTTPException(401, "Invalid 2FA code")
|
||||
# All checks passed — create session (password already verified above)
|
||||
token = await asyncio.to_thread(auth_manager.create_session_trusted, username)
|
||||
if not token:
|
||||
raise HTTPException(401, "Invalid credentials")
|
||||
cookie_kwargs = dict(
|
||||
key=SESSION_COOKIE,
|
||||
value=token,
|
||||
httponly=True,
|
||||
samesite="lax",
|
||||
secure=os.getenv("SECURE_COOKIES", "false").lower() == "true",
|
||||
path="/",
|
||||
)
|
||||
if body.remember:
|
||||
cookie_kwargs["max_age"] = TOKEN_TTL
|
||||
response.set_cookie(**cookie_kwargs)
|
||||
return {"ok": True, "username": username}
|
||||
|
||||
@router.post("/logout")
|
||||
async def logout(request: Request, response: Response):
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
if token:
|
||||
auth_manager.revoke_token(token)
|
||||
response.delete_cookie(SESSION_COOKIE, path="/")
|
||||
return {"ok": True}
|
||||
|
||||
@router.get("/status")
|
||||
async def auth_status(request: Request):
|
||||
token = request.cookies.get(SESSION_COOKIE)
|
||||
result = auth_manager.status(token)
|
||||
result["signup_enabled"] = auth_manager.signup_enabled
|
||||
# Include the caller's effective privileges so the frontend can
|
||||
# hide / dim UI controls the user isn't allowed to use. Admins get
|
||||
# ADMIN_PRIVILEGES (everything on), regular users get their stored
|
||||
# set merged with DEFAULT_PRIVILEGES.
|
||||
try:
|
||||
u = result.get("username")
|
||||
if u:
|
||||
result["privileges"] = auth_manager.get_privileges(u)
|
||||
except Exception:
|
||||
pass
|
||||
return result
|
||||
|
||||
@router.get("/policy")
|
||||
async def auth_policy():
|
||||
"""Return public auth policy constants for the frontend."""
|
||||
return auth_manager.policy()
|
||||
|
||||
@router.post("/change-password")
|
||||
async def change_password(body: ChangePasswordRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if len(body.new_password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
current_token = request.cookies.get(SESSION_COOKIE)
|
||||
ok = await asyncio.to_thread(auth_manager.change_password, user, body.current_password, body.new_password)
|
||||
if not ok:
|
||||
raise HTTPException(400, "Current password is incorrect")
|
||||
await asyncio.to_thread(auth_manager.revoke_user_sessions, user, current_token)
|
||||
return {"ok": True}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Two-factor authentication
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@router.post("/2fa/setup")
|
||||
async def totp_setup(request: Request):
|
||||
"""Generate a TOTP secret and return the QR code URI."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if auth_manager.totp_enabled(user):
|
||||
raise HTTPException(400, "2FA is already enabled")
|
||||
secret = auth_manager.totp_generate_secret(user)
|
||||
if not secret:
|
||||
raise HTTPException(500, "Failed to generate secret")
|
||||
uri = auth_manager.totp_get_provisioning_uri(user, secret)
|
||||
# Generate QR code as base64 PNG
|
||||
import qrcode, io, base64
|
||||
qr = qrcode.make(uri, box_size=6, border=2)
|
||||
buf = io.BytesIO()
|
||||
qr.save(buf, format="PNG")
|
||||
qr_b64 = base64.b64encode(buf.getvalue()).decode("ascii")
|
||||
return {"secret": secret, "uri": uri, "qr_code": f"data:image/png;base64,{qr_b64}"}
|
||||
|
||||
class TotpVerifyRequest(BaseModel):
|
||||
code: str
|
||||
|
||||
@router.post("/2fa/confirm")
|
||||
async def totp_confirm(body: TotpVerifyRequest, request: Request):
|
||||
"""Verify a TOTP code to confirm 2FA setup. Returns backup codes."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if not auth_manager.totp_confirm_enable(user, body.code):
|
||||
raise HTTPException(400, "Invalid code — try again")
|
||||
backup = auth_manager.users.get(user, {}).get("totp_backup_codes", [])
|
||||
return {"ok": True, "backup_codes": backup}
|
||||
|
||||
class TotpDisableRequest(BaseModel):
|
||||
password: str
|
||||
|
||||
@router.post("/2fa/disable")
|
||||
async def totp_disable(body: TotpDisableRequest, request: Request):
|
||||
"""Disable 2FA. Requires password confirmation."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
if not auth_manager.totp_disable(user, body.password):
|
||||
raise HTTPException(400, "Invalid password")
|
||||
return {"ok": True}
|
||||
|
||||
@router.get("/2fa/status")
|
||||
async def totp_status(request: Request):
|
||||
"""Check if 2FA is enabled for the current user."""
|
||||
user = _get_current_user(request)
|
||||
if not user:
|
||||
raise HTTPException(401, "Not authenticated")
|
||||
return {"enabled": auth_manager.totp_enabled(user)}
|
||||
|
||||
# Admin-only routes
|
||||
@router.get("/users")
|
||||
async def list_users(request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
return {"users": auth_manager.list_users()}
|
||||
|
||||
@router.post("/users")
|
||||
async def admin_create_user(body: CreateUserRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
if len(body.password) < PASSWORD_MIN_LENGTH:
|
||||
raise HTTPException(400, f"Password must be at least {PASSWORD_MIN_LENGTH} characters")
|
||||
if len(body.username.strip()) < 1:
|
||||
raise HTTPException(400, "Username is required")
|
||||
if body.username.lower() in RESERVED_USERNAMES:
|
||||
raise HTTPException(403, "Username is reserved")
|
||||
ok = auth_manager.create_user(body.username, body.password, body.is_admin)
|
||||
if not ok:
|
||||
raise HTTPException(409, "Username already taken")
|
||||
return {"ok": True}
|
||||
|
||||
@router.put("/users/{username}/privileges")
|
||||
async def update_user_privileges(username: str, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
ok = auth_manager.set_privileges(username, body)
|
||||
if not ok:
|
||||
raise HTTPException(404, "User not found or is admin")
|
||||
return {"ok": True, "privileges": auth_manager.get_privileges(username)}
|
||||
|
||||
@router.put("/users/{username}/rename")
|
||||
async def rename_user(username: str, body: RenameUserRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
old_username = (username or "").strip().lower()
|
||||
new_username = (body.username or "").strip().lower()
|
||||
if not new_username:
|
||||
raise HTTPException(400, "Username required")
|
||||
if old_username == new_username:
|
||||
return {"ok": True, "username": new_username, "renamed_self": old_username == user}
|
||||
if old_username not in auth_manager.users:
|
||||
raise HTTPException(404, "User not found")
|
||||
if new_username in auth_manager.users:
|
||||
raise HTTPException(409, "Username already taken")
|
||||
|
||||
# Gate on auth first. Every mutation below is contingent on this
|
||||
# succeeding — doing it last meant a rejected rename (e.g. reserved
|
||||
# username) left file-backed owner fields already rewritten with no
|
||||
# way to roll them back.
|
||||
ok = auth_manager.rename_user(old_username, new_username, user)
|
||||
if not ok:
|
||||
raise HTTPException(400, "Cannot rename user")
|
||||
|
||||
def _rollback_auth_rename() -> bool:
|
||||
# On self-rename the admin session has already moved to the new
|
||||
# username, so the rollback must authenticate as the new user.
|
||||
rollback_user = new_username if user == old_username else user
|
||||
try:
|
||||
return bool(auth_manager.rename_user(new_username, old_username, rollback_user))
|
||||
except Exception as rollback_err:
|
||||
logger.error(
|
||||
"Failed to roll back auth rename %s -> %s after owner migration failure: %s",
|
||||
new_username, old_username, rollback_err,
|
||||
)
|
||||
return False
|
||||
|
||||
# Usernames are ownership keys for user data. Rename the common
|
||||
# owner-scoped DB rows so the account keeps access to its sessions,
|
||||
# docs, email accounts, tasks, etc.
|
||||
try:
|
||||
from sqlalchemy import func
|
||||
from core.database import Base, SessionLocal
|
||||
db = SessionLocal()
|
||||
try:
|
||||
for mapper in Base.registry.mappers:
|
||||
model = mapper.class_
|
||||
if not hasattr(model, "owner"):
|
||||
continue
|
||||
(
|
||||
db.query(model)
|
||||
.filter(func.lower(model.owner) == old_username)
|
||||
.update({"owner": new_username}, synchronize_session=False)
|
||||
)
|
||||
db.commit()
|
||||
except Exception:
|
||||
db.rollback()
|
||||
raise
|
||||
finally:
|
||||
db.close()
|
||||
except Exception as e:
|
||||
logger.error("Failed to rename owner references %s -> %s: %s", old_username, new_username, e)
|
||||
if not _rollback_auth_rename():
|
||||
logger.error(
|
||||
"Auth rename %s -> %s could not be rolled back after owner migration failure",
|
||||
old_username, new_username,
|
||||
)
|
||||
raise HTTPException(500, "Failed to rename user data")
|
||||
|
||||
# Per-user prefs are JSON-backed, not SQL-backed.
|
||||
try:
|
||||
from routes.prefs_routes import _load as _load_prefs, _save as _save_prefs
|
||||
prefs = _load_prefs()
|
||||
users = prefs.get("_users") if isinstance(prefs, dict) else None
|
||||
if isinstance(users, dict):
|
||||
prefs_key = next(
|
||||
(k for k in users if str(k).strip().lower() == old_username),
|
||||
None,
|
||||
)
|
||||
new_taken = any(str(k).strip().lower() == new_username for k in users)
|
||||
if prefs_key is not None and not new_taken:
|
||||
users[new_username] = users.pop(prefs_key)
|
||||
_save_prefs(prefs)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename user prefs %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# In-flight deep-research tasks live in the process-local
|
||||
# ResearchHandler registry. They are not covered by the persisted JSON
|
||||
# migration above, but the research routes filter and cancel by this
|
||||
# owner field while the job is running. Do this before sweeping
|
||||
# completed JSON files so a job that finishes during the rename saves
|
||||
# with the new owner or is caught by the disk sweep below.
|
||||
try:
|
||||
rh = getattr(request.app.state, "research_handler", None)
|
||||
rename_owner = getattr(rh, "rename_owner", None)
|
||||
if callable(rename_owner):
|
||||
rename_owner(old_username, new_username)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename active research tasks %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# deep_research: each completed report is a standalone JSON file with
|
||||
# an `owner` field. research_routes filters by d.get("owner") == user,
|
||||
# so a stale owner makes every report invisible to the renamed user.
|
||||
try:
|
||||
dr_dir = Path(DEEP_RESEARCH_DIR)
|
||||
if dr_dir.is_dir():
|
||||
for p in dr_dir.glob("*.json"):
|
||||
try:
|
||||
d = json.loads(p.read_text(encoding="utf-8"))
|
||||
if str(d.get("owner", "")).strip().lower() == old_username:
|
||||
d["owner"] = new_username
|
||||
atomic_write_json(str(p), d)
|
||||
except Exception as err:
|
||||
logger.warning("Failed to update research owner in %s: %s", p.name, err)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename research owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# memory.json: a flat JSON array where each entry carries an `owner`
|
||||
# field. memory_manager.load(owner=user) filters on it, so stale
|
||||
# entries disappear from the memory panel.
|
||||
try:
|
||||
if os.path.isfile(MEMORY_FILE):
|
||||
with open(MEMORY_FILE, encoding="utf-8") as fh:
|
||||
entries = json.loads(fh.read())
|
||||
if isinstance(entries, list):
|
||||
changed = False
|
||||
for entry in entries:
|
||||
if isinstance(entry, dict) and str(entry.get("owner", "")).strip().lower() == old_username:
|
||||
entry["owner"] = new_username
|
||||
changed = True
|
||||
if changed:
|
||||
atomic_write_json(MEMORY_FILE, entries)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename memory.json owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# uploads.json: upload rows use owner metadata for access checks and
|
||||
# owner-prefixed index keys for dedupe. Rename both so attachments keep
|
||||
# resolving after the account username changes.
|
||||
try:
|
||||
upload_handler = getattr(request.app.state, "upload_handler", None)
|
||||
rename_owner = getattr(upload_handler, "rename_owner", None)
|
||||
if callable(rename_owner):
|
||||
rename_owner(old_username, new_username)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename upload owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# direct personal RAG uploads live in per-owner directories and the
|
||||
# vector metadata also carries the username used for owner-filtered
|
||||
# search. Keep both in sync with the auth rename.
|
||||
try:
|
||||
from routes.personal_routes import rename_personal_upload_owner
|
||||
personal_docs_manager = getattr(request.app.state, "personal_docs_manager", None)
|
||||
if personal_docs_manager is not None:
|
||||
rag_manager = getattr(personal_docs_manager, "rag_manager", None)
|
||||
rename_personal_upload_owner(
|
||||
old_username,
|
||||
new_username,
|
||||
personal_docs_manager=personal_docs_manager,
|
||||
rag_manager=rag_manager,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename personal RAG upload owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# skills: SKILL.md frontmatter carries owner: <username>; the usage
|
||||
# sidecar (_usage.json) keys entries as owner::skill-name. Both must
|
||||
# be updated or the renamed user's Skills panel goes empty.
|
||||
try:
|
||||
skills_root = Path(SKILLS_DIR)
|
||||
if skills_root.is_dir():
|
||||
_owner_re = re.compile(
|
||||
r'(?m)^(owner:\s*)' + re.escape(old_username) + r'\s*$',
|
||||
re.IGNORECASE,
|
||||
)
|
||||
for p in skills_root.rglob("SKILL.md"):
|
||||
try:
|
||||
text = p.read_text(encoding="utf-8")
|
||||
new_text = _owner_re.sub(r'\g<1>' + new_username, text)
|
||||
if new_text != text:
|
||||
atomic_write_text(str(p), new_text)
|
||||
except Exception as err:
|
||||
logger.warning("Failed to update skill owner in %s: %s", p, err)
|
||||
usage_path = skills_root / "_usage.json"
|
||||
if usage_path.is_file():
|
||||
try:
|
||||
usage = json.loads(usage_path.read_text(encoding="utf-8"))
|
||||
if isinstance(usage, dict):
|
||||
new_usage = {}
|
||||
changed = False
|
||||
for k, v in usage.items():
|
||||
owner_part, sep, skill_part = k.partition("::")
|
||||
if sep and owner_part.lower() == old_username:
|
||||
new_usage[new_username + "::" + skill_part] = v
|
||||
changed = True
|
||||
else:
|
||||
new_usage[k] = v
|
||||
if changed:
|
||||
atomic_write_json(str(usage_path), new_usage)
|
||||
except Exception as err:
|
||||
logger.warning("Failed to update skills usage keys %s -> %s: %s", old_username, new_username, err)
|
||||
except Exception as e:
|
||||
logger.warning("Failed to rename skills owner references %s -> %s: %s", old_username, new_username, e)
|
||||
|
||||
# The in-memory session cache (session_manager.sessions) stores each
|
||||
# session's owner at load time. Without this patch the renamed user's
|
||||
# sessions are invisible on the next /api/sessions call because
|
||||
# get_sessions_for_user does an exact `s.owner == username` comparison
|
||||
# against stale in-memory values.
|
||||
sm = getattr(request.app.state, "session_manager", None)
|
||||
if sm is not None:
|
||||
for sess in list(getattr(sm, "sessions", {}).values()):
|
||||
if str(getattr(sess, "owner", None) or "").strip().lower() == old_username:
|
||||
sess.owner = new_username
|
||||
|
||||
# The owner-rename loop above updated ApiToken.owner in the DB, but the
|
||||
# bearer-token cache still maps each token to the OLD owner. Without
|
||||
# refreshing it, the renamed user's API tokens resolve to the old (now
|
||||
# non-existent) owner and stop reaching their data until the cache next
|
||||
# goes dirty. Invalidate it now, like the token CRUD routes do.
|
||||
invalidator = getattr(request.app.state, "invalidate_token_cache", None)
|
||||
if callable(invalidator):
|
||||
invalidator()
|
||||
return {"ok": True, "username": new_username, "renamed_self": old_username == user}
|
||||
|
||||
@router.put("/users/{username}/admin")
|
||||
async def set_user_admin(username: str, body: SetAdminRequest, request: Request):
|
||||
"""Promote/demote a user to/from admin. Admin only.
|
||||
|
||||
The last remaining admin can't be demoted (no lockout). Self-demotion
|
||||
is allowed while another admin exists; the `self` flag tells the UI to
|
||||
reload the acting user into the normal-user view.
|
||||
"""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
result = auth_manager.set_admin(username, body.is_admin, user)
|
||||
if result is SetAdminResult.USER_NOT_FOUND:
|
||||
raise HTTPException(404, "User not found")
|
||||
if result is SetAdminResult.NOT_AUTHORIZED:
|
||||
raise HTTPException(403, "Admin only")
|
||||
if result is SetAdminResult.LAST_ADMIN:
|
||||
raise HTTPException(400, "Cannot demote the last admin")
|
||||
target = (username or "").strip().lower()
|
||||
return {
|
||||
"ok": True,
|
||||
"is_admin": body.is_admin,
|
||||
"self": target == (user or "").strip().lower(),
|
||||
}
|
||||
|
||||
@router.post("/signup-toggle", deprecated=True)
|
||||
async def toggle_signup(request: Request):
|
||||
"""
|
||||
Toggle open registration on/off. Admin only.
|
||||
|
||||
DEPRECATED: This endpoint uses toggle semantics which can lead to unsafe state changes.
|
||||
Use PUT /open-signup instead.
|
||||
|
||||
This endpoint is kept for backward compatibility and may be removed in future versions.
|
||||
"""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
auth_manager.signup_enabled = not auth_manager.signup_enabled
|
||||
return {"ok": True, "signup_enabled": auth_manager.signup_enabled}
|
||||
|
||||
@router.put("/open-signup")
|
||||
async def set_signup_enabled(body: SetOpenRegistrationRequest, request: Request):
|
||||
"""Set open signup enabled state. Admin only."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
auth_manager.signup_enabled = body.enabled
|
||||
return {"ok": True,"signup_enabled": auth_manager.signup_enabled}
|
||||
|
||||
@router.delete("/users")
|
||||
async def admin_delete_user(body: DeleteUserRequest, request: Request):
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
|
||||
def _invalidate_api_token_cache():
|
||||
try:
|
||||
invalidator = getattr(request.app.state, "invalidate_token_cache", None)
|
||||
if invalidator:
|
||||
invalidator()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
ok = auth_manager.delete_user(body.username, user)
|
||||
except Exception:
|
||||
# delete_user can touch ApiToken rows before a later auth-store write
|
||||
# fails. Dirty the bearer cache anyway so a partial token purge does
|
||||
# not leave already-cached tokens authenticating until restart.
|
||||
_invalidate_api_token_cache()
|
||||
raise
|
||||
if not ok:
|
||||
raise HTTPException(400, "Cannot delete user")
|
||||
# delete_user removes the user's ApiToken rows, but the bearer-auth
|
||||
# middleware serves from an in-memory prefix->token cache that only
|
||||
# rebuilds when flagged dirty. Without this, a deleted user's already
|
||||
# cached token keeps authenticating until some other token op or a
|
||||
# restart clears the cache. Mirror what the token routes do.
|
||||
_invalidate_api_token_cache()
|
||||
return {"ok": True}
|
||||
|
||||
# ---- Feature visibility (admin-managed) ----
|
||||
|
||||
@router.get("/features")
|
||||
async def get_features():
|
||||
"""Public: returns which UI features are enabled."""
|
||||
return _load_features()
|
||||
|
||||
@router.post("/features")
|
||||
async def set_features(request: Request):
|
||||
"""Admin only: update feature toggles."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
current = _load_features()
|
||||
for key in current:
|
||||
if key in body and isinstance(body[key], bool):
|
||||
current[key] = body[key]
|
||||
_save_features(current)
|
||||
return current
|
||||
|
||||
# ---- App settings (admin-managed) ----
|
||||
|
||||
@router.get("/settings")
|
||||
async def get_settings(request: Request):
|
||||
"""Returns app settings. Admins get the full set; non-admins get
|
||||
a scrubbed copy with secret keys blanked. The frontend uses this
|
||||
for keybinds + TTS prefs, so it stays callable without admin."""
|
||||
user = _get_current_user(request)
|
||||
settings = _load_settings()
|
||||
if user and auth_manager.is_admin(user):
|
||||
return settings
|
||||
return scrub_settings(settings)
|
||||
|
||||
@router.post("/settings")
|
||||
async def set_settings(request: Request):
|
||||
"""Admin only: update app settings."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
current = _load_settings()
|
||||
# Per-key validation for numeric settings: coerce to int and clamp to a
|
||||
# sane range so a bad value can't disable the agent or let it run away.
|
||||
_INT_RANGES = {
|
||||
"agent_max_rounds": (1, 200),
|
||||
"agent_max_tool_calls": (0, 1000), # 0 = unlimited
|
||||
}
|
||||
for key in DEFAULT_SETTINGS:
|
||||
if key not in body:
|
||||
continue
|
||||
val = body[key]
|
||||
if key in _INT_RANGES:
|
||||
lo, hi = _INT_RANGES[key]
|
||||
try:
|
||||
val = int(val)
|
||||
except (TypeError, ValueError):
|
||||
raise HTTPException(400, f"{key} must be an integer")
|
||||
val = max(lo, min(val, hi))
|
||||
current[key] = val
|
||||
_save_settings(current)
|
||||
return current
|
||||
|
||||
# ---- Integrations CRUD ----
|
||||
|
||||
# Run migration on startup
|
||||
migrate_from_settings()
|
||||
|
||||
@router.get("/integrations")
|
||||
async def list_integrations_route(request: Request):
|
||||
"""List all integrations (admin only, keys masked)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
items = load_integrations()
|
||||
# Mask API keys for frontend display
|
||||
safe = [mask_integration_secret(item) for item in items]
|
||||
return {"integrations": safe}
|
||||
|
||||
@router.get("/integrations/presets")
|
||||
async def list_presets():
|
||||
"""List available integration presets."""
|
||||
return {"presets": {k: {kk: vv for kk, vv in v.items() if kk != "api_key"} for k, v in INTEGRATION_PRESETS.items()}}
|
||||
|
||||
@router.post("/integrations")
|
||||
async def create_integration(request: Request):
|
||||
"""Create a new integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
item = add_integration(body)
|
||||
return {"ok": True, "integration": mask_integration_secret(item)}
|
||||
|
||||
@router.put("/integrations/{integration_id}")
|
||||
async def update_integration_route(integration_id: str, request: Request):
|
||||
"""Update an existing integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
body = await request.json()
|
||||
item = update_integration(integration_id, body)
|
||||
if not item:
|
||||
raise HTTPException(404, "Integration not found")
|
||||
return {"ok": True, "integration": mask_integration_secret(item)}
|
||||
|
||||
@router.delete("/integrations/{integration_id}")
|
||||
async def delete_integration_route(integration_id: str, request: Request):
|
||||
"""Delete an integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
ok = delete_integration(integration_id)
|
||||
if not ok:
|
||||
raise HTTPException(404, "Integration not found")
|
||||
return {"ok": True}
|
||||
|
||||
@router.post("/integrations/{integration_id}/test")
|
||||
async def test_integration_route(integration_id: str, request: Request):
|
||||
"""Test connectivity to an integration (admin only)."""
|
||||
user = _get_current_user(request)
|
||||
if not user or not auth_manager.is_admin(user):
|
||||
raise HTTPException(403, "Admin only")
|
||||
integ = get_integration(integration_id)
|
||||
if not integ:
|
||||
raise HTTPException(404, "Integration not found")
|
||||
preset = (integ.get("preset") or integ.get("name", "")).lower()
|
||||
|
||||
# ntfy is special: a GET / proves the server is reachable but
|
||||
# publishes nothing, so the user has no way to know whether
|
||||
# subscribers will actually receive notifications. Instead, do
|
||||
# the real thing — POST a one-line "connectivity test" message
|
||||
# to the topic the Reminders panel is configured to use. If the
|
||||
# subscriber app is wired up correctly, this is what the green
|
||||
# checkmark + a phone ping confirms together.
|
||||
if preset == "ntfy":
|
||||
import httpx
|
||||
from urllib.parse import urlparse
|
||||
# Strip any path/query the user accidentally pasted in the
|
||||
# base URL (e.g. `http://host:8091/odysseus`) — otherwise
|
||||
# the topic gets appended after the path and we publish to
|
||||
# `/odysseus/odysseus` (which ntfy 404s on). ntfy itself
|
||||
# only ever serves from the root.
|
||||
raw_base = (integ.get("base_url") or "").strip()
|
||||
parsed = urlparse(raw_base)
|
||||
base = f"{parsed.scheme}://{parsed.netloc}" if parsed.scheme and parsed.netloc else raw_base.rstrip("/")
|
||||
settings = _load_settings()
|
||||
topic = (settings.get("reminder_ntfy_topic") or "reminders").strip() or "reminders"
|
||||
full_url = f"{base}/{topic}"
|
||||
api_key = integ.get("api_key", "")
|
||||
auth_type = (integ.get("auth_type") or "none").lower()
|
||||
headers = {
|
||||
"Title": "Odysseus connectivity test",
|
||||
"Tags": "white_check_mark",
|
||||
"Priority": "default",
|
||||
}
|
||||
if api_key:
|
||||
if auth_type == "bearer":
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
elif auth_type == "header":
|
||||
headers[integ.get("auth_header") or "Authorization"] = api_key
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=8.0) as client:
|
||||
r = await client.post(
|
||||
full_url,
|
||||
content="Connectivity test from Odysseus. If you see this on your phone, ntfy is wired up correctly.",
|
||||
headers=headers,
|
||||
)
|
||||
if r.is_success:
|
||||
# Tell the user EXACTLY where it went and what to
|
||||
# subscribe to on their phone, so they can match
|
||||
# without guesswork. The doubled-topic / wrong-host
|
||||
# mistakes are easier to spot when the actual URL
|
||||
# is right there in the success line.
|
||||
return {
|
||||
"ok": True,
|
||||
"message": (
|
||||
f"Sent to {full_url} — on your ntfy app, "
|
||||
f"subscribe to topic \"{topic}\" with server "
|
||||
f"\"{base}\" (or paste the full URL: {full_url})."
|
||||
),
|
||||
}
|
||||
return {"ok": False, "message": f"ntfy returned HTTP {r.status_code} from {full_url}: {r.text[:200]}"}
|
||||
except Exception as e:
|
||||
hint = ""
|
||||
if parsed.hostname not in ("127.0.0.1", "localhost"):
|
||||
hint = " If this is Docker Compose ntfy, set NTFY_BIND to that host/Tailscale IP and NTFY_BASE_URL to the same server URL in .env, then recreate ntfy."
|
||||
return {"ok": False, "message": f"ntfy publish to {full_url} failed: {e}.{hint}"[:500]}
|
||||
|
||||
if preset == "discord_webhook":
|
||||
import httpx
|
||||
webhook_url = (integ.get("base_url") or "").strip()
|
||||
if not webhook_url:
|
||||
return {"ok": False, "message": "No webhook URL set — paste the full Discord webhook URL into the Base URL field."}
|
||||
payload = {
|
||||
"embeds": [{
|
||||
"title": "Odysseus connectivity test",
|
||||
"description": "If you see this, your Discord Webhook integration is wired up correctly.",
|
||||
"color": 5793266,
|
||||
}]
|
||||
}
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=8.0) as client:
|
||||
r = await client.post(webhook_url, json=payload)
|
||||
if r.is_success:
|
||||
return {"ok": True, "message": "Test embed sent — check your Discord channel to confirm it arrived."}
|
||||
return {"ok": False, "message": f"Discord returned HTTP {r.status_code}: {r.text[:200]}"}
|
||||
except Exception as e:
|
||||
return {"ok": False, "message": f"Request failed: {e}"[:400]}
|
||||
|
||||
# All other presets: GET against a known health endpoint.
|
||||
# Fall back to detecting from name if preset is missing.
|
||||
health_paths = {
|
||||
"miniflux": "/v1/me",
|
||||
"gitea": "/api/v1/version",
|
||||
"linkding": "/api/tags/",
|
||||
"homeassistant": "/api/",
|
||||
"home assistant": "/api/",
|
||||
}
|
||||
path = health_paths.get(preset, "/")
|
||||
result = await execute_api_call(integration_id, "GET", path)
|
||||
if result.get("exit_code", 1) == 0:
|
||||
return {"ok": True, "message": "Connection successful"}
|
||||
return {"ok": False, "message": (result.get("error") or "Connection failed")[:300]}
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
|
|
|||
|
|
@ -1,193 +1,16 @@
|
|||
"""Shared OAuth/device-flow route scaffolding for provider setup."""
|
||||
"""Backward-compat shim — canonical location is routes/auth/device_flow.py.
|
||||
|
||||
from __future__ import annotations
|
||||
This module is replaced in ``sys.modules`` by the canonical module object so
|
||||
that ``import routes.device_flow``, ``from routes.device_flow import X``, and
|
||||
the ``monkeypatch.setattr(device_flow, "require_admin", ...)`` pattern in
|
||||
test_device_flow_routes.py all operate on the *same* object. Also keeps the
|
||||
external importers (routes/copilot_routes.py, routes/
|
||||
chatgpt_subscription_routes.py) working unchanged. Keeps existing import
|
||||
paths working after slice 2n (#4082/#4071).
|
||||
"""
|
||||
|
||||
import inspect
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Iterable, Mapping, Optional
|
||||
import sys as _sys
|
||||
|
||||
from fastapi import APIRouter, Form, HTTPException, Request
|
||||
from routes.auth import device_flow as _canonical # noqa: F401
|
||||
|
||||
from core.middleware import require_admin
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeviceFlowStart:
|
||||
"""Provider-specific start result consumed by the shared route wrapper."""
|
||||
|
||||
pending: Mapping[str, Any]
|
||||
response: Mapping[str, Any]
|
||||
interval: int = 5
|
||||
expires_in: int = 900
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class DeviceFlowPoll:
|
||||
"""Normalized provider poll outcome."""
|
||||
|
||||
status: str
|
||||
endpoint: Optional[Mapping[str, Any]] = None
|
||||
error: Optional[str] = None
|
||||
detail: Optional[str] = None
|
||||
interval: Optional[int] = None
|
||||
|
||||
@classmethod
|
||||
def pending(cls, detail: Optional[str] = None) -> "DeviceFlowPoll":
|
||||
return cls(status="pending", detail=detail)
|
||||
|
||||
@classmethod
|
||||
def slow_down(cls, interval: Optional[int] = None, detail: Optional[str] = None) -> "DeviceFlowPoll":
|
||||
return cls(status="slow_down", interval=interval, detail=detail)
|
||||
|
||||
@classmethod
|
||||
def authorized(cls, endpoint: Mapping[str, Any]) -> "DeviceFlowPoll":
|
||||
return cls(status="authorized", endpoint=endpoint)
|
||||
|
||||
@classmethod
|
||||
def failed(cls, error: str) -> "DeviceFlowPoll":
|
||||
return cls(status="failed", error=error)
|
||||
|
||||
|
||||
class PendingDeviceFlowStore:
|
||||
"""Thread-safe in-memory pending device-flow store.
|
||||
|
||||
Device codes and provider-side secrets stay inside this process. Each entry
|
||||
stores provider payload separately from poll metadata so provider callbacks
|
||||
only receive the fields they created.
|
||||
"""
|
||||
|
||||
def __init__(self, *, time_func: Callable[[], float] = time.time):
|
||||
self._pending: dict[str, dict[str, Any]] = {}
|
||||
self._lock = threading.Lock()
|
||||
self._time = time_func
|
||||
|
||||
def _now(self) -> float:
|
||||
return float(self._time())
|
||||
|
||||
def prune_expired(self) -> None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
for key in [k for k, v in self._pending.items() if v.get("expires_at", 0) < now]:
|
||||
self._pending.pop(key, None)
|
||||
|
||||
def add(self, payload: Mapping[str, Any], *, interval: int, expires_in: int) -> str:
|
||||
self.prune_expired()
|
||||
poll_id = uuid.uuid4().hex
|
||||
with self._lock:
|
||||
self._pending[poll_id] = {
|
||||
"payload": dict(payload),
|
||||
"interval": max(int(interval or 5), 1),
|
||||
"expires_at": self._now() + max(int(expires_in or 900), 1),
|
||||
"next_poll_at": 0.0,
|
||||
}
|
||||
return poll_id
|
||||
|
||||
def get_payload(self, poll_id: str) -> Optional[dict[str, Any]]:
|
||||
self.prune_expired()
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
if entry is None:
|
||||
return None
|
||||
return dict(entry.get("payload") or {})
|
||||
|
||||
def is_throttled(self, poll_id: str) -> bool:
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
return bool(entry and self._now() < float(entry.get("next_poll_at") or 0))
|
||||
|
||||
def schedule_next(self, poll_id: str) -> None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
if entry is not None:
|
||||
entry["next_poll_at"] = now + int(entry.get("interval") or 5)
|
||||
|
||||
def slow_down(self, poll_id: str, interval: Optional[int] = None) -> None:
|
||||
now = self._now()
|
||||
with self._lock:
|
||||
entry = self._pending.get(poll_id)
|
||||
if entry is not None:
|
||||
new_interval = int(interval or (int(entry.get("interval") or 5) + 5))
|
||||
entry["interval"] = max(new_interval, 1)
|
||||
entry["next_poll_at"] = now + entry["interval"]
|
||||
|
||||
def pop(self, poll_id: str) -> None:
|
||||
with self._lock:
|
||||
self._pending.pop(poll_id, None)
|
||||
|
||||
|
||||
async def _maybe_await(value: Any) -> Any:
|
||||
if inspect.isawaitable(value):
|
||||
return await value
|
||||
return value
|
||||
|
||||
|
||||
def _pending_response(detail: Optional[str] = None) -> dict[str, Any]:
|
||||
response: dict[str, Any] = {"status": "pending"}
|
||||
if detail:
|
||||
response["detail"] = detail
|
||||
return response
|
||||
|
||||
|
||||
def create_device_flow_router(
|
||||
*,
|
||||
prefix: str,
|
||||
tags: Iterable[str],
|
||||
store: PendingDeviceFlowStore,
|
||||
start_flow: Callable[[Request, Mapping[str, Any]], DeviceFlowStart],
|
||||
poll_flow: Callable[[Request, Mapping[str, Any]], DeviceFlowPoll],
|
||||
) -> APIRouter:
|
||||
"""Create standard `/device/start|poll|cancel` routes for a provider."""
|
||||
|
||||
router = APIRouter(prefix=prefix, tags=list(tags))
|
||||
|
||||
@router.post("/device/start")
|
||||
async def device_start(request: Request):
|
||||
require_admin(request)
|
||||
form = await request.form()
|
||||
start = await _maybe_await(start_flow(request, form))
|
||||
interval = int(start.interval or 5)
|
||||
expires_in = int(start.expires_in or 900)
|
||||
poll_id = store.add(start.pending, interval=interval, expires_in=expires_in)
|
||||
response = dict(start.response)
|
||||
response.update({"poll_id": poll_id, "interval": interval, "expires_in": expires_in})
|
||||
return response
|
||||
|
||||
@router.post("/device/poll")
|
||||
async def device_poll(request: Request, poll_id: str = Form(...)):
|
||||
require_admin(request)
|
||||
payload = store.get_payload(poll_id)
|
||||
if payload is None:
|
||||
raise HTTPException(404, "Unknown or expired login session")
|
||||
if store.is_throttled(poll_id):
|
||||
return {"status": "pending"}
|
||||
|
||||
try:
|
||||
outcome = await _maybe_await(poll_flow(request, payload))
|
||||
except Exception:
|
||||
store.pop(poll_id)
|
||||
raise
|
||||
|
||||
if outcome.status == "authorized":
|
||||
store.pop(poll_id)
|
||||
return {"status": "authorized", "endpoint": dict(outcome.endpoint or {})}
|
||||
if outcome.status == "failed":
|
||||
store.pop(poll_id)
|
||||
return {"status": "failed", "error": outcome.error or "denied"}
|
||||
if outcome.status == "slow_down":
|
||||
store.slow_down(poll_id, outcome.interval)
|
||||
return _pending_response(outcome.detail)
|
||||
|
||||
store.schedule_next(poll_id)
|
||||
return _pending_response(outcome.detail)
|
||||
|
||||
@router.post("/device/cancel")
|
||||
def device_cancel(request: Request, poll_id: str = Form(...)):
|
||||
require_admin(request)
|
||||
store.pop(poll_id)
|
||||
return {"status": "cancelled"}
|
||||
|
||||
return router
|
||||
_sys.modules[__name__] = _canonical
|
||||
|
|
|
|||
41
tests/test_auth_routes_shim.py
Normal file
41
tests/test_auth_routes_shim.py
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
"""Regression test for the auth route shims (slice 2n, #4082/#4071).
|
||||
|
||||
The backward-compat shims at ``routes/auth_routes.py``,
|
||||
``routes/api_token_routes.py``, and ``routes/device_flow.py`` use
|
||||
``sys.modules`` replacement so the legacy import paths and the canonical
|
||||
``routes.auth.*`` paths resolve to the *same* module objects. This is
|
||||
required because:
|
||||
|
||||
- ``test_auth_policy.py`` / ``test_auth_session_revocation.py`` do
|
||||
``sys.modules.pop("routes.auth_routes")`` + re-import
|
||||
- ``test_api_token_routes.py`` does ``monkeypatch.delitem(sys.modules,
|
||||
"routes.api_token_routes")`` + re-import
|
||||
- ``test_rename_user_owner_sync.py`` does ``import ... as ar`` +
|
||||
``monkeypatch.setattr(ar, "DEEP_RESEARCH_DIR", ...)``
|
||||
- ``test_device_flow_routes.py`` does ``setattr(device_flow,
|
||||
"require_admin", ...)``
|
||||
"""
|
||||
|
||||
import importlib
|
||||
|
||||
import routes.api_token_routes as _shim_api_token # noqa: F401
|
||||
import routes.auth_routes as _shim_auth # noqa: F401
|
||||
import routes.device_flow as _shim_device_flow # noqa: F401
|
||||
|
||||
|
||||
def test_legacy_and_canonical_auth_routes_are_same_object():
|
||||
legacy = importlib.import_module("routes.auth_routes")
|
||||
canonical = importlib.import_module("routes.auth.auth_routes")
|
||||
assert legacy is canonical
|
||||
|
||||
|
||||
def test_legacy_and_canonical_api_token_routes_are_same_object():
|
||||
legacy = importlib.import_module("routes.api_token_routes")
|
||||
canonical = importlib.import_module("routes.auth.api_token_routes")
|
||||
assert legacy is canonical
|
||||
|
||||
|
||||
def test_legacy_and_canonical_device_flow_are_same_object():
|
||||
legacy = importlib.import_module("routes.device_flow")
|
||||
canonical = importlib.import_module("routes.auth.device_flow")
|
||||
assert legacy is canonical
|
||||
Loading…
Add table
Add a link
Reference in a new issue