Harden auth with revocable sessions, cookies, and single-login default.

Issue short-lived access JWTs backed by AuthSession rows, HttpOnly cookies with refresh rotation, idle timeout, session management UI, WebCRT ownership caps, and optional Redis login rate limits; new logins revoke other sessions by default.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
oliver 2026-08-06 14:13:37 +08:00
parent ba33ab5c4f
commit 20c2fcd496
34 changed files with 1462 additions and 149 deletions

View file

@ -0,0 +1,40 @@
"""Add refresh token columns on auth_session.
Revision ID: 20260806_auth_refresh
Revises: 20260806_auth_session
Create Date: 2026-08-06
"""
from __future__ import annotations
from typing import Sequence, Union
from alembic import op
revision: str = "20260806_auth_refresh"
down_revision: Union[str, Sequence[str], None] = "20260806_auth_session"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
from netx_api.schema_patches import apply_auth_schema_patches
apply_auth_schema_patches(op.get_bind())
def downgrade() -> None:
bind = op.get_bind()
dialect = bind.dialect.name
if dialect == "postgresql":
op.execute("ALTER TABLE auth_session DROP COLUMN IF EXISTS refresh_expires_at")
op.execute("ALTER TABLE auth_session DROP COLUMN IF EXISTS refresh_token_hash")
else:
try:
op.drop_column("auth_session", "refresh_expires_at")
except Exception:
pass
try:
op.drop_column("auth_session", "refresh_token_hash")
except Exception:
pass

View file

@ -0,0 +1,35 @@
"""Create auth_session table for revocable JWT logins.
Revision ID: 20260806_auth_session
Revises: 20260802_legacy
Create Date: 2026-08-06
"""
from __future__ import annotations
from typing import Sequence, Union
from alembic import op
revision: str = "20260806_auth_session"
down_revision: Union[str, Sequence[str], None] = "20260802_legacy"
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
from netx_api.schema_patches import apply_auth_schema_patches
apply_auth_schema_patches(op.get_bind())
def downgrade() -> None:
bind = op.get_bind()
dialect = bind.dialect.name
if dialect == "postgresql":
op.execute("DROP TABLE IF EXISTS auth_session")
else:
try:
op.drop_table("auth_session")
except Exception:
pass

View file

@ -43,3 +43,5 @@ If Alembic history was never stamped and you prefer to mark current without re-r
|----------|---------| |----------|---------|
| `20260802_scopes` | `app_user.scopes` / `api_token.scopes` | | `20260802_scopes` | `app_user.scopes` / `api_token.scopes` |
| `20260802_legacy` | Shared brownfield patches (alarms, inventory, managed_ne, topology, port traffic, key-alert, …) | | `20260802_legacy` | Shared brownfield patches (alarms, inventory, managed_ne, topology, port traffic, key-alert, …) |
| `20260806_auth_session` | Revocable JWT login sessions (`auth_session`) |
| `20260806_auth_refresh` | Refresh token columns on `auth_session` |

94
netx_api/auth_cookies.py Normal file
View file

@ -0,0 +1,94 @@
"""Auth cookie helpers (HttpOnly access + refresh)."""
from __future__ import annotations
from typing import Any
from fastapi import Response
from starlette.requests import Request
from .config import settings
COOKIE_ACCESS = "netx_at"
COOKIE_REFRESH = "netx_rt"
def _cookie_secure(request: Request | None = None) -> bool:
configured = getattr(settings, "auth_cookie_secure", None)
if configured is True:
return True
if configured is False:
return False
# Auto: secure when request is HTTPS (or behind TLS-terminating proxy).
if request is None:
return False
if request.url.scheme == "https":
return True
fwd = str(request.headers.get("x-forwarded-proto") or "").split(",")[0].strip().lower()
return fwd == "https"
def _cookie_samesite() -> str:
raw = str(getattr(settings, "auth_cookie_samesite", "lax") or "lax").strip().lower()
if raw in ("lax", "strict", "none"):
return raw
return "lax"
def access_cookie_max_age() -> int:
return max(300, int(getattr(settings, "auth_token_ttl_sec", 3600) or 3600))
def refresh_cookie_max_age() -> int:
return max(3600, int(getattr(settings, "auth_refresh_ttl_sec", 604800) or 604800))
def set_auth_cookies(
response: Response,
*,
access_token: str,
refresh_token: str,
request: Request | None = None,
) -> None:
if not bool(getattr(settings, "auth_cookie_enabled", True)):
return
secure = _cookie_secure(request)
samesite = _cookie_samesite()
# SameSite=None requires Secure.
if samesite == "none":
secure = True
common: dict[str, Any] = {
"httponly": True,
"secure": secure,
"samesite": samesite,
"path": "/",
}
response.set_cookie(
COOKIE_ACCESS,
access_token,
max_age=access_cookie_max_age(),
**common,
)
response.set_cookie(
COOKIE_REFRESH,
refresh_token,
max_age=refresh_cookie_max_age(),
**common,
)
def clear_auth_cookies(response: Response, *, request: Request | None = None) -> None:
secure = _cookie_secure(request)
samesite = _cookie_samesite()
if samesite == "none":
secure = True
for name in (COOKIE_ACCESS, COOKIE_REFRESH):
response.delete_cookie(name, path="/", secure=secure, httponly=True, samesite=samesite)
def read_access_cookie(request: Request) -> str:
return str(request.cookies.get(COOKIE_ACCESS) or "").strip()
def read_refresh_cookie(request: Request) -> str:
return str(request.cookies.get(COOKIE_REFRESH) or "").strip()

View file

@ -8,6 +8,7 @@ from typing import Annotated, Callable
from fastapi import Depends, HTTPException, Request from fastapi import Depends, HTTPException, Request
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .auth_cookies import read_access_cookie
from .auth_scopes import ( from .auth_scopes import (
ALL_SCOPES, ALL_SCOPES,
effective_token_scopes, effective_token_scopes,
@ -16,7 +17,7 @@ from .auth_scopes import (
has_scope, has_scope,
normalize_scopes, normalize_scopes,
) )
from .auth_service import get_user_by_id, resolve_api_token_row from .auth_service import get_auth_session, get_user_by_id, resolve_api_token_row, touch_auth_session
from .auth_tokens import decode_access_token from .auth_tokens import decode_access_token
from .config import settings from .config import settings
from .db import get_db from .db import get_db
@ -29,23 +30,25 @@ class AuthContext:
auth_via: str # jwt | api_token | disabled auth_via: str # jwt | api_token | disabled
scopes: frozenset[str] = field(default_factory=frozenset) scopes: frozenset[str] = field(default_factory=frozenset)
api_token_id: str = "" api_token_id: str = ""
session_jti: str = ""
def _extract_bearer(request: Request) -> str: def _extract_bearer(request: Request) -> str:
auth = str(request.headers.get("authorization") or "").strip() auth = str(request.headers.get("authorization") or "").strip()
if auth.lower().startswith("bearer "): if auth.lower().startswith("bearer "):
return auth[7:].strip() return auth[7:].strip()
# Prefer Header; query access_token is deprecated (WebSocket may still use short-lived tickets). # Prefer Authorization; fall back to HttpOnly access cookie for browser sessions.
q = request.query_params.get("access_token") return read_access_cookie(request)
return str(q or "").strip()
def user_scopes(user: AppUser) -> frozenset[str]: def user_scopes(user: AppUser) -> frozenset[str]:
return effective_user_scopes(role=str(user.role or "user"), override=getattr(user, "scopes", None) or []) return effective_user_scopes(role=str(user.role or "user"), override=getattr(user, "scopes", None) or [])
def resolve_user_from_token(db: Session, token: str) -> tuple[AppUser, str, frozenset[str], str] | None: def resolve_user_from_token(
"""Return (user, via, scopes, api_token_id) or None.""" db: Session, token: str
) -> tuple[AppUser, str, frozenset[str], str, str] | None:
"""Return (user, via, scopes, api_token_id, session_jti) or None."""
raw = str(token or "").strip() raw = str(token or "").strip()
if not raw: if not raw:
return None return None
@ -60,17 +63,26 @@ def resolve_user_from_token(db: Session, token: str) -> tuple[AppUser, str, froz
user_scopes=user_scopes(user), user_scopes=user_scopes(user),
token_scopes=getattr(row, "scopes", None) or [], token_scopes=getattr(row, "scopes", None) or [],
) )
return user, "api_token", scopes, str(row.id) return user, "api_token", scopes, str(row.id), ""
try: try:
payload = decode_access_token(raw) payload = decode_access_token(raw)
except Exception: except Exception:
return None return None
if str(payload.get("typ") or "") not in ("", "access"): if str(payload.get("typ") or "") not in ("", "access"):
return None return None
jti = str(payload.get("jti") or "").strip()
if not jti:
return None
sess = get_auth_session(db, jti)
if sess is None:
return None
user = get_user_by_id(db, str(payload.get("sub") or "")) user = get_user_by_id(db, str(payload.get("sub") or ""))
if user is None or not user.is_active: if user is None or not user.is_active:
return None return None
return user, "jwt", user_scopes(user), "" if str(sess.user_id) != str(user.id):
return None
touch_auth_session(db, jti)
return user, "jwt", user_scopes(user), "", jti
def get_optional_user( def get_optional_user(
@ -86,19 +98,25 @@ def get_optional_user(
if not isinstance(scopes, frozenset): if not isinstance(scopes, frozenset):
scopes = user_scopes(cached) scopes = user_scopes(cached)
token_id = str(getattr(request.state, "auth_api_token_id", "") or "") token_id = str(getattr(request.state, "auth_api_token_id", "") or "")
return AuthContext(user=cached, auth_via=via, scopes=scopes, api_token_id=token_id) jti = str(getattr(request.state, "auth_session_jti", "") or "")
return AuthContext(
user=cached, auth_via=via, scopes=scopes, api_token_id=token_id, session_jti=jti
)
token = _extract_bearer(request) token = _extract_bearer(request)
if not token: if not token:
return None return None
resolved = resolve_user_from_token(db, token) resolved = resolve_user_from_token(db, token)
if resolved is None: if resolved is None:
return None return None
user, via, scopes, token_id = resolved user, via, scopes, token_id, jti = resolved
request.state.auth_user = user request.state.auth_user = user
request.state.auth_via = via request.state.auth_via = via
request.state.auth_scopes = scopes request.state.auth_scopes = scopes
request.state.auth_api_token_id = token_id request.state.auth_api_token_id = token_id
return AuthContext(user=user, auth_via=via, scopes=scopes, api_token_id=token_id) request.state.auth_session_jti = jti
return AuthContext(
user=user, auth_via=via, scopes=scopes, api_token_id=token_id, session_jti=jti
)
def require_user( def require_user(

View file

@ -29,12 +29,23 @@ _PUBLIC_EXACT = frozenset(
"/v1/metrics/json", "/v1/metrics/json",
"/favicon.ico", "/favicon.ico",
"/v1/auth/login", "/v1/auth/login",
"/v1/auth/refresh",
} }
) )
_PUBLIC_PREFIXES = ( _PUBLIC_PREFIXES = (
"/assets", "/assets",
) )
# While must_change_password is true, only these authenticated endpoints are allowed.
_PASSWORD_CHANGE_ALLOW = frozenset(
{
("GET", "/v1/auth/me"),
("POST", "/v1/auth/change-password"),
("POST", "/v1/auth/logout"),
("GET", "/v1/auth/sessions"),
}
)
def _docs_public() -> bool: def _docs_public() -> bool:
return bool(getattr(settings, "docs_enabled", False)) return bool(getattr(settings, "docs_enabled", False))
@ -92,10 +103,11 @@ class AuthAuditMiddleware(BaseHTTPMiddleware):
auth = str(request.headers.get("authorization") or "").strip() auth = str(request.headers.get("authorization") or "").strip()
if auth.lower().startswith("bearer "): if auth.lower().startswith("bearer "):
token = auth[7:].strip() token = auth[7:].strip()
# Query access_token: only allow for non-webcrt paths as deprecated fallback; if not token:
# WebCRT HTTP must use Authorization header (see webcrt_router). from .auth_cookies import read_access_cookie
if not token and not path.startswith("/v1/webcrt"):
token = str(request.query_params.get("access_token") or "").strip() token = read_access_cookie(request)
# Query-string access_token is no longer accepted (leaks via logs/proxies).
db = SessionLocal() db = SessionLocal()
try: try:
@ -112,7 +124,26 @@ class AuthAuditMiddleware(BaseHTTPMiddleware):
detail={}, detail={},
) )
return JSONResponse(status_code=401, content={"detail": "unauthorized"}) return JSONResponse(status_code=401, content={"detail": "unauthorized"})
user, via, scopes, token_id = resolved user, via, scopes, token_id, jti = resolved
if bool(getattr(user, "must_change_password", False)):
allow_key = (request.method.upper(), path.rstrip("/") if len(path) > 1 else path)
if allow_key not in _PASSWORD_CHANGE_ALLOW:
write_audit(
db,
action="auth.password_change_required",
actor_user_id=str(user.id),
actor_username=str(user.username),
method=request.method,
path=path,
status_code=403,
client_ip=_client_ip(request),
user_agent=str(request.headers.get("user-agent") or "")[:512],
detail={"auth_via": via},
)
return JSONResponse(
status_code=403,
content={"detail": "password_change_required"},
)
need = required_scope_for_request(request.method, path) need = required_scope_for_request(request.method, path)
if need and not has_scope(scopes, need): if need and not has_scope(scopes, need):
write_audit( write_audit(
@ -141,6 +172,7 @@ class AuthAuditMiddleware(BaseHTTPMiddleware):
request.state.auth_via = via request.state.auth_via = via
request.state.auth_scopes = scopes request.state.auth_scopes = scopes
request.state.auth_api_token_id = token_id request.state.auth_api_token_id = token_id
request.state.auth_session_jti = jti
actor_id = str(user.id) actor_id = str(user.id)
actor_name = str(user.username) actor_name = str(user.username)
auth_via = via auth_via = via

139
netx_api/auth_rate_limit.py Normal file
View file

@ -0,0 +1,139 @@
"""Login brute-force protection: Redis when configured, else in-process."""
from __future__ import annotations
import logging
import threading
import time
from dataclasses import dataclass
from typing import Any
from .config import settings
_log = logging.getLogger("netx.auth.rate")
_lock = threading.Lock()
@dataclass
class _Bucket:
window_start: float
failures: int
locked_until: float = 0.0
_buckets: dict[str, _Bucket] = {}
_redis_client: Any | None = None
_redis_failed = False
def _key(username: str, client_ip: str) -> str:
return f"{str(client_ip or '').strip().lower()}|{str(username or '').strip().lower()}"
def _limits() -> tuple[int, int, int]:
max_fail = max(3, int(getattr(settings, "auth_login_max_failures", 10) or 10))
window = max(60, int(getattr(settings, "auth_login_window_sec", 300) or 300))
lockout = max(60, int(getattr(settings, "auth_login_lockout_sec", 900) or 900))
return max_fail, window, lockout
def _get_redis() -> Any | None:
global _redis_client, _redis_failed
url = str(getattr(settings, "auth_redis_url", "") or "").strip()
if not url:
return None
if _redis_failed:
return None
if _redis_client is not None:
return _redis_client
try:
import redis # type: ignore
client = redis.Redis.from_url(url, decode_responses=True, socket_connect_timeout=0.5)
client.ping()
_redis_client = client
_log.info("auth login rate-limit using Redis")
return _redis_client
except Exception:
_redis_failed = True
_log.warning("auth Redis unavailable; falling back to in-process rate-limit", exc_info=True)
return None
def login_lock_remaining(username: str, client_ip: str) -> float:
"""Seconds remaining on lockout, or 0 if not locked."""
k = _key(username, client_ip)
r = _get_redis()
if r is not None:
try:
ttl = r.ttl(f"netx:auth:lock:{k}")
if isinstance(ttl, int) and ttl > 0:
return float(ttl)
return 0.0
except Exception:
_log.debug("redis lock_remaining failed", exc_info=True)
now = time.time()
with _lock:
row = _buckets.get(k)
if row is None:
return 0.0
return max(0.0, float(row.locked_until) - now)
def register_login_failure(username: str, client_ip: str) -> float:
"""Record a failed login. Returns lock remaining seconds (0 if not yet locked)."""
max_fail, window, lockout = _limits()
k = _key(username, client_ip)
r = _get_redis()
if r is not None:
try:
lock_key = f"netx:auth:lock:{k}"
fail_key = f"netx:auth:fail:{k}"
existing = r.ttl(lock_key)
if isinstance(existing, int) and existing > 0:
return float(existing)
count = int(r.incr(fail_key))
if count == 1:
r.expire(fail_key, window)
if count >= max_fail:
r.setex(lock_key, lockout, "1")
r.delete(fail_key)
return float(lockout)
return 0.0
except Exception:
_log.debug("redis register_failure failed; using memory", exc_info=True)
now = time.time()
with _lock:
row = _buckets.get(k)
if row is None or (now - row.window_start) > window:
row = _Bucket(window_start=now, failures=0)
_buckets[k] = row
if row.locked_until > now:
return float(row.locked_until - now)
row.failures += 1
if row.failures >= max_fail:
row.locked_until = now + lockout
return float(lockout)
return 0.0
def clear_login_failures(username: str, client_ip: str) -> None:
k = _key(username, client_ip)
r = _get_redis()
if r is not None:
try:
r.delete(f"netx:auth:fail:{k}", f"netx:auth:lock:{k}")
except Exception:
_log.debug("redis clear_failures failed", exc_info=True)
with _lock:
_buckets.pop(k, None)
def reset_login_rate_limit_for_tests() -> None:
global _redis_client, _redis_failed
with _lock:
_buckets.clear()
_redis_client = None
_redis_failed = False

View file

@ -4,15 +4,17 @@ from __future__ import annotations
from typing import Annotated, Any from typing import Annotated, Any
from fastapi import APIRouter, Depends, Query, Request from fastapi import APIRouter, Depends, Query, Request, Response
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
from .auth_cookies import clear_auth_cookies, read_refresh_cookie, set_auth_cookies
from .auth_deps import AuthContext, require_admin, require_user from .auth_deps import AuthContext, require_admin, require_user
from .auth_schemas import ( from .auth_schemas import (
ApiTokenCreateRequest, ApiTokenCreateRequest,
ApiTokenUpdateRequest, ApiTokenUpdateRequest,
ChangePasswordRequest, ChangePasswordRequest,
LoginRequest, LoginRequest,
RefreshRequest,
UserCreateRequest, UserCreateRequest,
UserUpdateRequest, UserUpdateRequest,
) )
@ -23,14 +25,23 @@ from .auth_service import (
create_user, create_user,
list_api_tokens, list_api_tokens,
list_audit_logs, list_audit_logs,
list_auth_sessions,
list_users, list_users,
login_issue_token, login_issue_token,
refresh_login_tokens,
revoke_api_token, revoke_api_token,
revoke_auth_session_for_user,
revoke_auth_sessions,
update_api_token, update_api_token,
update_user, update_user,
user_public, user_public,
write_audit, write_audit,
) )
from .auth_rate_limit import (
clear_login_failures,
login_lock_remaining,
register_login_failure,
)
from .db import get_db from .db import get_db
router = APIRouter(tags=["auth"]) router = APIRouter(tags=["auth"])
@ -42,11 +53,45 @@ def _client_meta(request: Request) -> tuple[str, str]:
return ip, ua return ip, ua
def _token_response(response: Response, request: Request, out: dict[str, Any]) -> dict[str, Any]:
access = str(out.get("access_token") or "")
refresh = str(out.get("refresh_token") or "")
if access and refresh:
set_auth_cookies(response, access_token=access, refresh_token=refresh, request=request)
# Keep tokens in JSON for API clients / scripts; browsers rely on HttpOnly cookies.
return out
@router.post("/v1/auth/login") @router.post("/v1/auth/login")
def api_login(body: LoginRequest, request: Request, db: Session = Depends(get_db)) -> dict[str, Any]: def api_login(
body: LoginRequest,
request: Request,
response: Response,
db: Session = Depends(get_db),
) -> dict[str, Any]:
from fastapi import HTTPException
ip, ua = _client_meta(request) ip, ua = _client_meta(request)
locked = login_lock_remaining(body.username, ip)
if locked > 0:
write_audit(
db,
action="auth.login_locked",
actor_username=str(body.username or "").strip(),
method="POST",
path="/v1/auth/login",
status_code=429,
client_ip=ip,
user_agent=ua,
detail={"retry_after_sec": int(locked)},
)
raise HTTPException(
status_code=429,
detail={"error": "login_locked", "retry_after_sec": int(locked)},
)
user = authenticate_user(db, body.username, body.password) user = authenticate_user(db, body.username, body.password)
if user is None: if user is None:
remaining = register_login_failure(body.username, ip)
write_audit( write_audit(
db, db,
action="auth.login_failed", action="auth.login_failed",
@ -56,12 +101,16 @@ def api_login(body: LoginRequest, request: Request, db: Session = Depends(get_db
status_code=401, status_code=401,
client_ip=ip, client_ip=ip,
user_agent=ua, user_agent=ua,
detail={}, detail={"locked": remaining > 0, "retry_after_sec": int(remaining)},
)
if remaining > 0:
raise HTTPException(
status_code=429,
detail={"error": "login_locked", "retry_after_sec": int(remaining)},
) )
from fastapi import HTTPException
raise HTTPException(status_code=401, detail="invalid_credentials") raise HTTPException(status_code=401, detail="invalid_credentials")
out = login_issue_token(user) clear_login_failures(body.username, ip)
out = login_issue_token(db, user, client_ip=ip, user_agent=ua)
write_audit( write_audit(
db, db,
action="auth.login", action="auth.login",
@ -74,16 +123,71 @@ def api_login(body: LoginRequest, request: Request, db: Session = Depends(get_db
user_agent=ua, user_agent=ua,
detail={"role": user.role}, detail={"role": user.role},
) )
return out return _token_response(response, request, out)
@router.post("/v1/auth/refresh")
def api_refresh(
body: RefreshRequest,
request: Request,
response: Response,
db: Session = Depends(get_db),
) -> dict[str, Any]:
from fastapi import HTTPException
ip, ua = _client_meta(request)
refresh = str(body.refresh_token or "").strip() or read_refresh_cookie(request)
if not refresh:
raise HTTPException(status_code=401, detail="invalid_refresh_token")
try:
out = refresh_login_tokens(db, refresh_token=refresh, client_ip=ip, user_agent=ua)
except Exception:
write_audit(
db,
action="auth.refresh_failed",
method="POST",
path="/v1/auth/refresh",
status_code=401,
client_ip=ip,
user_agent=ua,
detail={},
)
raise
user = out.get("user") or {}
write_audit(
db,
action="auth.refresh",
actor_user_id=str(user.get("id") or ""),
actor_username=str(user.get("username") or ""),
method="POST",
path="/v1/auth/refresh",
status_code=200,
client_ip=ip,
user_agent=ua,
detail={},
)
return _token_response(response, request, out)
@router.post("/v1/auth/logout") @router.post("/v1/auth/logout")
def api_logout( def api_logout(
request: Request, request: Request,
response: Response,
ctx: Annotated[AuthContext, Depends(require_user)], ctx: Annotated[AuthContext, Depends(require_user)],
db: Session = Depends(get_db), db: Session = Depends(get_db),
) -> dict[str, Any]: ) -> dict[str, Any]:
ip, ua = _client_meta(request) ip, ua = _client_meta(request)
revoked = 0
if ctx.auth_via == "jwt" and ctx.session_jti:
revoked = revoke_auth_sessions(db, user_id=str(ctx.user.id), only_jti=ctx.session_jti)
closed_webcrt = 0
try:
from .webcrt_session_registry import close_sessions_for_user
closed_webcrt = close_sessions_for_user(str(ctx.user.id), reason="auth_logout")
except Exception:
pass
clear_auth_cookies(response, request=request)
write_audit( write_audit(
db, db,
action="auth.logout", action="auth.logout",
@ -94,9 +198,75 @@ def api_logout(
status_code=200, status_code=200,
client_ip=ip, client_ip=ip,
user_agent=ua, user_agent=ua,
detail={"auth_via": ctx.auth_via}, detail={"auth_via": ctx.auth_via, "revoked": revoked, "webcrt_closed": closed_webcrt},
) )
return {"ok": True} return {"ok": True, "revoked": revoked, "webcrt_closed": closed_webcrt}
@router.get("/v1/auth/sessions")
def api_list_sessions(ctx: Annotated[AuthContext, Depends(require_user)], db: Session = Depends(get_db)) -> dict[str, Any]:
items = list_auth_sessions(db, user_id=str(ctx.user.id), current_jti=ctx.session_jti)
return {"items": items, "total": len(items)}
@router.delete("/v1/auth/sessions/{session_id}")
def api_revoke_session(
session_id: str,
request: Request,
ctx: Annotated[AuthContext, Depends(require_user)],
db: Session = Depends(get_db),
) -> dict[str, Any]:
from fastapi import HTTPException
ok = revoke_auth_session_for_user(
db,
user_id=str(ctx.user.id),
session_id=session_id,
current_jti=ctx.session_jti,
)
if not ok:
raise HTTPException(status_code=404, detail="session_not_found")
ip, ua = _client_meta(request)
write_audit(
db,
action="auth.session_revoke",
actor_user_id=ctx.user.id,
actor_username=ctx.user.username,
method="DELETE",
path=f"/v1/auth/sessions/{session_id}",
status_code=200,
client_ip=ip,
user_agent=ua,
detail={"session_id": session_id, "current": session_id == ctx.session_jti},
)
return {"ok": True, "revoked": True, "session_id": session_id}
@router.post("/v1/auth/sessions/revoke-others")
def api_revoke_other_sessions(
request: Request,
ctx: Annotated[AuthContext, Depends(require_user)],
db: Session = Depends(get_db),
) -> dict[str, Any]:
if not ctx.session_jti:
# API-token auth has no JWT session to keep.
n = revoke_auth_sessions(db, user_id=str(ctx.user.id))
else:
n = revoke_auth_sessions(db, user_id=str(ctx.user.id), except_jti=ctx.session_jti)
ip, ua = _client_meta(request)
write_audit(
db,
action="auth.session_revoke_others",
actor_user_id=ctx.user.id,
actor_username=ctx.user.username,
method="POST",
path="/v1/auth/sessions/revoke-others",
status_code=200,
client_ip=ip,
user_agent=ua,
detail={"revoked": n},
)
return {"ok": True, "revoked": n}
@router.get("/v1/auth/me") @router.get("/v1/auth/me")
@ -118,7 +288,13 @@ def api_change_password(
ctx: Annotated[AuthContext, Depends(require_user)], ctx: Annotated[AuthContext, Depends(require_user)],
db: Session = Depends(get_db), db: Session = Depends(get_db),
) -> dict[str, Any]: ) -> dict[str, Any]:
change_password(db, user=ctx.user, old_password=body.old_password, new_password=body.new_password) change_password(
db,
user=ctx.user,
old_password=body.old_password,
new_password=body.new_password,
keep_jti=ctx.session_jti or None,
)
ip, ua = _client_meta(request) ip, ua = _client_meta(request)
write_audit( write_audit(
db, db,

View file

@ -12,12 +12,17 @@ class LoginRequest(BaseModel):
class ChangePasswordRequest(BaseModel): class ChangePasswordRequest(BaseModel):
old_password: str = Field(min_length=1, max_length=256) old_password: str = Field(min_length=1, max_length=256)
new_password: str = Field(min_length=6, max_length=256) new_password: str = Field(min_length=8, max_length=256)
class RefreshRequest(BaseModel):
# Optional when refresh token is sent via HttpOnly cookie.
refresh_token: str = Field(default="", max_length=512)
class UserCreateRequest(BaseModel): class UserCreateRequest(BaseModel):
username: str = Field(min_length=2, max_length=64) username: str = Field(min_length=2, max_length=64)
password: str = Field(min_length=6, max_length=256) password: str = Field(min_length=8, max_length=256)
role: str = Field(default="user") role: str = Field(default="user")
scopes: list[str] | None = None scopes: list[str] | None = None
@ -25,7 +30,7 @@ class UserCreateRequest(BaseModel):
class UserUpdateRequest(BaseModel): class UserUpdateRequest(BaseModel):
is_active: bool | None = None is_active: bool | None = None
role: str | None = None role: str | None = None
password: str | None = Field(default=None, min_length=6, max_length=256) password: str | None = Field(default=None, min_length=8, max_length=256)
scopes: list[str] | None = None scopes: list[str] | None = None

View file

@ -18,9 +18,14 @@ from .auth_scopes import (
effective_user_scopes, effective_user_scopes,
normalize_scopes, normalize_scopes,
) )
from .auth_tokens import hash_api_token, issue_access_token, new_api_token_plaintext from .auth_tokens import (
hash_api_token,
issue_access_token,
new_api_token_plaintext,
new_refresh_token_plaintext,
)
from .config import settings from .config import settings
from .models import ApiToken, AppUser, AuditLog from .models import ApiToken, AppUser, AuditLog, AuthSession
from .timeutil import utcnow_naive from .timeutil import utcnow_naive
_log = logging.getLogger("netx.auth") _log = logging.getLogger("netx.auth")
@ -33,6 +38,7 @@ _SECRET_KEYS = frozenset(
"hop_password", "hop_password",
"enable_secret", "enable_secret",
"access_token", "access_token",
"refresh_token",
"token", "token",
"authorization", "authorization",
"secret", "secret",
@ -41,6 +47,17 @@ _SECRET_KEYS = frozenset(
) )
def _password_min_len() -> int:
return max(8, int(getattr(settings, "auth_password_min_len", 8) or 8))
def _require_password_strength(pwd: str) -> str:
raw = str(pwd or "")
if len(raw) < _password_min_len():
raise HTTPException(status_code=400, detail="password_too_short")
return raw
def user_public(user: AppUser) -> dict[str, Any]: def user_public(user: AppUser) -> dict[str, Any]:
scopes = sorted( scopes = sorted(
effective_user_scopes(role=str(user.role or "user"), override=getattr(user, "scopes", None) or []) effective_user_scopes(role=str(user.role or "user"), override=getattr(user, "scopes", None) or [])
@ -256,15 +273,201 @@ def authenticate_user(db: Session, username: str, password: str) -> AppUser | No
return user return user
def login_issue_token(user: AppUser) -> dict[str, Any]: def revoke_auth_sessions(
token = issue_access_token(user_id=user.id, username=user.username, role=user.role) db: Session,
*,
user_id: str,
except_jti: str | None = None,
only_jti: str | None = None,
) -> int:
"""Revoke JWT sessions. Returns count newly revoked."""
now = utcnow_naive()
q = db.query(AuthSession).filter(
AuthSession.user_id == str(user_id),
AuthSession.revoked_at.is_(None),
)
if only_jti:
q = q.filter(AuthSession.id == str(only_jti))
if except_jti:
q = q.filter(AuthSession.id != str(except_jti))
rows = q.all()
for row in rows:
row.revoked_at = now
if rows:
db.commit()
return len(rows)
def get_auth_session(db: Session, jti: str) -> AuthSession | None:
sid = str(jti or "").strip()
if not sid:
return None
row = db.query(AuthSession).filter(AuthSession.id == sid).first()
if row is None:
return None
if row.revoked_at is not None:
return None
now = utcnow_naive()
exp = row.expires_at
if exp is not None and exp < now:
return None
idle = max(0, int(getattr(settings, "auth_idle_timeout_sec", 7200) or 0))
if idle > 0:
seen = row.last_seen_at or row.created_at
if seen is not None and (now - seen).total_seconds() > idle:
row.revoked_at = now
try:
db.commit()
except Exception:
db.rollback()
return None
return row
def touch_auth_session(db: Session, jti: str) -> None:
row = db.query(AuthSession).filter(AuthSession.id == str(jti)).first()
if row is None or row.revoked_at is not None:
return
row.last_seen_at = utcnow_naive()
try:
db.commit()
except Exception:
db.rollback()
def list_auth_sessions(db: Session, *, user_id: str, current_jti: str = "") -> list[dict[str, Any]]:
now = utcnow_naive()
rows = (
db.query(AuthSession)
.filter(AuthSession.user_id == str(user_id), AuthSession.revoked_at.is_(None))
.order_by(AuthSession.created_at.desc())
.limit(100)
.all()
)
out: list[dict[str, Any]] = []
for row in rows:
if row.expires_at is not None and row.expires_at < now:
continue
refresh_alive = bool(row.refresh_expires_at and row.refresh_expires_at >= now)
access_alive = bool(row.expires_at and row.expires_at >= now)
if not access_alive and not refresh_alive:
continue
out.append(
{
"id": row.id,
"created_at": row.created_at.isoformat() if row.created_at else None,
"expires_at": row.expires_at.isoformat() if row.expires_at else None,
"refresh_expires_at": row.refresh_expires_at.isoformat() if row.refresh_expires_at else None,
"last_seen_at": row.last_seen_at.isoformat() if row.last_seen_at else None,
"client_ip": row.client_ip or "",
"user_agent": (row.user_agent or "")[:200],
"current": bool(current_jti) and row.id == str(current_jti),
}
)
return out
def revoke_auth_session_for_user(
db: Session,
*,
user_id: str,
session_id: str,
current_jti: str = "",
) -> bool:
"""Revoke one session owned by user. Returns True if newly revoked."""
sid = str(session_id or "").strip()
if not sid:
return False
row = (
db.query(AuthSession)
.filter(
AuthSession.id == sid,
AuthSession.user_id == str(user_id),
AuthSession.revoked_at.is_(None),
)
.first()
)
if row is None:
return False
row.revoked_at = utcnow_naive()
db.commit()
return True
def login_issue_token(
db: Session,
user: AppUser,
*,
client_ip: str = "",
user_agent: str = "",
) -> dict[str, Any]:
token, jti, ttl = issue_access_token(
user_id=user.id, username=user.username, role=user.role
)
if bool(getattr(settings, "auth_single_session", False)):
revoke_auth_sessions(db, user_id=str(user.id), except_jti=jti)
refresh_ttl = max(3600, int(getattr(settings, "auth_refresh_ttl_sec", 604800) or 604800))
refresh_plain = new_refresh_token_plaintext()
now = utcnow_naive()
db.add(
AuthSession(
id=jti,
user_id=str(user.id),
created_at=now,
expires_at=now + timedelta(seconds=ttl),
client_ip=str(client_ip or "")[:128],
user_agent=str(user_agent or "")[:512],
last_seen_at=now,
refresh_token_hash=hash_api_token(refresh_plain),
refresh_expires_at=now + timedelta(seconds=refresh_ttl),
)
)
db.commit()
return { return {
"access_token": token, "access_token": token,
"refresh_token": refresh_plain,
"token_type": "bearer", "token_type": "bearer",
"expires_in": ttl,
"refresh_expires_in": refresh_ttl,
"user": user_public(user), "user": user_public(user),
} }
def refresh_login_tokens(
db: Session,
*,
refresh_token: str,
client_ip: str = "",
user_agent: str = "",
) -> dict[str, Any]:
"""Rotate refresh token and mint a new access JWT (old session revoked)."""
raw = str(refresh_token or "").strip()
if not raw.startswith("nxr_"):
raise HTTPException(status_code=401, detail="invalid_refresh_token")
th = hash_api_token(raw)
now = utcnow_naive()
row = (
db.query(AuthSession)
.filter(AuthSession.refresh_token_hash == th, AuthSession.revoked_at.is_(None))
.first()
)
if row is None:
raise HTTPException(status_code=401, detail="invalid_refresh_token")
refresh_exp = row.refresh_expires_at
if refresh_exp is None or refresh_exp < now:
row.revoked_at = now
db.commit()
raise HTTPException(status_code=401, detail="refresh_token_expired")
user = get_user_by_id(db, str(row.user_id))
if user is None or not user.is_active:
row.revoked_at = now
db.commit()
raise HTTPException(status_code=401, detail="invalid_refresh_token")
row.revoked_at = now
db.commit()
return login_issue_token(db, user, client_ip=client_ip, user_agent=user_agent)
def list_users(db: Session) -> list[dict[str, Any]]: def list_users(db: Session) -> list[dict[str, Any]]:
rows = db.query(AppUser).order_by(AppUser.created_at.asc()).all() rows = db.query(AppUser).order_by(AppUser.created_at.asc()).all()
return [user_public(u) for u in rows] return [user_public(u) for u in rows]
@ -282,9 +485,7 @@ def create_user(
name = str(username or "").strip() name = str(username or "").strip()
if not _USERNAME_RE.match(name): if not _USERNAME_RE.match(name):
raise HTTPException(status_code=400, detail="invalid_username") raise HTTPException(status_code=400, detail="invalid_username")
pwd = str(password or "") pwd = _require_password_strength(password)
if len(pwd) < 6:
raise HTTPException(status_code=400, detail="password_too_short")
role_n = str(role or "user").strip().lower() role_n = str(role or "user").strip().lower()
if role_n not in ("admin", "user"): if role_n not in ("admin", "user"):
raise HTTPException(status_code=400, detail="invalid_role") raise HTTPException(status_code=400, detail="invalid_role")
@ -327,29 +528,38 @@ def update_user(
if user.id == actor.id and role_n != "admin": if user.id == actor.id and role_n != "admin":
raise HTTPException(status_code=400, detail="cannot_demote_self") raise HTTPException(status_code=400, detail="cannot_demote_self")
user.role = role_n user.role = role_n
revoke_all = False
if is_active is not None: if is_active is not None:
user.is_active = bool(is_active) user.is_active = bool(is_active)
if not user.is_active:
revoke_all = True
if password is not None: if password is not None:
pwd = str(password) pwd = _require_password_strength(password)
if len(pwd) < 6:
raise HTTPException(status_code=400, detail="password_too_short")
user.password_hash = hash_password(pwd) user.password_hash = hash_password(pwd)
user.must_change_password = True user.must_change_password = True
revoke_all = True
if scopes is not None: if scopes is not None:
user.scopes = normalize_scopes(scopes) user.scopes = normalize_scopes(scopes)
user.updated_at = utcnow_naive() user.updated_at = utcnow_naive()
db.commit() db.commit()
db.refresh(user) db.refresh(user)
if revoke_all:
revoke_auth_sessions(db, user_id=str(user.id))
return user return user
def change_password(db: Session, *, user: AppUser, old_password: str, new_password: str) -> None: def change_password(
db: Session,
*,
user: AppUser,
old_password: str,
new_password: str,
keep_jti: str | None = None,
) -> None:
row = get_user_by_id(db, str(user.id)) or user row = get_user_by_id(db, str(user.id)) or user
if not verify_password(old_password, row.password_hash): if not verify_password(old_password, row.password_hash):
raise HTTPException(status_code=400, detail="old_password_incorrect") raise HTTPException(status_code=400, detail="old_password_incorrect")
pwd = str(new_password or "") pwd = _require_password_strength(new_password)
if len(pwd) < 6:
raise HTTPException(status_code=400, detail="password_too_short")
default_pwd = str(settings.bootstrap_admin_password or "admin123").strip() or "admin123" default_pwd = str(settings.bootstrap_admin_password or "admin123").strip() or "admin123"
if pwd == default_pwd or pwd == old_password: if pwd == default_pwd or pwd == old_password:
raise HTTPException(status_code=400, detail="password_must_differ_from_default") raise HTTPException(status_code=400, detail="password_must_differ_from_default")
@ -357,6 +567,8 @@ def change_password(db: Session, *, user: AppUser, old_password: str, new_passwo
row.must_change_password = False row.must_change_password = False
row.updated_at = utcnow_naive() row.updated_at = utcnow_naive()
db.commit() db.commit()
# Drop other browser sessions; keep current jti so force-change flow can continue.
revoke_auth_sessions(db, user_id=str(row.id), except_jti=keep_jti)
def create_api_token( def create_api_token(

View file

@ -81,18 +81,29 @@ def auth_secret() -> str:
return ensure_auth_secret() return ensure_auth_secret()
def issue_access_token(*, user_id: str, username: str, role: str) -> str: def issue_access_token(
*,
user_id: str,
username: str,
role: str,
jti: str | None = None,
) -> tuple[str, str, int]:
"""Return (token, jti, ttl_sec). jti is required for server-side revocation."""
ttl = max(300, int(settings.auth_token_ttl_sec or 86400)) ttl = max(300, int(settings.auth_token_ttl_sec or 86400))
now = datetime.now(timezone.utc) now = datetime.now(timezone.utc)
sid = str(jti or secrets.token_urlsafe(24)).strip()
if not sid:
sid = secrets.token_urlsafe(24)
payload = { payload = {
"sub": str(user_id), "sub": str(user_id),
"username": str(username), "username": str(username),
"role": str(role), "role": str(role),
"typ": "access", "typ": "access",
"jti": sid,
"iat": int(now.timestamp()), "iat": int(now.timestamp()),
"exp": int((now + timedelta(seconds=ttl)).timestamp()), "exp": int((now + timedelta(seconds=ttl)).timestamp()),
} }
return jwt.encode(payload, auth_secret(), algorithm="HS256") return jwt.encode(payload, auth_secret(), algorithm="HS256"), sid, ttl
def decode_access_token(token: str) -> dict[str, Any]: def decode_access_token(token: str) -> dict[str, Any]:
@ -100,7 +111,7 @@ def decode_access_token(token: str) -> dict[str, Any]:
str(token or ""), str(token or ""),
auth_secret(), auth_secret(),
algorithms=["HS256"], algorithms=["HS256"],
options={"require": ["exp", "sub"]}, options={"require": ["exp", "sub", "jti"]},
) )
@ -109,5 +120,10 @@ def new_api_token_plaintext() -> str:
return "nxt_" + secrets.token_urlsafe(32) return "nxt_" + secrets.token_urlsafe(32)
def new_refresh_token_plaintext() -> str:
"""Opaque refresh token (shown once). Prefix distinguishes from API tokens."""
return "nxr_" + secrets.token_urlsafe(32)
def hash_api_token(plaintext: str) -> str: def hash_api_token(plaintext: str) -> str:
return hashlib.sha256(str(plaintext or "").encode("utf-8")).hexdigest() return hashlib.sha256(str(plaintext or "").encode("utf-8")).hexdigest()

View file

@ -97,6 +97,8 @@ class Settings(BaseSettings):
ne_exec_max_commands: int = 5 ne_exec_max_commands: int = 5
# WebCRT interactive terminal sessions (multi-operator concurrent terminals). # WebCRT interactive terminal sessions (multi-operator concurrent terminals).
webcrt_max_sessions: int = 40 webcrt_max_sessions: int = 40
# Per-user cap (0 = unlimited beyond global max).
webcrt_max_sessions_per_user: int = 5
webcrt_idle_timeout_sec: int = 1800 webcrt_idle_timeout_sec: int = 1800
webcrt_connect_timeout_sec: int = 90 webcrt_connect_timeout_sec: int = 90
webcrt_attach_timeout_sec: int = 60 webcrt_attach_timeout_sec: int = 60
@ -126,7 +128,25 @@ class Settings(BaseSettings):
# Set explicitly only when you want a shared/ops-managed secret. # Set explicitly only when you want a shared/ops-managed secret.
auth_secret: str = "" auth_secret: str = ""
auth_secret_file: str = "data/auth/jwt_secret" auth_secret_file: str = "data/auth/jwt_secret"
auth_token_ttl_sec: int = 86400 auth_token_ttl_sec: int = 3600
# Refresh token lifetime (default 7 days). Used with POST /v1/auth/refresh.
auth_refresh_ttl_sec: int = 604800
# When true, each login revokes other JWT sessions for that user (single active browser login).
auth_single_session: bool = True
# Login brute-force protection (in-process; resets on restart).
auth_login_max_failures: int = 10
auth_login_window_sec: int = 300
auth_login_lockout_sec: int = 900
auth_password_min_len: int = 8
# Idle revoke: if last_seen_at older than this, JWT session is revoked (0 = off).
auth_idle_timeout_sec: int = 7200
# Browser session cookies (HttpOnly). Bearer header still works for API tokens / scripts.
auth_cookie_enabled: bool = True
# None = auto (HTTPS / X-Forwarded-Proto); True/False force.
auth_cookie_secure: bool | None = None
auth_cookie_samesite: str = "lax"
# Optional Redis URL for shared login rate-limit (empty = in-process only).
auth_redis_url: str = ""
bootstrap_admin_username: str = "admin" bootstrap_admin_username: str = "admin"
bootstrap_admin_password: str = "admin123" bootstrap_admin_password: str = "admin123"
# Written on first boot for MCP; path relative to cwd / absolute # Written on first boot for MCP; path relative to cwd / absolute

View file

@ -8,7 +8,7 @@ from .alarms import (
ImportErrorRow, ImportErrorRow,
ImportJob, ImportJob,
) )
from .auth import ApiToken, AppUser, AuditLog from .auth import ApiToken, AppUser, AuditLog, AuthSession
from .config_sync import ( from .config_sync import (
ConfigSyncCycle, ConfigSyncCycle,
ConfigSyncPolicy, ConfigSyncPolicy,
@ -93,6 +93,7 @@ __all__ = [
"AppUser", "AppUser",
"AuditLog", "AuditLog",
"ApiToken", "ApiToken",
"AuthSession",
"ConfigSyncPolicy", "ConfigSyncPolicy",
"ConfigSyncCycle", "ConfigSyncCycle",
"ConfigSyncTask", "ConfigSyncTask",

View file

@ -61,3 +61,21 @@ class ApiToken(Base):
expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True) expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True)
last_used_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) last_used_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
revoked_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True) revoked_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
class AuthSession(Base):
"""Server-side JWT session (jti). Logout / password change can revoke without waiting for exp."""
__tablename__ = "auth_session"
id: Mapped[str] = mapped_column(String(64), primary_key=True) # JWT jti
user_id: Mapped[str] = mapped_column(String(64), index=True)
created_at: Mapped[datetime] = mapped_column(DateTime, default=utcnow_naive)
expires_at: Mapped[datetime] = mapped_column(DateTime, index=True)
revoked_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True)
client_ip: Mapped[str] = mapped_column(String(128), default="")
user_agent: Mapped[str] = mapped_column(String(512), default="")
last_seen_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True)
# Opaque refresh token (hashed); longer-lived than access JWT.
refresh_token_hash: Mapped[str] = mapped_column(String(128), default="", index=True)
refresh_expires_at: Mapped[datetime | None] = mapped_column(DateTime, nullable=True, index=True)

View file

@ -91,6 +91,36 @@ def apply_auth_schema_patches(conn: Connection) -> None:
_run_sql(conn, "ALTER TABLE app_user ADD COLUMN IF NOT EXISTS scopes JSON DEFAULT '[]'") _run_sql(conn, "ALTER TABLE app_user ADD COLUMN IF NOT EXISTS scopes JSON DEFAULT '[]'")
_run_sql(conn, "ALTER TABLE api_token ADD COLUMN IF NOT EXISTS expires_at TIMESTAMP") _run_sql(conn, "ALTER TABLE api_token ADD COLUMN IF NOT EXISTS expires_at TIMESTAMP")
_run_sql(conn, "ALTER TABLE api_token ADD COLUMN IF NOT EXISTS scopes JSON DEFAULT '[]'") _run_sql(conn, "ALTER TABLE api_token ADD COLUMN IF NOT EXISTS scopes JSON DEFAULT '[]'")
_run_sql(
conn,
"""
CREATE TABLE IF NOT EXISTS auth_session (
id VARCHAR(64) PRIMARY KEY,
user_id VARCHAR(64) NOT NULL,
created_at TIMESTAMP,
expires_at TIMESTAMP,
revoked_at TIMESTAMP,
client_ip VARCHAR(128) DEFAULT '',
user_agent VARCHAR(512) DEFAULT '',
last_seen_at TIMESTAMP
)
""",
)
_run_sql(conn, "CREATE INDEX IF NOT EXISTS ix_auth_session_user_id ON auth_session (user_id)")
_run_sql(conn, "CREATE INDEX IF NOT EXISTS ix_auth_session_expires_at ON auth_session (expires_at)")
_run_sql(conn, "CREATE INDEX IF NOT EXISTS ix_auth_session_revoked_at ON auth_session (revoked_at)")
_run_sql(
conn,
"ALTER TABLE auth_session ADD COLUMN IF NOT EXISTS refresh_token_hash VARCHAR(128) DEFAULT ''",
)
_run_sql(
conn,
"ALTER TABLE auth_session ADD COLUMN IF NOT EXISTS refresh_expires_at TIMESTAMP",
)
_run_sql(
conn,
"CREATE INDEX IF NOT EXISTS ix_auth_session_refresh_token_hash ON auth_session (refresh_token_hash)",
)
def apply_key_alert_schema_patches( def apply_key_alert_schema_patches(

View file

@ -26,6 +26,7 @@ from .webcrt_service import (
list_sessions, list_sessions,
mark_attached, mark_attached,
read_session_log_tail, read_session_log_tail,
session_access_allowed,
wait_session_ready, wait_session_ready,
_decode_bytes, _decode_bytes,
_encode_text, _encode_text,
@ -134,8 +135,9 @@ def _client_label(request: Request | None = None, websocket: WebSocket | None =
@router.get("/sessions") @router.get("/sessions")
def api_list_sessions() -> dict[str, Any]: def api_list_sessions(ctx: AuthContext = Depends(require_user)) -> dict[str, Any]:
return list_sessions() is_admin = str(ctx.user.role or "") == "admin"
return list_sessions(for_user_id=str(ctx.user.id), admin=is_admin)
@router.get("/meta/device-types") @router.get("/meta/device-types")
@ -150,6 +152,7 @@ def api_create_session(
body: WebcrtSessionCreate, body: WebcrtSessionCreate,
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
ctx: AuthContext = Depends(require_user),
) -> dict[str, Any]: ) -> dict[str, Any]:
mid = str(body.ne_id or "").strip() mid = str(body.ne_id or "").strip()
uid = str(body.ume_ne_id or "").strip() uid = str(body.ume_ne_id or "").strip()
@ -168,6 +171,8 @@ def api_create_session(
async_connect=bool(body.async_connect), async_connect=bool(body.async_connect),
username_override=body.username, username_override=body.username,
password_override=body.password, password_override=body.password,
owner_user_id=str(ctx.user.id),
owner_username=str(ctx.user.username),
) )
@ -176,6 +181,7 @@ def api_quick_connect(
body: WebcrtQuickConnectBody, body: WebcrtQuickConnectBody,
request: Request, request: Request,
db: Session = Depends(get_db), db: Session = Depends(get_db),
ctx: AuthContext = Depends(require_user),
) -> dict[str, Any]: ) -> dict[str, Any]:
from .ne_service import upsert_webcrt_session_host from .ne_service import upsert_webcrt_session_host
@ -217,6 +223,8 @@ def api_quick_connect(
async_connect=async_connect, async_connect=async_connect,
username_override=user_override, username_override=user_override,
password_override=pwd_override, password_override=pwd_override,
owner_user_id=str(ctx.user.id),
owner_username=str(ctx.user.username),
) )
except HTTPException as exc: except HTTPException as exc:
# NE row already exists; return it so the UI retries in place (no duplicate hosts). # NE row already exists; return it so the UI retries in place (no duplicate hosts).
@ -241,7 +249,16 @@ def api_quick_connect(
@router.delete("/sessions/{session_id}") @router.delete("/sessions/{session_id}")
def api_close_session(session_id: str, request: Request) -> dict[str, Any]: def api_close_session(
session_id: str,
request: Request,
ctx: AuthContext = Depends(require_user),
) -> dict[str, Any]:
sess = get_session(session_id)
if sess is not None:
is_admin = str(ctx.user.role or "") == "admin"
if not session_access_allowed(sess, user_id=str(ctx.user.id), is_admin=is_admin):
raise HTTPException(status_code=403, detail="webcrt_session_forbidden")
return close_session(session_id, reason="client_delete", client=_client_label(request=request)) return close_session(session_id, reason="client_delete", client=_client_label(request=request))
@ -359,6 +376,8 @@ async def api_sftp_upload(
@router.websocket("/sessions/{session_id}/ws") @router.websocket("/sessions/{session_id}/ws")
async def websocket_session(websocket: WebSocket, session_id: str) -> None: async def websocket_session(websocket: WebSocket, session_id: str) -> None:
actor_user_id = ""
actor_is_admin = False
if bool(settings.auth_enabled): if bool(settings.auth_enabled):
# Prefer one-time ws_ticket (query). Bearer header also accepted (non-browser clients). # Prefer one-time ws_ticket (query). Bearer header also accepted (non-browser clients).
# Long-lived access_token in query is rejected. # Long-lived access_token in query is rejected.
@ -368,6 +387,8 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None:
if info is None or not has_scope(info.scopes, SCOPE_WEBCRT): if info is None or not has_scope(info.scopes, SCOPE_WEBCRT):
await websocket.close(code=4403 if info is not None else 4401) await websocket.close(code=4403 if info is not None else 4401)
return return
actor_user_id = str(info.user_id)
actor_is_admin = has_scope(info.scopes, "admin:users")
else: else:
if str(websocket.query_params.get("access_token") or "").strip(): if str(websocket.query_params.get("access_token") or "").strip():
await websocket.close(code=4401) await websocket.close(code=4401)
@ -384,10 +405,21 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None:
if resolved is None: if resolved is None:
await websocket.close(code=4401) await websocket.close(code=4401)
return return
_user, _via, scopes, _tid = resolved user, _via, scopes, _tid, _jti = resolved
if not has_scope(scopes, SCOPE_WEBCRT): if not has_scope(scopes, SCOPE_WEBCRT):
await websocket.close(code=4403) await websocket.close(code=4403)
return return
actor_user_id = str(user.id)
actor_is_admin = str(user.role or "") == "admin" or has_scope(scopes, "admin:users")
# Ownership check before accept when session already exists.
existing = get_session(session_id)
if existing is not None and bool(settings.auth_enabled):
if not session_access_allowed(
existing, user_id=actor_user_id, is_admin=actor_is_admin
):
await websocket.close(code=4403)
return
await websocket.accept() await websocket.accept()
attach_gen = 0 attach_gen = 0
@ -398,6 +430,13 @@ async def websocket_session(websocket: WebSocket, session_id: str) -> None:
await websocket.close(code=4404 if exc.status_code == 404 else 4409) await websocket.close(code=4404 if exc.status_code == 404 else 4409)
return return
if bool(settings.auth_enabled) and not session_access_allowed(
sess, user_id=actor_user_id, is_admin=actor_is_admin
):
await websocket.send_json({"type": "status", "state": "error", "message": "forbidden"})
await websocket.close(code=4403)
return
await websocket.send_json( await websocket.send_json(
{ {
"type": "status", "type": "status",

View file

@ -29,6 +29,7 @@ from .webcrt_session import (
_webcrt_creds_ready, _webcrt_creds_ready,
active_session_count, active_session_count,
close_session, close_session,
close_sessions_for_user,
create_session, create_session,
detach_session, detach_session,
find_ssh_session_for_ne, find_ssh_session_for_ne,
@ -41,6 +42,7 @@ from .webcrt_session_registry import (
_reap_sessions, _reap_sessions,
_sessions, _sessions,
_sessions_lock, _sessions_lock,
session_access_allowed,
) )
__all__ = [ __all__ = [
@ -63,6 +65,7 @@ __all__ = [
"active_session_count", "active_session_count",
"channel_return", "channel_return",
"close_session", "close_session",
"close_sessions_for_user",
"create_session", "create_session",
"detach_session", "detach_session",
"find_ssh_session_for_ne", "find_ssh_session_for_ne",
@ -76,6 +79,7 @@ __all__ = [
"open_netmiko_connection", "open_netmiko_connection",
"prepare_bootstrap_output", "prepare_bootstrap_output",
"read_session_log_tail", "read_session_log_tail",
"session_access_allowed",
"settings", "settings",
"uses_network_cli_keymap", "uses_network_cli_keymap",
"wait_session_ready", "wait_session_ready",

View file

@ -6,12 +6,14 @@ from .webcrt_session_registry import (
_webcrt_creds_ready, _webcrt_creds_ready,
active_session_count, active_session_count,
close_session, close_session,
close_sessions_for_user,
create_session, create_session,
detach_session, detach_session,
find_ssh_session_for_ne, find_ssh_session_for_ne,
get_session, get_session,
list_sessions, list_sessions,
mark_attached, mark_attached,
session_access_allowed,
wait_session_ready, wait_session_ready,
) )
@ -20,11 +22,13 @@ __all__ = [
"_webcrt_creds_ready", "_webcrt_creds_ready",
"active_session_count", "active_session_count",
"close_session", "close_session",
"close_sessions_for_user",
"create_session", "create_session",
"detach_session", "detach_session",
"find_ssh_session_for_ne", "find_ssh_session_for_ne",
"get_session", "get_session",
"list_sessions", "list_sessions",
"mark_attached", "mark_attached",
"session_access_allowed",
"wait_session_ready", "wait_session_ready",
] ]

View file

@ -42,6 +42,9 @@ class WebcrtSession:
cli_keymap: bool = True cli_keymap: bool = True
encoding: str = "utf-8" encoding: str = "utf-8"
keepalive_sec: int = 0 keepalive_sec: int = 0
# Owning netx user; empty = legacy unbound (tests / auth_disabled).
owner_user_id: str = ""
owner_username: str = ""
conn: ConnectHandler | None = None conn: ConnectHandler | None = None
created_at: float = field(default_factory=time.time) created_at: float = field(default_factory=time.time)
last_activity: float = field(default_factory=time.time) last_activity: float = field(default_factory=time.time)

View file

@ -125,6 +125,37 @@ def active_session_count() -> int:
return sum(1 for s in _sessions.values() if not s.closed) return sum(1 for s in _sessions.values() if not s.closed)
def active_session_count_for_user(user_id: str) -> int:
uid = str(user_id or "").strip()
if not uid:
return 0
with _sessions_lock:
return sum(
1
for s in _sessions.values()
if (not s.closed) and str(s.owner_user_id or "").strip() == uid
)
def close_sessions_for_user(user_id: str, *, reason: str = "owner_logout") -> int:
"""Close all WebCRT sessions owned by user_id. Returns count closed."""
uid = str(user_id or "").strip()
if not uid:
return 0
with _sessions_lock:
ids = [
sid
for sid, s in _sessions.items()
if (not s.closed) and str(s.owner_user_id or "").strip() == uid
]
closed = 0
for sid in ids:
out = close_session(sid, reason=reason, client="auth_logout")
if out.get("closed"):
closed += 1
return closed
def get_session(session_id: str) -> WebcrtSession | None: def get_session(session_id: str) -> WebcrtSession | None:
with _sessions_lock: with _sessions_lock:
sess = _sessions.get(session_id) sess = _sessions.get(session_id)
@ -374,6 +405,8 @@ def create_session(
async_connect: bool = True, async_connect: bool = True,
username_override: str | None = None, username_override: str | None = None,
password_override: str | None = None, password_override: str | None = None,
owner_user_id: str = "",
owner_username: str = "",
) -> dict[str, Any]: ) -> dict[str, Any]:
from .cli_resolve import resolve_cli_target from .cli_resolve import resolve_cli_target
@ -381,6 +414,10 @@ def create_session(
max_sessions = max(1, int(settings.webcrt_max_sessions or 20)) max_sessions = max(1, int(settings.webcrt_max_sessions or 20))
if active_session_count() >= max_sessions: if active_session_count() >= max_sessions:
raise HTTPException(status_code=429, detail="webcrt_session_limit") raise HTTPException(status_code=429, detail="webcrt_session_limit")
owner_id = str(owner_user_id or "").strip()
per_user = int(getattr(settings, "webcrt_max_sessions_per_user", 5) or 0)
if owner_id and per_user > 0 and active_session_count_for_user(owner_id) >= per_user:
raise HTTPException(status_code=429, detail="webcrt_user_session_limit")
mid = str(ne_id or "").strip() mid = str(ne_id or "").strip()
uid = str(ume_ne_id or "").strip() uid = str(ume_ne_id or "").strip()
@ -425,6 +462,7 @@ def create_session(
else: else:
ka = max(0, min(600, int(keepalive_sec))) ka = max(0, min(600, int(keepalive_sec)))
owner_name = str(owner_username or "").strip()
sess = WebcrtSession( sess = WebcrtSession(
session_id=session_id, session_id=session_id,
ne_id=target_id, ne_id=target_id,
@ -438,6 +476,8 @@ def create_session(
cli_keymap=cli_keymap, cli_keymap=cli_keymap,
encoding=enc, encoding=enc,
keepalive_sec=ka, keepalive_sec=ka,
owner_user_id=owner_id,
owner_username=owner_name,
state="connecting", state="connecting",
post_login_commands=list(post_login_commands or [])[:20], post_login_commands=list(post_login_commands or [])[:20],
) )
@ -497,9 +537,26 @@ def create_session(
"ws_path": f"/v1/webcrt/sessions/{session_id}/ws", "ws_path": f"/v1/webcrt/sessions/{session_id}/ws",
"cli_hop": bool(sess.cli_hop_guard), "cli_hop": bool(sess.cli_hop_guard),
"sftp_ready": bool(sess.sftp_ready), "sftp_ready": bool(sess.sftp_ready),
"owner_user_id": sess.owner_user_id,
"owner_username": sess.owner_username,
} }
def session_access_allowed(
sess: WebcrtSession,
*,
user_id: str,
is_admin: bool = False,
) -> bool:
"""Owner or admin may attach/close. Unbound sessions (empty owner) stay open for lab/tests."""
owner = str(sess.owner_user_id or "").strip()
if not owner:
return True
if is_admin:
return True
return owner == str(user_id or "").strip()
def mark_attached(session_id: str) -> tuple[WebcrtSession, int]: def mark_attached(session_id: str) -> tuple[WebcrtSession, int]:
sess = get_session(session_id) sess = get_session(session_id)
if sess is None: if sess is None:
@ -595,12 +652,21 @@ def close_all_sessions(*, reason: str = "shutdown") -> int:
return closed return closed
def list_sessions() -> dict[str, Any]: def list_sessions(
*,
for_user_id: str | None = None,
admin: bool = False,
) -> dict[str, Any]:
"""List active sessions. Non-admin callers only see their own owned sessions."""
viewer = str(for_user_id or "").strip()
with _sessions_lock: with _sessions_lock:
items = [] items = []
for s in _sessions.values(): for s in _sessions.values():
if s.closed: if s.closed:
continue continue
owner = str(s.owner_user_id or "").strip()
if viewer and not admin and owner and owner != viewer:
continue
state = str(s.state or "unknown") state = str(s.state or "unknown")
attached = bool(s.attached) attached = bool(s.attached)
# Lifecycle for ops UI: distinguish login vs live vs grace-period detach. # Lifecycle for ops UI: distinguish login vs live vs grace-period detach.
@ -643,6 +709,8 @@ def list_sessions() -> dict[str, Any]:
if s.connect_finished_at if s.connect_finished_at
else None else None
), ),
"owner_user_id": s.owner_user_id,
"owner_username": s.owner_username,
} }
) )
return { return {

View file

@ -29,6 +29,7 @@ dependencies = [
"PyJWT>=2.8.0", "PyJWT>=2.8.0",
"alembic>=1.13.0", "alembic>=1.13.0",
"psutil>=5.9.0", "psutil>=5.9.0",
"redis>=5.0.0",
] ]
[project.optional-dependencies] [project.optional-dependencies]

View file

@ -21,3 +21,4 @@ bcrypt>=4.1.0
PyJWT>=2.8.0 PyJWT>=2.8.0
alembic>=1.13.0 alembic>=1.13.0
psutil>=5.9.0 psutil>=5.9.0
redis>=5.0.0

View file

@ -30,11 +30,13 @@ class AuthUnitTests(unittest.TestCase):
with patch("netx_api.auth_tokens.settings") as st: with patch("netx_api.auth_tokens.settings") as st:
st.auth_secret = "test-secret-key-for-jwt" st.auth_secret = "test-secret-key-for-jwt"
st.auth_token_ttl_sec = 3600 st.auth_token_ttl_sec = 3600
tok = issue_access_token(user_id="u1", username="admin", role="admin") tok, jti, ttl = issue_access_token(user_id="u1", username="admin", role="admin")
payload = decode_access_token(tok) payload = decode_access_token(tok)
self.assertEqual(payload["sub"], "u1") self.assertEqual(payload["sub"], "u1")
self.assertEqual(payload["username"], "admin") self.assertEqual(payload["username"], "admin")
self.assertEqual(payload["role"], "admin") self.assertEqual(payload["role"], "admin")
self.assertEqual(payload["jti"], jti)
self.assertEqual(ttl, 3600)
class AuthApiTests(unittest.TestCase): class AuthApiTests(unittest.TestCase):
@ -81,6 +83,10 @@ class AuthApiTests(unittest.TestCase):
db = self.Session() db = self.Session()
try: try:
bootstrap_admin_if_needed(db) bootstrap_admin_if_needed(db)
# Most tests exercise normal APIs; password-change gate is covered separately.
admin = db.query(AppUser).filter(AppUser.username == "admin").one()
admin.must_change_password = False
db.commit()
finally: finally:
db.close() db.close()
@ -114,12 +120,17 @@ class AuthApiTests(unittest.TestCase):
db = self.Session() db = self.Session()
try: try:
admin = db.query(AppUser).filter(AppUser.username == "admin").one() admin = db.query(AppUser).filter(AppUser.username == "admin").one()
admin.must_change_password = True
db.commit()
self.assertTrue(admin.must_change_password) self.assertTrue(admin.must_change_password)
finally: finally:
db.close() db.close()
token = self._login() token = self._login()
me = self.client.get("/v1/auth/me", headers={"Authorization": f"Bearer {token}"}) me = self.client.get("/v1/auth/me", headers={"Authorization": f"Bearer {token}"})
self.assertTrue(me.json()["user"]["must_change_password"]) self.assertTrue(me.json()["user"]["must_change_password"])
blocked = self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(blocked.status_code, 403)
self.assertEqual(blocked.json()["detail"], "password_change_required")
bad = self.client.post( bad = self.client.post(
"/v1/auth/change-password", "/v1/auth/change-password",
headers={"Authorization": f"Bearer {token}"}, headers={"Authorization": f"Bearer {token}"},
@ -134,6 +145,154 @@ class AuthApiTests(unittest.TestCase):
self.assertEqual(ok.status_code, 200, ok.text) self.assertEqual(ok.status_code, 200, ok.text)
me2 = self.client.get("/v1/auth/me", headers={"Authorization": f"Bearer {token}"}) me2 = self.client.get("/v1/auth/me", headers={"Authorization": f"Bearer {token}"})
self.assertFalse(me2.json()["user"]["must_change_password"]) self.assertFalse(me2.json()["user"]["must_change_password"])
probe = self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(probe.status_code, 200)
def test_logout_revokes_jwt(self) -> None:
token = self._login()
r = self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(r.status_code, 200)
out = self.client.post("/v1/auth/logout", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(out.status_code, 200, out.text)
self.assertGreaterEqual(int(out.json().get("revoked") or 0), 1)
r2 = self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(r2.status_code, 401)
def test_refresh_rotates_tokens(self) -> None:
login = self.client.post("/v1/auth/login", json={"username": "admin", "password": "adminpass"})
self.assertEqual(login.status_code, 200, login.text)
body = login.json()
access = body["access_token"]
refresh = body["refresh_token"]
self.assertTrue(str(refresh).startswith("nxr_"))
# Access works
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {access}"}).status_code,
200,
)
rotated = self.client.post("/v1/auth/refresh", json={"refresh_token": refresh})
self.assertEqual(rotated.status_code, 200, rotated.text)
new_access = rotated.json()["access_token"]
new_refresh = rotated.json()["refresh_token"]
self.assertNotEqual(access, new_access)
self.assertNotEqual(refresh, new_refresh)
# Old access revoked
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {access}"}).status_code,
401,
)
# New access works
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {new_access}"}).status_code,
200,
)
# Old refresh cannot be reused
reuse = self.client.post("/v1/auth/refresh", json={"refresh_token": refresh})
self.assertEqual(reuse.status_code, 401)
def test_single_session_login_revokes_others(self) -> None:
token = self._login()
token2 = self._login()
# Default auth_single_session=True: first login is kicked.
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"}).status_code,
401,
)
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token2}"}).status_code,
200,
)
listed = self.client.get("/v1/auth/sessions", headers={"Authorization": f"Bearer {token2}"})
self.assertEqual(listed.status_code, 200, listed.text)
self.assertEqual(listed.json()["total"], 1)
self.assertTrue(listed.json()["items"][0].get("current"))
def test_list_and_revoke_sessions(self) -> None:
with patch("netx_api.auth_service.settings.auth_single_session", False):
token = self._login()
listed = self.client.get("/v1/auth/sessions", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(listed.status_code, 200, listed.text)
items = listed.json()["items"]
self.assertGreaterEqual(len(items), 1)
self.assertTrue(any(i.get("current") for i in items))
# Multi-session mode: second login keeps the first alive.
token2 = self._login()
listed2 = self.client.get("/v1/auth/sessions", headers={"Authorization": f"Bearer {token2}"})
self.assertGreaterEqual(listed2.json()["total"], 2)
revoked = self.client.post(
"/v1/auth/sessions/revoke-others",
headers={"Authorization": f"Bearer {token2}"},
)
self.assertEqual(revoked.status_code, 200, revoked.text)
self.assertGreaterEqual(int(revoked.json().get("revoked") or 0), 1)
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"}).status_code,
401,
)
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token2}"}).status_code,
200,
)
def test_idle_timeout_revokes(self) -> None:
from datetime import timedelta
from netx_api.models import AuthSession
from netx_api.timeutil import utcnow_naive
token = self._login()
with patch("netx_api.auth_service.settings.auth_idle_timeout_sec", 60):
db = self.Session()
try:
row = db.query(AuthSession).filter(AuthSession.revoked_at.is_(None)).first()
self.assertIsNotNone(row)
row.last_seen_at = utcnow_naive() - timedelta(seconds=120)
db.commit()
finally:
db.close()
self.assertEqual(
self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"}).status_code,
401,
)
def test_login_sets_auth_cookies(self) -> None:
r = self.client.post("/v1/auth/login", json={"username": "admin", "password": "adminpass"})
self.assertEqual(r.status_code, 200, r.text)
# Starlette TestClient exposes set cookies
self.assertIn("netx_at", r.cookies)
self.assertIn("netx_rt", r.cookies)
me = self.client.get("/v1/auth/me") # cookie auth
self.assertEqual(me.status_code, 200, me.text)
self.assertEqual(me.json()["user"]["username"], "admin")
def test_query_access_token_rejected(self) -> None:
token = self._login()
# Drop HttpOnly session cookies so only the deprecated query param remains.
self.client.cookies.clear()
r = self.client.get(f"/v1/probe?access_token={token}")
self.assertEqual(r.status_code, 401)
# Same token still works via Bearer header.
r2 = self.client.get("/v1/probe", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(r2.status_code, 200)
def test_login_lockout(self) -> None:
from netx_api.auth_rate_limit import reset_login_rate_limit_for_tests
reset_login_rate_limit_for_tests()
with patch("netx_api.auth_rate_limit.settings.auth_login_max_failures", 3):
with patch("netx_api.auth_rate_limit.settings.auth_login_lockout_sec", 120):
for _ in range(3):
r = self.client.post(
"/v1/auth/login", json={"username": "admin", "password": "wrong"}
)
self.assertIn(r.status_code, (401, 429))
locked = self.client.post(
"/v1/auth/login", json={"username": "admin", "password": "wrong"}
)
self.assertEqual(locked.status_code, 429)
detail = locked.json()["detail"]
self.assertEqual(detail["error"], "login_locked")
reset_login_rate_limit_for_tests()
def test_login_and_me(self) -> None: def test_login_and_me(self) -> None:
token = self._login() token = self._login()
@ -167,10 +326,10 @@ class AuthApiTests(unittest.TestCase):
db = self.Session() db = self.Session()
try: try:
admin = db.query(AppUser).filter(AppUser.username == "admin").one() admin = db.query(AppUser).filter(AppUser.username == "admin").one()
create_user(db, username="alice", password="alice12", role="user", actor=admin) create_user(db, username="alice", password="alice123", role="user", actor=admin)
finally: finally:
db.close() db.close()
token = self._login("alice", "alice12") token = self._login("alice", "alice123")
r = self.client.post( r = self.client.post(
"/v1/users", "/v1/users",
headers={"Authorization": f"Bearer {token}"}, headers={"Authorization": f"Bearer {token}"},
@ -209,7 +368,7 @@ class AuthApiTests(unittest.TestCase):
self.client.post( self.client.post(
"/v1/users", "/v1/users",
headers={"Authorization": f"Bearer {token}"}, headers={"Authorization": f"Bearer {token}"},
json={"username": "carol", "password": "carol12", "role": "user"}, json={"username": "carol", "password": "carol123", "role": "user"},
) )
users = self.client.get("/v1/users", headers={"Authorization": f"Bearer {token}"}) users = self.client.get("/v1/users", headers={"Authorization": f"Bearer {token}"})
carol_id = next(u["id"] for u in users.json()["items"] if u["username"] == "carol") carol_id = next(u["id"] for u in users.json()["items"] if u["username"] == "carol")

View file

@ -167,7 +167,9 @@ class RbacApiTests(unittest.TestCase):
try: try:
bootstrap_admin_if_needed(db) bootstrap_admin_if_needed(db)
admin = db.query(AppUser).filter(AppUser.username == "admin").one() admin = db.query(AppUser).filter(AppUser.username == "admin").one()
create_user(db, username="alice", password="alice12", role="user", actor=admin) admin.must_change_password = False
db.commit()
create_user(db, username="alice", password="alice123", role="user", actor=admin)
finally: finally:
db.close() db.close()
self.client = TestClient(self.app) self.client = TestClient(self.app)
@ -184,7 +186,7 @@ class RbacApiTests(unittest.TestCase):
return str(r.json()["access_token"]) return str(r.json()["access_token"])
def test_user_denied_webcrt_and_sql(self) -> None: def test_user_denied_webcrt_and_sql(self) -> None:
token = self._login("alice", "alice12") token = self._login("alice", "alice123")
h = {"Authorization": f"Bearer {token}"} h = {"Authorization": f"Bearer {token}"}
self.assertEqual(self.client.post("/v1/webcrt/sessions", headers=h).status_code, 403) self.assertEqual(self.client.post("/v1/webcrt/sessions", headers=h).status_code, 403)
self.assertEqual( self.assertEqual(
@ -209,7 +211,7 @@ class RbacApiTests(unittest.TestCase):
self.assertNotEqual(r.status_code, 403) self.assertNotEqual(r.status_code, 403)
def test_me_returns_scopes(self) -> None: def test_me_returns_scopes(self) -> None:
token = self._login("alice", "alice12") token = self._login("alice", "alice123")
me = self.client.get("/v1/auth/me", headers={"Authorization": f"Bearer {token}"}) me = self.client.get("/v1/auth/me", headers={"Authorization": f"Bearer {token}"})
self.assertEqual(me.status_code, 200) self.assertEqual(me.status_code, 200)
scopes = me.json()["scopes"] scopes = me.json()["scopes"]

View file

@ -32,6 +32,7 @@ class SchemaPatchesTests(unittest.TestCase):
self.assertIn("must_change_password", user_cols) self.assertIn("must_change_password", user_cols)
self.assertIn("scopes", token_cols) self.assertIn("scopes", token_cols)
self.assertIn("expires_at", token_cols) self.assertIn("expires_at", token_cols)
self.assertIn("auth_session", insp.get_table_names())
def test_domain_patches_do_not_raise(self) -> None: def test_domain_patches_do_not_raise(self) -> None:
with self.engine.begin() as conn: with self.engine.begin() as conn:
@ -49,6 +50,8 @@ class SchemaPatchesTests(unittest.TestCase):
files = sorted(p.name for p in versions.glob("*.py") if p.name != "__init__.py") files = sorted(p.name for p in versions.glob("*.py") if p.name != "__init__.py")
self.assertIn("20260802_scopes.py", files) self.assertIn("20260802_scopes.py", files)
self.assertIn("20260802_legacy_schema.py", files) self.assertIn("20260802_legacy_schema.py", files)
self.assertIn("20260806_auth_session.py", files)
self.assertIn("20260806_auth_refresh.py", files)
text_legacy = (versions / "20260802_legacy_schema.py").read_text(encoding="utf-8") text_legacy = (versions / "20260802_legacy_schema.py").read_text(encoding="utf-8")
self.assertIn('down_revision', text_legacy) self.assertIn('down_revision', text_legacy)
self.assertIn("20260802_scopes", text_legacy) self.assertIn("20260802_scopes", text_legacy)

View file

@ -54,6 +54,9 @@ const AuditPage = lazy(() => import("./pages/AuditPage").then((m) => ({ default:
const ApiTokensPage = lazy(() => const ApiTokensPage = lazy(() =>
import("./pages/ApiTokensPage").then((m) => ({ default: m.ApiTokensPage })), import("./pages/ApiTokensPage").then((m) => ({ default: m.ApiTokensPage })),
); );
const SessionsPage = lazy(() =>
import("./pages/SessionsPage").then((m) => ({ default: m.SessionsPage })),
);
/** Preserve query when redirecting legacy /network/webcrt → /webcrt. */ /** Preserve query when redirecting legacy /network/webcrt → /webcrt. */
function NetworkWebcrtRedirect() { function NetworkWebcrtRedirect() {
@ -129,6 +132,7 @@ function ProtectedApp() {
<Route path="logs" element={<AuditPage />} /> <Route path="logs" element={<AuditPage />} />
</Route> </Route>
<Route path="/api-keys" element={<ApiTokensPage />} /> <Route path="/api-keys" element={<ApiTokensPage />} />
<Route path="/sessions" element={<SessionsPage />} />
<Route path="*" element={<Navigate to="/" replace />} /> <Route path="*" element={<Navigate to="/" replace />} />
</Routes> </Routes>
</Suspense> </Suspense>

View file

@ -7,14 +7,7 @@ import {
useState, useState,
type ReactNode, type ReactNode,
} from "react"; } from "react";
import { import { apiGet, apiPost, clearAuthToken } from "../services/api";
AUTH_TOKEN_KEY,
apiGet,
apiPost,
clearAuthToken,
getAuthToken,
setAuthToken,
} from "../services/api";
export type AuthUser = { export type AuthUser = {
id: string; id: string;
@ -44,25 +37,19 @@ const AuthContext = createContext<AuthState | null>(null);
export function AuthProvider({ children }: { children: ReactNode }) { export function AuthProvider({ children }: { children: ReactNode }) {
const [ready, setReady] = useState(false); const [ready, setReady] = useState(false);
const [token, setToken] = useState<string | null>(() => getAuthToken()); // token is opaque for UI; cookie session means we only care about user presence.
const [token, setToken] = useState<string | null>(null);
const [user, setUser] = useState<AuthUser | null>(null); const [user, setUser] = useState<AuthUser | null>(null);
const [scopes, setScopes] = useState<string[]>([]); const [scopes, setScopes] = useState<string[]>([]);
const refreshMe = useCallback(async () => { const refreshMe = useCallback(async () => {
const tok = getAuthToken(); clearAuthToken(); // drop any legacy localStorage tokens
if (!tok) {
setToken(null);
setUser(null);
setScopes([]);
return;
}
try { try {
const data = await apiGet<{ user: AuthUser; scopes?: string[] }>("/v1/auth/me"); const data = await apiGet<{ user: AuthUser; scopes?: string[] }>("/v1/auth/me");
setToken(tok); setToken("cookie");
setUser(data.user); setUser(data.user);
setScopes(data.scopes || data.user.scopes || []); setScopes(data.scopes || data.user.scopes || []);
} catch { } catch {
clearAuthToken();
setToken(null); setToken(null);
setUser(null); setUser(null);
setScopes([]); setScopes([]);
@ -76,40 +63,20 @@ export function AuthProvider({ children }: { children: ReactNode }) {
})(); })();
}, [refreshMe]); }, [refreshMe]);
// Other tabs keep React auth state until they hear localStorage change.
// `storage` fires only in *other* documents — used to sync logout/login.
useEffect(() => {
const onStorage = (ev: StorageEvent) => {
if (ev.storageArea && ev.storageArea !== localStorage) return;
if (ev.key !== null && ev.key !== AUTH_TOKEN_KEY) return;
if (ev.key === null || ev.newValue == null || ev.newValue === "") {
setToken(null);
setUser(null);
setScopes([]);
return;
}
void refreshMe();
};
window.addEventListener("storage", onStorage);
return () => window.removeEventListener("storage", onStorage);
}, [refreshMe]);
const login = useCallback(async (username: string, password: string) => { const login = useCallback(async (username: string, password: string) => {
const data = await apiPost<{ access_token: string; user: AuthUser }>("/v1/auth/login", { const data = await apiPost<{ user: AuthUser }>("/v1/auth/login", {
username, username,
password, password,
}); });
setAuthToken(data.access_token); clearAuthToken();
setToken(data.access_token); setToken("cookie");
setUser(data.user); setUser(data.user);
setScopes(data.user.scopes || []); setScopes(data.user.scopes || []);
}, []); }, []);
const logout = useCallback(async () => { const logout = useCallback(async () => {
try { try {
if (getAuthToken()) {
await apiPost("/v1/auth/logout", {}); await apiPost("/v1/auth/logout", {});
}
} catch { } catch {
// ignore // ignore
} }

View file

@ -130,6 +130,16 @@ export const MODULES: readonly ModuleDefinition[] = [
iconKind: "key", iconKind: "key",
titleKey: "layout.titleApiKeys", titleKey: "layout.titleApiKeys",
}, },
{
moduleId: "sessions",
path: "/sessions",
section: "system",
labelKey: "workbench.cards.sessions",
descKey: "workbench.cards.sessionsDesc",
iconTone: "slate",
iconKind: "key",
titleKey: "layout.titleSessions",
},
] as const satisfies readonly ModuleDefinition[]; ] as const satisfies readonly ModuleDefinition[];
export function getModuleById(moduleId: string): ModuleDefinition | undefined { export function getModuleById(moduleId: string): ModuleDefinition | undefined {

View file

@ -61,6 +61,7 @@ const en = {
logs: "Audit logs", logs: "Audit logs",
users: "User admin", users: "User admin",
apiKeys: "API keys", apiKeys: "API keys",
sessions: "Sessions",
}, },
announce: { announce: {
a1: "Dark workbench shell is live — other modules follow the same palette.", a1: "Dark workbench shell is live — other modules follow the same palette.",
@ -97,6 +98,8 @@ const en = {
auditDesc: "Live task overview and operation logs", auditDesc: "Live task overview and operation logs",
apiKeys: "API Keys", apiKeys: "API Keys",
apiKeysDesc: "Issue MCP/script tokens per user with expiry", apiKeysDesc: "Issue MCP/script tokens per user with expiry",
sessions: "Login sessions",
sessionsDesc: "Review and revoke logins on other devices",
}, },
}, },
network: { network: {
@ -519,6 +522,7 @@ const en = {
titleUsers: "Users", titleUsers: "Users",
titleAudit: "Audit", titleAudit: "Audit",
titleApiKeys: "API Keys", titleApiKeys: "API Keys",
titleSessions: "Login sessions",
navUme: "UME", navUme: "UME",
netxApi: "netx api", netxApi: "netx api",
oclawBridge: "oclaw WSS", oclawBridge: "oclaw WSS",
@ -538,6 +542,18 @@ const en = {
loggingIn: "Signing in…", loggingIn: "Signing in…",
loginFailed: "Login failed", loginFailed: "Login failed",
logout: "Sign out", logout: "Sign out",
sessionsTitle: "Login sessions",
sessionsHint: "Manage browser/device logins for this account. Revoked sessions must sign in again.",
sessionsEmpty: "No active sessions.",
revokeOtherSessions: "Revoke other sessions",
revokeSession: "Revoke",
sessionCurrent: "Current",
sessionRevoked: "Session revoked",
sessionsRevokedOthers: "Revoked {{count}} other session(s)",
revokeCurrentConfirm: "This is your current session; revoking it requires signing in again. Continue?",
colSession: "Session",
colLastSeen: "Last seen",
colCreated: "Created",
usersTitle: "User management", usersTitle: "User management",
usersHint: "Only admins can create and manage local accounts.", usersHint: "Only admins can create and manage local accounts.",
addUser: "Add user", addUser: "Add user",
@ -614,7 +630,7 @@ const en = {
confirmPassword: "Confirm new password", confirmPassword: "Confirm new password",
savePassword: "Save new password", savePassword: "Save new password",
savingPassword: "Saving…", savingPassword: "Saving…",
passwordTooShort: "New password must be at least 6 characters", passwordTooShort: "New password must be at least 8 characters",
passwordMismatch: "New passwords do not match", passwordMismatch: "New passwords do not match",
passwordMustChange: "New password must differ from the default/old password", passwordMustChange: "New password must differ from the default/old password",
}, },

View file

@ -61,6 +61,7 @@ const zh = {
logs: "操作日志", logs: "操作日志",
users: "用户管理", users: "用户管理",
apiKeys: "API Key", apiKeys: "API Key",
sessions: "登录会话",
}, },
announce: { announce: {
a1: "深色工作台已上线,其它模块将沿用同一套深色体系。", a1: "深色工作台已上线,其它模块将沿用同一套深色体系。",
@ -97,6 +98,8 @@ const zh = {
auditDesc: "任务概览与操作日志", auditDesc: "任务概览与操作日志",
apiKeys: "API Key", apiKeys: "API Key",
apiKeysDesc: "为用户生成 MCP/脚本用 Token,可设有效期", apiKeysDesc: "为用户生成 MCP/脚本用 Token,可设有效期",
sessions: "登录会话",
sessionsDesc: "查看并踢掉其他设备上的登录",
}, },
}, },
network: { network: {
@ -515,6 +518,7 @@ const zh = {
titleUsers: "用户管理", titleUsers: "用户管理",
titleAudit: "操作审计", titleAudit: "操作审计",
titleApiKeys: "API Key", titleApiKeys: "API Key",
titleSessions: "登录会话",
navUme: "UME 对接", navUme: "UME 对接",
netxApi: "netx api", netxApi: "netx api",
oclawBridge: "oclaw WSS", oclawBridge: "oclaw WSS",
@ -534,6 +538,18 @@ const zh = {
loggingIn: "登录中…", loggingIn: "登录中…",
loginFailed: "登录失败", loginFailed: "登录失败",
logout: "退出", logout: "退出",
sessionsTitle: "登录会话",
sessionsHint: "管理当前账号在各浏览器/设备上的登录。踢掉会话后对方需重新登录。",
sessionsEmpty: "当前没有活跃会话。",
revokeOtherSessions: "踢掉其他会话",
revokeSession: "踢掉",
sessionCurrent: "当前",
sessionRevoked: "会话已吊销",
sessionsRevokedOthers: "已踢掉 {{count}} 个其他会话",
revokeCurrentConfirm: "这是当前会话,踢掉后需要重新登录。继续?",
colSession: "会话",
colLastSeen: "最近活动",
colCreated: "创建时间",
usersTitle: "用户管理", usersTitle: "用户管理",
usersHint: "仅管理员可创建与管理本地账号。", usersHint: "仅管理员可创建与管理本地账号。",
addUser: "添加用户", addUser: "添加用户",
@ -609,7 +625,7 @@ const zh = {
confirmPassword: "确认新密码", confirmPassword: "确认新密码",
savePassword: "保存新密码", savePassword: "保存新密码",
savingPassword: "保存中…", savingPassword: "保存中…",
passwordTooShort: "新密码至少 6 位", passwordTooShort: "新密码至少 8 位",
passwordMismatch: "两次输入的新密码不一致", passwordMismatch: "两次输入的新密码不一致",
passwordMustChange: "新密码不能与默认/旧密码相同", passwordMustChange: "新密码不能与默认/旧密码相同",
}, },

View file

@ -16,7 +16,7 @@ export function ForceChangePasswordPage() {
const onSubmit = async (e: FormEvent) => { const onSubmit = async (e: FormEvent) => {
e.preventDefault(); e.preventDefault();
setError(""); setError("");
if (newPassword.length < 6) { if (newPassword.length < 8) {
setError(t("auth.passwordTooShort")); setError(t("auth.passwordTooShort"));
return; return;
} }
@ -77,7 +77,7 @@ export function ForceChangePasswordPage() {
onChange={(e) => setNewPassword(e.target.value)} onChange={(e) => setNewPassword(e.target.value)}
disabled={busy} disabled={busy}
required required
minLength={6} minLength={8}
/> />
</label> </label>
<label className="login-card__label"> <label className="login-card__label">
@ -90,7 +90,7 @@ export function ForceChangePasswordPage() {
onChange={(e) => setConfirm(e.target.value)} onChange={(e) => setConfirm(e.target.value)}
disabled={busy} disabled={busy}
required required
minLength={6} minLength={8}
/> />
</label> </label>
{error ? ( {error ? (

View file

@ -0,0 +1,118 @@
import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query";
import { useI18n } from "../i18n";
import { useToast } from "../hooks/useToast";
import { apiDelete, apiGet, apiPost } from "../services/api";
import { formatSystemTime } from "../utils/time";
type SessionRow = {
id: string;
created_at: string | null;
expires_at: string | null;
refresh_expires_at: string | null;
last_seen_at: string | null;
client_ip: string;
user_agent: string;
current: boolean;
};
export function SessionsPage() {
const { t } = useI18n();
const { showOk, showError } = useToast();
const qc = useQueryClient();
const sessionsQuery = useQuery({
queryKey: ["auth-sessions"],
queryFn: () => apiGet<{ items: SessionRow[]; total: number }>("/v1/auth/sessions"),
});
const revokeMut = useMutation({
mutationFn: (id: string) => apiDelete(`/v1/auth/sessions/${encodeURIComponent(id)}`),
onSuccess: () => {
showOk(t("auth.sessionRevoked"));
void qc.invalidateQueries({ queryKey: ["auth-sessions"] });
},
onError: (err) => showError(String(err instanceof Error ? err.message : err)),
});
const revokeOthersMut = useMutation({
mutationFn: () => apiPost<{ revoked: number }>("/v1/auth/sessions/revoke-others", {}),
onSuccess: (data) => {
showOk(t("auth.sessionsRevokedOthers", { count: data.revoked ?? 0 }));
void qc.invalidateQueries({ queryKey: ["auth-sessions"] });
},
onError: (err) => showError(String(err instanceof Error ? err.message : err)),
});
const items = sessionsQuery.data?.items || [];
return (
<div className="page">
<header className="page-header">
<h1>{t("auth.sessionsTitle")}</h1>
<p className="panel__hint">{t("auth.sessionsHint")}</p>
</header>
<div className="panel">
<div className="filter-inline" style={{ marginBottom: 12 }}>
<button
type="button"
disabled={revokeOthersMut.isPending || items.filter((s) => !s.current).length === 0}
onClick={() => revokeOthersMut.mutate()}
>
{t("auth.revokeOtherSessions")}
</button>
<button type="button" onClick={() => void sessionsQuery.refetch()} disabled={sessionsQuery.isFetching}>
{t("common.refresh")}
</button>
</div>
{sessionsQuery.isLoading ? <p className="muted">{t("common.refreshing")}</p> : null}
{!items.length && !sessionsQuery.isLoading ? (
<p className="muted">{t("auth.sessionsEmpty")}</p>
) : (
<table className="data-table">
<thead>
<tr>
<th>{t("auth.colSession")}</th>
<th>{t("auth.colIp")}</th>
<th>{t("auth.colLastSeen")}</th>
<th>{t("auth.colCreated")}</th>
<th>{t("auth.actions")}</th>
</tr>
</thead>
<tbody>
{items.map((row) => (
<tr key={row.id}>
<td>
<code title={row.user_agent}>{row.id.slice(0, 10)}…</code>
{row.current ? (
<span className="pt-list-status pt-list-status--ok" style={{ marginLeft: 8 }}>
{t("auth.sessionCurrent")}
</span>
) : null}
</td>
<td>{row.client_ip || "—"}</td>
<td>{row.last_seen_at ? formatSystemTime(row.last_seen_at) : "—"}</td>
<td>{row.created_at ? formatSystemTime(row.created_at) : "—"}</td>
<td>
<button
type="button"
disabled={revokeMut.isPending}
onClick={() => {
if (row.current && !window.confirm(t("auth.revokeCurrentConfirm"))) return;
revokeMut.mutate(row.id);
}}
>
{t("auth.revokeSession")}
</button>
</td>
</tr>
))}
</tbody>
</table>
)}
</div>
</div>
);
}

View file

@ -83,7 +83,7 @@ export function UsersPage() {
value={password} value={password}
onChange={(e) => setPassword(e.target.value)} onChange={(e) => setPassword(e.target.value)}
required required
minLength={6} minLength={8}
/> />
<select value={role} onChange={(e) => setRole(e.target.value)}> <select value={role} onChange={(e) => setRole(e.target.value)}>
<option value="user">{t("auth.roleUser")}</option> <option value="user">{t("auth.roleUser")}</option>
@ -155,7 +155,7 @@ export function UsersPage() {
/> />
<button <button
type="button" type="button"
disabled={!resetPwd[u.id] || resetPwd[u.id].length < 6} disabled={!resetPwd[u.id] || resetPwd[u.id].length < 8}
onClick={() => { onClick={() => {
const pwd = resetPwd[u.id]; const pwd = resetPwd[u.id];
patchMut.mutate({ id: u.id, body: { password: pwd } }); patchMut.mutate({ id: u.id, body: { password: pwd } });

View file

@ -50,36 +50,40 @@ import type {
export const AUTH_TOKEN_KEY = "netx_access_token"; export const AUTH_TOKEN_KEY = "netx_access_token";
export const AUTH_REFRESH_KEY = "netx_refresh_token";
export const getAuthToken = (): string | null => { /** Clear legacy localStorage tokens; browser auth now uses HttpOnly cookies. */
try {
return localStorage.getItem(AUTH_TOKEN_KEY);
} catch {
return null;
}
};
export const setAuthToken = (token: string): void => {
localStorage.setItem(AUTH_TOKEN_KEY, String(token || ""));
};
export const clearAuthToken = (): void => { export const clearAuthToken = (): void => {
try { try {
localStorage.removeItem(AUTH_TOKEN_KEY); localStorage.removeItem(AUTH_TOKEN_KEY);
localStorage.removeItem(AUTH_REFRESH_KEY);
} catch { } catch {
// ignore // ignore
} }
}; };
/** @deprecated Cookie session — always null for browser UI. */
export const getAuthToken = (): string | null => null;
/** No-op kept for call-site compatibility during cookie migration. */
export const setAuthToken = (_token: string): void => {
clearAuthToken();
};
export const setAuthTokens = (_access: string, _refresh?: string | null): void => {
clearAuthToken();
};
const fetchCreds: RequestCredentials = "include";
const authHeaders = (extra?: Record<string, string>): Record<string, string> => { const authHeaders = (extra?: Record<string, string>): Record<string, string> => {
const h: Record<string, string> = { accept: "application/json", ...(extra || {}) }; // Browser JWT rides HttpOnly cookies (credentials: include).
const tok = getAuthToken(); // Authorization Bearer is only needed for non-browser API tokens if ever injected.
if (tok) h.authorization = `Bearer ${tok}`; return { accept: "application/json", ...(extra || {}) };
return h;
}; };
const handleUnauthorized = (path: string): void => { const handleUnauthorized = (path: string): void => {
if (path.startsWith("/v1/auth/login")) return; if (path.startsWith("/v1/auth/login") || path.startsWith("/v1/auth/refresh")) return;
clearAuthToken(); clearAuthToken();
if (typeof window !== "undefined" && !window.location.pathname.startsWith("/login")) { if (typeof window !== "undefined" && !window.location.pathname.startsWith("/login")) {
const next = `${window.location.pathname}${window.location.search || ""}`; const next = `${window.location.pathname}${window.location.search || ""}`;
@ -87,6 +91,38 @@ const handleUnauthorized = (path: string): void => {
} }
}; };
let refreshInFlight: Promise<boolean> | null = null;
const tryRefreshAccessToken = async (): Promise<boolean> => {
if (!refreshInFlight) {
refreshInFlight = (async () => {
try {
// Refresh token comes from HttpOnly cookie when body is empty.
const res = await fetch("/v1/auth/refresh", {
method: "POST",
credentials: fetchCreds,
headers: { accept: "application/json", "content-type": "application/json" },
body: JSON.stringify({}),
});
if (!res.ok) {
clearAuthToken();
return false;
}
return true;
} catch {
clearAuthToken();
return false;
} finally {
refreshInFlight = null;
}
})();
}
return refreshInFlight;
};
const shouldAttemptRefresh = (path: string): boolean =>
!path.startsWith("/v1/auth/login") && !path.startsWith("/v1/auth/refresh");
const parseApiResponse = async (res: Response): Promise<Record<string, unknown>> => { const parseApiResponse = async (res: Response): Promise<Record<string, unknown>> => {
const text = await res.text(); const text = await res.text();
if (!text) return {}; if (!text) return {};
@ -139,7 +175,10 @@ export function formatApiDetail(detail: unknown): string {
} }
export const apiGet = async <T,>(path: string): Promise<T> => { export const apiGet = async <T,>(path: string): Promise<T> => {
const res = await fetch(path, { headers: authHeaders() }); let res = await fetch(path, { headers: authHeaders(), credentials: fetchCreds });
if (res.status === 401 && shouldAttemptRefresh(path) && (await tryRefreshAccessToken())) {
res = await fetch(path, { headers: authHeaders(), credentials: fetchCreds });
}
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
throw new Error("401 unauthorized"); throw new Error("401 unauthorized");
@ -149,11 +188,17 @@ export const apiGet = async <T,>(path: string): Promise<T> => {
}; };
export const apiPost = async <T,>(path: string, body: unknown): Promise<T> => { export const apiPost = async <T,>(path: string, body: unknown): Promise<T> => {
const res = await fetch(path, { const doFetch = () =>
fetch(path, {
method: "POST", method: "POST",
credentials: fetchCreds,
headers: authHeaders({ "content-type": "application/json" }), headers: authHeaders({ "content-type": "application/json" }),
body: JSON.stringify(body), body: JSON.stringify(body),
}); });
let res = await doFetch();
if (res.status === 401 && shouldAttemptRefresh(path) && (await tryRefreshAccessToken())) {
res = await doFetch();
}
const data = await parseApiResponse(res); const data = await parseApiResponse(res);
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
@ -164,11 +209,17 @@ export const apiPost = async <T,>(path: string, body: unknown): Promise<T> => {
}; };
export const apiPatch = async <T,>(path: string, body: unknown): Promise<T> => { export const apiPatch = async <T,>(path: string, body: unknown): Promise<T> => {
const res = await fetch(path, { const doFetch = () =>
fetch(path, {
method: "PATCH", method: "PATCH",
credentials: fetchCreds,
headers: authHeaders({ "content-type": "application/json" }), headers: authHeaders({ "content-type": "application/json" }),
body: JSON.stringify(body), body: JSON.stringify(body),
}); });
let res = await doFetch();
if (res.status === 401 && shouldAttemptRefresh(path) && (await tryRefreshAccessToken())) {
res = await doFetch();
}
const data = await parseApiResponse(res); const data = await parseApiResponse(res);
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
@ -179,7 +230,10 @@ export const apiPatch = async <T,>(path: string, body: unknown): Promise<T> => {
}; };
export const apiDelete = async <T,>(path: string): Promise<T> => { export const apiDelete = async <T,>(path: string): Promise<T> => {
const res = await fetch(path, { method: "DELETE", headers: authHeaders() }); let res = await fetch(path, { method: "DELETE", headers: authHeaders(), credentials: fetchCreds });
if (res.status === 401 && shouldAttemptRefresh(path) && (await tryRefreshAccessToken())) {
res = await fetch(path, { method: "DELETE", headers: authHeaders(), credentials: fetchCreds });
}
const data = await parseApiResponse(res); const data = await parseApiResponse(res);
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
@ -190,11 +244,17 @@ export const apiDelete = async <T,>(path: string): Promise<T> => {
}; };
export const apiPut = async <T,>(path: string, body: unknown): Promise<T> => { export const apiPut = async <T,>(path: string, body: unknown): Promise<T> => {
const res = await fetch(path, { const doFetch = () =>
fetch(path, {
method: "PUT", method: "PUT",
credentials: fetchCreds,
headers: authHeaders({ "content-type": "application/json" }), headers: authHeaders({ "content-type": "application/json" }),
body: JSON.stringify(body), body: JSON.stringify(body),
}); });
let res = await doFetch();
if (res.status === 401 && shouldAttemptRefresh(path) && (await tryRefreshAccessToken())) {
res = await doFetch();
}
const data = await parseApiResponse(res); const data = await parseApiResponse(res);
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
@ -499,7 +559,7 @@ export const managedNeImportTemplateUrl = (format: "xlsx" | "csv" = "xlsx") =>
export const downloadManagedNeImportTemplate = async (format: "xlsx" | "csv" = "xlsx"): Promise<void> => { export const downloadManagedNeImportTemplate = async (format: "xlsx" | "csv" = "xlsx"): Promise<void> => {
const path = managedNeImportTemplateUrl(format); const path = managedNeImportTemplateUrl(format);
const res = await fetch(path, { headers: authHeaders() }); const res = await fetch(path, { headers: authHeaders(), credentials: fetchCreds });
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
throw new Error("unauthorized"); throw new Error("unauthorized");
@ -527,7 +587,7 @@ export const importManagedNe = async (file: File): Promise<ManagedNeImportResult
form.append("file", file); form.append("file", file);
const path = "/v1/managed-ne/import"; const path = "/v1/managed-ne/import";
// Do not set content-type — browser must add multipart boundary. // Do not set content-type — browser must add multipart boundary.
const res = await fetch(path, { method: "POST", headers: authHeaders(), body: form }); const res = await fetch(path, { method: "POST", headers: authHeaders(), body: form, credentials: fetchCreds });
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
throw new Error("unauthorized"); throw new Error("unauthorized");
@ -673,7 +733,7 @@ export function closeWebcrtSessionsKeepalive(sessionIds: string[]): void {
for (const id of ids) { for (const id of ids) {
const path = `/v1/webcrt/sessions/${encodeURIComponent(id)}`; const path = `/v1/webcrt/sessions/${encodeURIComponent(id)}`;
try { try {
void fetch(path, { method: "DELETE", headers, keepalive: true }); void fetch(path, { method: "DELETE", headers, credentials: fetchCreds, keepalive: true });
} catch { } catch {
/* ignore unload failures */ /* ignore unload failures */
} }
@ -772,6 +832,7 @@ async function webcrtSftpDownloadOnce(
const path = "/v1/webcrt/sftp/download"; const path = "/v1/webcrt/sftp/download";
const res = await fetch(path, { const res = await fetch(path, {
method: "POST", method: "POST",
credentials: fetchCreds,
headers: authHeaders({ "Content-Type": "application/json" }), headers: authHeaders({ "Content-Type": "application/json" }),
body: JSON.stringify(body), body: JSON.stringify(body),
signal: opts?.signal, signal: opts?.signal,
@ -863,8 +924,7 @@ function webcrtSftpUploadOnce(
opts.signal.addEventListener("abort", onAbort, { once: true }); opts.signal.addEventListener("abort", onAbort, { once: true });
} }
xhr.open("POST", path); xhr.open("POST", path);
const tok = getAuthToken(); xhr.withCredentials = true;
if (tok) xhr.setRequestHeader("Authorization", `Bearer ${tok}`);
xhr.responseType = "text"; xhr.responseType = "text";
xhr.upload.onprogress = (ev) => { xhr.upload.onprogress = (ev) => {
if (!opts?.onProgress) return; if (!opts?.onProgress) return;
@ -1414,7 +1474,7 @@ export const downloadNeConfigSnapshot = async (
const path = const path =
`/v1/config-sync/snapshots/${encodeURIComponent(source)}/${encodeURIComponent(targetId)}` + `/v1/config-sync/snapshots/${encodeURIComponent(source)}/${encodeURIComponent(targetId)}` +
`/download?field=${encodeURIComponent(field)}`; `/download?field=${encodeURIComponent(field)}`;
const res = await fetch(path, { headers: authHeaders() }); const res = await fetch(path, { headers: authHeaders(), credentials: fetchCreds });
if (res.status === 401) { if (res.status === 401) {
handleUnauthorized(path); handleUnauthorized(path);
throw new Error("unauthorized"); throw new Error("unauthorized");