deerflow-code/offline-backend-20260512/backend/app/gateway/routers/auth.py
2026-09-07 18:24:55 +08:00

1032 lines
42 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Authentication endpoints."""
import logging
import os
import time
from ipaddress import ip_address, ip_network
from typing import Literal
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response, status
from fastapi.security import OAuth2PasswordRequestForm
from pydantic import BaseModel, EmailStr, Field, field_validator
from app.gateway.auth import (
UserResponse,
create_access_token,
)
from app.gateway.auth.config import get_auth_config
from app.gateway.auth.errors import AuthErrorCode, AuthErrorResponse
from app.gateway.auth.models import User
from app.gateway.csrf_middleware import is_secure_request
from app.gateway.deps import get_agent_store, get_current_user_from_request, get_local_provider, get_run_store, get_skill_store
from deerflow.config import get_app_config
from deerflow.config.system_settings import load_system_settings
from deerflow.persistence.agents.base import AgentStore
from deerflow.persistence.skills.base import SkillStore
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/v1/auth", tags=["auth"])
# ── Request/Response Models ──────────────────────────────────────────────
class LoginResponse(BaseModel):
"""Response model for login — token only lives in HttpOnly cookie."""
expires_in: int # seconds
needs_setup: bool = False
# Top common-password blocklist. Drawn from the public SecLists "10k worst
# passwords" set, lowercased + length>=8 only (shorter ones already fail
# the min_length check). Kept tight on purpose: this is the **lower bound**
# defense, not a full HIBP / passlib check, and runs in-process per request.
_COMMON_PASSWORDS: frozenset[str] = frozenset(
{
"password",
"password1",
"password12",
"password123",
"password1234",
"12345678",
"123456789",
"1234567890",
"qwerty12",
"qwertyui",
"qwerty123",
"abc12345",
"abcd1234",
"iloveyou",
"letmein1",
"welcome1",
"welcome123",
"admin123",
"administrator",
"passw0rd",
"p@ssw0rd",
"monkey12",
"trustno1",
"sunshine",
"princess",
"football",
"baseball",
"superman",
"batman123",
"starwars",
"dragon123",
"master123",
"shadow12",
"michael1",
"jennifer",
"computer",
}
)
def _password_is_common(password: str) -> bool:
"""Case-insensitive blocklist check.
Lowercases the input so trivial mutations like ``Password`` /
``PASSWORD`` are also rejected. Does not normalize digit substitutions
(``p@ssw0rd`` is included as a literal entry instead) — keeping the
rule cheap and predictable.
"""
return password.lower() in _COMMON_PASSWORDS
def _validate_strong_password(value: str) -> str:
"""Pydantic field-validator body shared by Register + ChangePassword.
Constraint = function, not type-level mixin. The two request models
have no "is-a" relationship; they only share the password-strength
rule. Lifting it into a free function lets each model bind it via
``@field_validator(field_name)`` without inheritance gymnastics.
"""
if _password_is_common(value):
raise ValueError("Password is too common; choose a stronger password.")
return value
class RegisterRequest(BaseModel):
"""Request model for user registration."""
email: EmailStr
password: str = Field(..., min_length=8)
_strong_password = field_validator("password")(classmethod(lambda cls, v: _validate_strong_password(v)))
class UsernameLoginRequest(BaseModel):
"""Request model for passwordless username login.
Behavior:
- Existing user: direct login without password
- Missing user: auto-register and login without password
"""
username: str = Field(..., min_length=1, max_length=128)
class UsernameLoginResponse(BaseModel):
"""Response model for passwordless username login.
Returns the access token directly in the body so non-browser clients
can use it as a Bearer credential, while a session cookie is still set
for browser-based callers.
"""
access_token: str
token_type: str = "bearer"
expires_in: int
user_id: str
email: str
system_role: Literal["admin", "user"] = "user"
needs_setup: bool = False
created: bool = Field(default=False, description="True when the user was auto-registered by this call")
username: str | None = Field(
default=None,
description=(
"The login username this call resolved (token login: the value read "
"from getTokenInfo; username login: the supplied username). Used by "
"the frontend to append ?username= on task-workspace jump links; the "
"stored email local-part is lossy (lowercased/normalized)."
),
)
class ChangePasswordRequest(BaseModel):
"""Request model for password change (also handles setup flow)."""
current_password: str
new_password: str = Field(..., min_length=8)
new_email: EmailStr | None = None
_strong_password = field_validator("new_password")(classmethod(lambda cls, v: _validate_strong_password(v)))
class MessageResponse(BaseModel):
"""Generic message response."""
message: str
# ── Helpers ───────────────────────────────────────────────────────────────
def _normalize_username_local_part(username: str) -> str | None:
"""Return the canonical local-part for a username, or None if user typed an email.
Used both to build the synthetic email and to look up legacy accounts that
were registered with a different default domain. Returning ``None`` for
raw-email input tells callers to skip the domain-agnostic fallback —
if you typed ``alice@old.com``, you meant that specific row.
"""
raw = username.strip()
if not raw:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="用户名不能为空")
if "@" in raw:
return None
normalized = "".join(ch if (ch.isalnum() or ch in "._-") else "-" for ch in raw.lower()).strip("._-")
if not normalized:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="用户名不合法")
return normalized
def _username_to_email(username: str) -> str:
"""Map a frontend username to a valid email used by user storage.
- If username already looks like an email (`contains @`), use it directly.
- Otherwise generate `<normalized-username>@<domain>`.
"""
local = _normalize_username_local_part(username)
if local is None:
return username.strip().lower()
# Default mail domain for passwordless / username-only registrations.
# Kept overridable via env so individual deployments can swap it out, but the
# built-in default targets the cm.com tenant.
domain = os.getenv("DEER_FLOW_USERNAME_LOGIN_DOMAIN", "cm.com").strip().lower() or "cm.com"
return f"{local}@{domain}"
async def _sync_user_position(request: Request, user_id: str) -> None:
"""Best-effort: materialize the user's 岗位 (position) grants on login.
Ensures a user who has never been synced — e.g. a brand-new account that
inherits the ``is_default`` position — gets the default agents / tasks /
prompts loaded into their "我的" tables. Failures never block login.
"""
try:
sync = getattr(request.app.state, "position_sync", None)
if sync is not None:
await sync.sync_user(str(user_id))
except Exception:
logger.debug("position sync on login failed for user %s", user_id, exc_info=True)
def _is_admin_username(username: str) -> bool:
"""内置管理员用户名判定:用户名包含 ``admin``(如 ``admin``、``b1admin``)。
命中的用户名首次自动注册时直接落 ``system_role="admin"``;历史已注册但仍是
普通用户的,登录时自动提升为管理员(与 ``admin`` 账号原有行为一致)。判定
不区分大小写;显式邮箱登录按整串匹配(邮箱本地部分含 admin 同样命中)。
"""
return "admin" in (username or "").lower()
async def _get_or_create_user_by_username(username: str):
"""Resolve a username to a local user, auto-registering when absent.
Shared by ``/login/username`` and ``/login/token`` — both turn an
externally-supplied username into a local account and sign it in. Returns
``(user, created)`` where ``created`` is True only when this call inserted
the row. Mirrors the legacy domain-agnostic fallback and the admin-username
promotion so both entry points behave identically.
"""
provider = get_local_provider()
email = _username_to_email(username)
local_part = _normalize_username_local_part(username)
created = False
# A bare username is globally unique across synthetic email domains. Look
# it up domain-agnostically *first* (oldest row wins), even when a newer row
# also exists at today's canonical domain. Choosing the exact canonical
# row first made the same username alternate backend user_ids after a
# domain migration, which in turn made its existing threads inaccessible.
if local_part is not None:
user = await provider.find_user_by_email_local_part(local_part)
if user is not None and user.email != email:
logger.info("Matched username %r to stable legacy account %s", username, user.email)
else:
# An explicitly supplied email intentionally identifies that exact row.
user = await provider.get_user_by_email(email)
if user is None:
try:
role = "admin" if _is_admin_username(username) else "user"
# 黑名单模式:从未登录过的新用户(此处即首次自动注册)默认放行(approved=True)。
# 管理员可在白名单管理页对个别用户「取消放行」来封禁——仅当登录白名单开关打开时生效。
user = await provider.create_user(email=email, password=None, system_role=role, needs_setup=False, approved=True)
created = True
logger.info("Auto-registered passwordless user: %s", email)
except ValueError:
# Concurrent first login may create the same user in parallel.
user = await provider.get_user_by_email(email)
if user is None and local_part is not None:
user = await provider.find_user_by_email_local_part(local_part)
if user is None:
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="创建或加载用户失败")
elif _is_admin_username(username) and user.system_role != "admin":
# Historical accounts created before the username-based role assignment
# may still have system_role="user" — promote admin-username accounts
# (admin / b1admin / …) so both login entry points behave consistently.
user.system_role = "admin"
user = await provider.update_user(user)
logger.info("Promoted historical admin-username account %r to system_role=admin: %s", username, email)
return user, created
async def _lookup_user_by_username(username: str) -> User | None:
"""Resolve a username to an **existing** local user, without creating one.
Mirrors the lookup half of ``_get_or_create_user_by_username`` (canonical
``<username>@<domain>`` email + legacy local-part fallback) so the
per-account 登录口令 gate can read the target account's ``login_password``
before deciding whether to admit the login. Returns ``None`` for a username
that has never registered (the personal-password gate then simply does not
apply, and ``_get_or_create_user_by_username`` later auto-registers it).
"""
provider = get_local_provider()
email = _username_to_email(username)
local_part = _normalize_username_local_part(username)
if local_part is not None:
return await provider.find_user_by_email_local_part(local_part)
return await provider.get_user_by_email(email)
def enforce_user_approved(user) -> None:
"""登录白名单拦截(**黑名单模式**):开关打开且用户被「取消放行」时拒绝登录。
新账号首次登录即默认放行(``approved=True``,见 ``_get_or_create_user_by_username``
与 ``create_user``),所以从未在本系统登录过的用户可以直接使用;白名单开关只用来
封禁管理员显式「取消放行」(``approved=False``) 的个别用户。
规则(三条硬约束防锁死):
- 开关关(``login_whitelist.enabled`` 默认 False)→ 直接放行,人人可登录。
- admin → 永远放行(避免管理员把自己关在门外)。
- 其余:``approved`` 为 False(被管理员取消放行)→ 403 ``PENDING_APPROVAL``;
前端据此提示「等待管理员开通」,且不写任何登录态。
放在各登录端点拿到 user 之后、签发 JWT 之前调用。
"""
try:
enabled = load_system_settings().login_whitelist.enabled
except Exception:
# 读配置失败时绝不误伤(白名单是附加闸门,故障时回到「不拦截」)。
logger.warning("读取登录白名单配置失败,本次登录不拦截", exc_info=True)
return
if not enabled:
return
if getattr(user, "system_role", None) == "admin":
return
if not getattr(user, "approved", False):
raise HTTPException(
status_code=status.HTTP_403_FORBIDDEN,
detail={"code": "PENDING_APPROVAL", "message": "您的账号已被限制登录,请联系管理员开通"},
)
async def _fetch_username_from_token(token: str) -> str:
"""Resolve an upstream access token to a username via the configured endpoint.
Calls ``auth_login.token_login.token_info_url`` with the caller-supplied
token placed directly in the ``Authorization: Bearer <token>`` header — the
upstream reads the bearer token from there, not from a cookie. Expects
``{"state": "200", "data": {"username": ...}}`` and returns ``data.username``.
Any upstream failure, non-200 ``state``, or missing username raises HTTP 401.
"""
config = get_app_config().auth_login.token_login
# Offline test shortcut: a token listed in ``mock_tokens`` resolves to its
# mapped username without ever calling the (possibly unreachable) upstream
# ``token_info_url``. Lets token login be exercised end-to-end where the real
# getTokenInfo host can't be reached (external networks). Empty in prod.
mocked = config.mock_tokens.get(token)
if mocked and str(mocked).strip():
return str(mocked).strip()
if not config.token_info_url:
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="token 登录未配置(token_info_url 为空)")
# TLS verification for the (HTTPS) upstream. A custom CA bundle wins; else
# the boolean toggle. UAT/internal hosts often need verify_ssl=false because
# they serve self-signed or private-CA certs — without this the call dies
# with an SSLError and every token login silently 401s.
verify: bool | str = config.ca_cert_path if config.ca_cert_path else config.verify_ssl
try:
async with httpx.AsyncClient(timeout=config.timeout_seconds, verify=verify) as client:
upstream = await client.post(
config.token_info_url,
headers={
"Authorization": f"Bearer {token}",
"Content-Type": "application/x-www-form-urlencoded",
},
)
except httpx.HTTPError as exc:
# Surface the concrete failure (SSL handshake, DNS, timeout, ...) in logs
# so HTTPS/cert misconfiguration is diagnosable; client still sees 401.
logger.warning(
"Token login: upstream call to %s failed (verify=%r): %s: %s",
config.token_info_url,
verify,
type(exc).__name__,
exc,
)
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="token 校验失败").model_dump())
if upstream.status_code != 200:
logger.warning("Token login: upstream returned HTTP %s", upstream.status_code)
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="token 无效").model_dump())
try:
payload = upstream.json()
except ValueError:
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="token 接口返回格式异常").model_dump())
# Upstream signals success with the string state "200" (see API docs).
if str(payload.get("state")) != "200":
logger.warning("Token login: upstream state=%r msg=%r", payload.get("state"), payload.get("msg"))
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="token 无效或已过期").model_dump())
username = (payload.get("data") or {}).get("username")
if not username or not str(username).strip():
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="未能从 token 解析出用户名").model_dump())
return str(username).strip()
def _set_session_cookie(response: Response, token: str, request: Request) -> None:
"""Set the access_token HttpOnly cookie on the response."""
config = get_auth_config()
is_https = is_secure_request(request)
response.set_cookie(
key="access_token",
value=token,
httponly=True,
secure=is_https,
samesite="lax",
max_age=config.token_expiry_days * 24 * 3600 if is_https else None,
)
# ── Rate Limiting ────────────────────────────────────────────────────────
# In-process dict — not shared across workers.
#
# **Limitation**: with multi-worker deployments (e.g., gunicorn -w N), each
# worker maintains its own lockout table, so an attacker effectively gets
# N × _MAX_LOGIN_ATTEMPTS guesses before being locked out everywhere. For
# production multi-worker setups, replace this with a shared store (Redis,
# database-backed counter) to enforce a true per-IP limit.
_MAX_LOGIN_ATTEMPTS = 100
_LOCKOUT_SECONDS = 300 # 5 minutes
# ip → (fail_count, lock_until_timestamp)
_login_attempts: dict[str, tuple[int, float]] = {}
def _trusted_proxies() -> list:
"""Parse ``AUTH_TRUSTED_PROXIES`` env var into a list of ip_network objects.
Comma-separated CIDR or single-IP entries. Empty / unset = no proxy is
trusted (direct mode). Invalid entries are skipped with a logger warning.
Read live so env-var overrides take effect immediately and tests can
``monkeypatch.setenv`` without poking a module-level cache.
"""
raw = os.getenv("AUTH_TRUSTED_PROXIES", "").strip()
if not raw:
return []
nets = []
for entry in raw.split(","):
entry = entry.strip()
if not entry:
continue
try:
nets.append(ip_network(entry, strict=False))
except ValueError:
logger.warning("AUTH_TRUSTED_PROXIES: ignoring invalid entry %r", entry)
return nets
def _get_client_ip(request: Request) -> str:
"""Extract the real client IP for rate limiting.
Trust model:
- The TCP peer (``request.client.host``) is always the baseline. It is
whatever the kernel reports as the connecting socket — unforgeable
by the client itself.
- ``X-Real-IP`` is **only** honored if the TCP peer is in the
``AUTH_TRUSTED_PROXIES`` allowlist (set via env var, comma-separated
CIDR or single IPs). When set, the gateway is assumed to be behind a
reverse proxy (nginx, Cloudflare, ALB, …) that overwrites
``X-Real-IP`` with the original client address.
- With no ``AUTH_TRUSTED_PROXIES`` set, ``X-Real-IP`` is silently
ignored — closing the bypass where any client could rotate the
header to dodge per-IP rate limits in dev / direct-gateway mode.
``X-Forwarded-For`` is intentionally NOT used because it is naturally
client-controlled at the *first* hop and the trust chain is harder to
audit per-request.
"""
peer_host = request.client.host if request.client else None
trusted = _trusted_proxies()
if trusted and peer_host:
try:
peer_ip = ip_address(peer_host)
if any(peer_ip in net for net in trusted):
real_ip = request.headers.get("x-real-ip", "").strip()
if real_ip:
return real_ip
except ValueError:
# peer_host wasn't a parseable IP (e.g. "unknown") — fall through
pass
return peer_host or "unknown"
def _check_rate_limit(ip: str) -> None:
"""Raise 429 if the IP is currently locked out."""
record = _login_attempts.get(ip)
if record is None:
return
fail_count, lock_until = record
if fail_count >= _MAX_LOGIN_ATTEMPTS:
if time.time() < lock_until:
raise HTTPException(
status_code=429,
detail="登录尝试次数过多,请稍后再试",
)
del _login_attempts[ip]
_MAX_TRACKED_IPS = 10000
def _record_login_failure(ip: str) -> None:
"""Record a failed login attempt for the given IP."""
# Evict expired lockouts when dict grows too large
if len(_login_attempts) >= _MAX_TRACKED_IPS:
now = time.time()
expired = [k for k, (c, t) in _login_attempts.items() if c >= _MAX_LOGIN_ATTEMPTS and now >= t]
for k in expired:
del _login_attempts[k]
# If still too large, evict cheapest-to-lose half: below-threshold
# IPs (lock_until=0.0) sort first, then earliest-expiring lockouts.
if len(_login_attempts) >= _MAX_TRACKED_IPS:
by_time = sorted(_login_attempts.items(), key=lambda kv: kv[1][1])
for k, _ in by_time[: len(by_time) // 2]:
del _login_attempts[k]
record = _login_attempts.get(ip)
if record is None:
_login_attempts[ip] = (1, 0.0)
else:
new_count = record[0] + 1
lock_until = time.time() + _LOCKOUT_SECONDS if new_count >= _MAX_LOGIN_ATTEMPTS else 0.0
_login_attempts[ip] = (new_count, lock_until)
def _record_login_success(ip: str) -> None:
"""Clear failure counter for the given IP on successful login."""
_login_attempts.pop(ip, None)
# ── Endpoints ─────────────────────────────────────────────────────────────
@router.post("/login/local", response_model=LoginResponse)
async def login_local(
request: Request,
response: Response,
form_data: OAuth2PasswordRequestForm = Depends(),
):
"""Local email/password login."""
client_ip = _get_client_ip(request)
_check_rate_limit(client_ip)
user = await get_local_provider().authenticate({"email": form_data.username, "password": form_data.password})
if user is None:
_record_login_failure(client_ip)
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="邮箱或密码错误").model_dump(),
)
_record_login_success(client_ip)
enforce_user_approved(user)
token = create_access_token(str(user.id), token_version=user.token_version)
_set_session_cookie(response, token, request)
await _sync_user_position(request, user.id)
return LoginResponse(
expires_in=get_auth_config().token_expiry_days * 24 * 3600,
needs_setup=user.needs_setup,
)
@router.post("/login/username", response_model=UsernameLoginResponse)
async def login_by_username(
request: Request,
response: Response,
body: UsernameLoginRequest,
password: str | None = Query(default=None, description="Shared gate secret; required only when auth_login.username_login.require_password is on."),
):
"""Passwordless username login with auto-registration.
Rules:
- If mapped account exists, sign in directly (no password required).
- If account does not exist, create a regular user and sign in.
Password gates (checked in this order):
1. **Per-account 登录口令** — if an admin assigned this specific account a
personal password (``users.login_password``, set on the 白名单管理 page),
the caller must pass ``?password=`` exactly matching it. This takes
precedence over the shared gate and is the「临时给某账号开通+设新口令」flow.
2. **Shared gate** — otherwise, when
``auth_login.username_login.require_password`` is enabled, callers must
pass ``?password=`` matching the configured shared secret (default
``123ewq`` when unset). A single shared password; does not change which
account is logged in.
Returns the JWT access token in the response body and also sets it as
a session cookie for browser callers.
"""
username_login_config = get_app_config().auth_login.username_login
# Per-account password takes precedence: look the account up (no create) and,
# if it carries a personal 登录口令, require an exact match — ignoring the
# shared gate entirely for that account.
existing = await _lookup_user_by_username(body.username)
personal_password = (getattr(existing, "login_password", None) or "").strip() if existing else ""
if personal_password:
client_ip = _get_client_ip(request)
_check_rate_limit(client_ip)
if (password or "") != personal_password:
_record_login_failure(client_ip)
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="口令错误").model_dump(),
)
_record_login_success(client_ip)
elif username_login_config.require_password:
client_ip = _get_client_ip(request)
_check_rate_limit(client_ip)
if password != username_login_config.effective_password:
_record_login_failure(client_ip)
raise HTTPException(
status_code=status.HTTP_401_UNAUTHORIZED,
detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="口令错误").model_dump(),
)
_record_login_success(client_ip)
user, created = await _get_or_create_user_by_username(body.username)
enforce_user_approved(user)
token = create_access_token(str(user.id), token_version=user.token_version)
_set_session_cookie(response, token, request)
await _sync_user_position(request, user.id)
return UsernameLoginResponse(
access_token=token,
token_type="bearer",
expires_in=get_auth_config().token_expiry_days * 24 * 3600,
user_id=str(user.id),
email=user.email,
system_role=user.system_role,
needs_setup=user.needs_setup,
created=created,
username=body.username,
)
@router.post("/login/token", response_model=UsernameLoginResponse)
async def login_by_token(
request: Request,
response: Response,
token: str = Query(..., min_length=1, description="Upstream access token to exchange for a username."),
):
"""Token-exchange login.
Verifies an upstream access token by calling the configured
``token_info_url`` (see ``auth_login.token_login``), reads the resolved
username, then signs the matching local account in — auto-registering it
when the username is not yet in the database. Gated by
``auth_login.token_login.enabled``.
Returns the JWT access token in the response body and also sets it as a
session cookie for browser callers.
"""
config = get_app_config().auth_login.token_login
if not config.enabled:
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="token 登录未开启")
client_ip = _get_client_ip(request)
_check_rate_limit(client_ip)
try:
username = await _fetch_username_from_token(token)
except HTTPException as exc:
if exc.status_code == status.HTTP_401_UNAUTHORIZED:
_record_login_failure(client_ip)
raise
_record_login_success(client_ip)
user, created = await _get_or_create_user_by_username(username)
enforce_user_approved(user)
access_token = create_access_token(str(user.id), token_version=user.token_version)
_set_session_cookie(response, access_token, request)
await _sync_user_position(request, user.id)
return UsernameLoginResponse(
access_token=access_token,
token_type="bearer",
expires_in=get_auth_config().token_expiry_days * 24 * 3600,
user_id=str(user.id),
email=user.email,
system_role=user.system_role,
needs_setup=user.needs_setup,
created=created,
# The raw username resolved from getTokenInfo (pre email-normalization),
# so the frontend can append the exact ?username= on jump links.
username=username,
)
@router.post("/register", response_model=UserResponse, status_code=status.HTTP_201_CREATED)
async def register(request: Request, response: Response, body: RegisterRequest):
"""Register a new user account (always 'user' role).
Admin is auto-created on first boot. This endpoint creates regular users.
Auto-login by setting the session cookie.
"""
try:
user = await get_local_provider().create_user(email=body.email, password=body.password, system_role="user")
except ValueError:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=AuthErrorResponse(code=AuthErrorCode.EMAIL_ALREADY_EXISTS, message="Email already registered").model_dump(),
)
token = create_access_token(str(user.id), token_version=user.token_version)
_set_session_cookie(response, token, request)
await _sync_user_position(request, user.id)
return UserResponse(id=str(user.id), email=user.email, system_role=user.system_role)
@router.post("/logout", response_model=MessageResponse)
async def logout(request: Request, response: Response):
"""Logout current user by clearing the cookie."""
response.delete_cookie(key="access_token", secure=is_secure_request(request), samesite="lax")
return MessageResponse(message="Successfully logged out")
@router.post("/change-password", response_model=MessageResponse)
async def change_password(request: Request, response: Response, body: ChangePasswordRequest):
"""Change password for the currently authenticated user.
Also handles the first-boot setup flow:
- If new_email is provided, updates email (checks uniqueness)
- If user.needs_setup is True and new_email is given, clears needs_setup
- Always increments token_version to invalidate old sessions
- Re-issues session cookie with new token_version
"""
from app.gateway.auth.password import hash_password_async, verify_password_async
user = await get_current_user_from_request(request)
if user.password_hash is None:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="OAuth users cannot change password").model_dump())
if not await verify_password_async(body.current_password, user.password_hash):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=AuthErrorResponse(code=AuthErrorCode.INVALID_CREDENTIALS, message="Current password is incorrect").model_dump())
provider = get_local_provider()
# Update email if provided
if body.new_email is not None:
existing = await provider.get_user_by_email(body.new_email)
if existing and str(existing.id) != str(user.id):
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=AuthErrorResponse(code=AuthErrorCode.EMAIL_ALREADY_EXISTS, message="Email already in use").model_dump())
user.email = body.new_email
# Update password + bump version
user.password_hash = await hash_password_async(body.new_password)
user.token_version += 1
# Clear setup flag if this is the setup flow
if user.needs_setup and body.new_email is not None:
user.needs_setup = False
await provider.update_user(user)
# Re-issue cookie with new token_version
token = create_access_token(str(user.id), token_version=user.token_version)
_set_session_cookie(response, token, request)
return MessageResponse(message="Password changed successfully")
@router.get("/me", response_model=UserResponse)
async def get_me(request: Request):
"""Get current authenticated user info."""
user = await get_current_user_from_request(request)
return UserResponse(id=str(user.id), email=user.email, system_role=user.system_role, needs_setup=user.needs_setup)
_SETUP_STATUS_COOLDOWN: dict[str, float] = {}
_SETUP_STATUS_COOLDOWN_SECONDS = 60
_MAX_TRACKED_SETUP_STATUS_IPS = 10000
@router.get("/setup-status")
async def setup_status(request: Request):
"""Check if an admin account exists. Returns needs_setup=True when no admin exists."""
client_ip = _get_client_ip(request)
now = time.time()
last_check = _SETUP_STATUS_COOLDOWN.get(client_ip, 0)
elapsed = now - last_check
if elapsed < _SETUP_STATUS_COOLDOWN_SECONDS:
retry_after = max(1, int(_SETUP_STATUS_COOLDOWN_SECONDS - elapsed))
raise HTTPException(
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
detail="Setup status check is rate limited",
headers={"Retry-After": str(retry_after)},
)
# Evict stale entries when dict grows too large to bound memory usage.
if len(_SETUP_STATUS_COOLDOWN) >= _MAX_TRACKED_SETUP_STATUS_IPS:
cutoff = now - _SETUP_STATUS_COOLDOWN_SECONDS
stale = [k for k, t in _SETUP_STATUS_COOLDOWN.items() if t < cutoff]
for k in stale:
del _SETUP_STATUS_COOLDOWN[k]
# If still too large after evicting expired entries, remove oldest half.
if len(_SETUP_STATUS_COOLDOWN) >= _MAX_TRACKED_SETUP_STATUS_IPS:
by_time = sorted(_SETUP_STATUS_COOLDOWN.items(), key=lambda kv: kv[1])
for k, _ in by_time[: len(by_time) // 2]:
del _SETUP_STATUS_COOLDOWN[k]
_SETUP_STATUS_COOLDOWN[client_ip] = now
admin_count = await get_local_provider().count_admin_users()
return {"needs_setup": admin_count == 0}
class InitializeAdminRequest(BaseModel):
"""Request model for first-boot admin account creation."""
email: EmailStr
password: str = Field(..., min_length=8)
_strong_password = field_validator("password")(classmethod(lambda cls, v: _validate_strong_password(v)))
@router.post("/initialize", response_model=UserResponse, status_code=status.HTTP_201_CREATED)
async def initialize_admin(request: Request, response: Response, body: InitializeAdminRequest):
"""Create the first admin account on initial system setup.
Only callable when no admin exists. Returns 409 Conflict if an admin
already exists.
On success, the admin account is created with ``needs_setup=False`` and
the session cookie is set.
"""
admin_count = await get_local_provider().count_admin_users()
if admin_count > 0:
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=AuthErrorResponse(code=AuthErrorCode.SYSTEM_ALREADY_INITIALIZED, message="System already initialized").model_dump(),
)
try:
user = await get_local_provider().create_user(email=body.email, password=body.password, system_role="admin", needs_setup=False, approved=True)
except ValueError:
# DB unique-constraint race: another concurrent request beat us.
raise HTTPException(
status_code=status.HTTP_409_CONFLICT,
detail=AuthErrorResponse(code=AuthErrorCode.SYSTEM_ALREADY_INITIALIZED, message="System already initialized").model_dump(),
)
token = create_access_token(str(user.id), token_version=user.token_version)
_set_session_cookie(response, token, request)
return UserResponse(id=str(user.id), email=user.email, system_role=user.system_role)
# ── Admin Endpoints ───────────────────────────────────────────────────────
class UserListItem(BaseModel):
id: str
email: str
system_role: Literal["admin", "user"]
skill_count: int
agent_count: int
qa_count: int = 0 # 问答次数: number of runs (one run = one question)
approved: bool = True # 登录白名单放行标记(白名单管理页据此放行/取消)
login_password: str | None = None # 账号级登录口令(明文,管理员可见/可复制登录链接)
@router.get("/users", response_model=list[UserListItem])
async def list_all_users(
request: Request,
skill_store: SkillStore = Depends(get_skill_store),
agent_store: AgentStore = Depends(get_agent_store),
):
"""List all users with their skill, agent and Q&A counts. Admin only."""
user = await get_current_user_from_request(request)
if user.system_role != "admin":
raise HTTPException(status_code=403, detail="Admin only")
users = await get_local_provider().list_users()
# One GROUP BY query for everyone's 问答次数, so the list stays a single
# aggregate read instead of a per-user count. Best-effort: if the run store
# is unavailable the column simply renders as 0 rather than failing the list.
try:
qa_counts = await get_run_store(request).count_by_user()
except Exception:
logger.warning("Failed to load per-user run counts for user list", exc_info=True)
qa_counts = {}
result = []
for u in users:
uid = str(u.id)
result.append(
UserListItem(
id=uid,
email=u.email,
system_role=u.system_role,
skill_count=await skill_store.count_owned(uid),
agent_count=await agent_store.count_owned(uid),
qa_count=qa_counts.get(uid, 0),
approved=getattr(u, "approved", True),
login_password=getattr(u, "login_password", None),
)
)
return result
class UserApprovalRequest(BaseModel):
approved: bool = Field(..., description="True=放行该用户登录;False=取消放行(下次登录被拦)")
@router.put("/users/{user_id}/approval", response_model=UserListItem)
async def set_user_approval(user_id: str, body: UserApprovalRequest, request: Request) -> UserListItem:
"""放行 / 取消放行某个注册用户(登录白名单)。Admin only。
取消放行只拦该用户的下次登录(不动 ``token_version``,现有会话不强制下线)。
"""
admin = await get_current_user_from_request(request)
if admin.system_role != "admin":
raise HTTPException(status_code=403, detail="Admin only")
provider = get_local_provider()
target = await provider.get_user(user_id)
if target is None:
raise HTTPException(status_code=404, detail="用户不存在")
target.approved = bool(body.approved)
saved = await provider.update_user(target)
return UserListItem(
id=str(saved.id),
email=saved.email,
system_role=saved.system_role,
skill_count=0,
agent_count=0,
qa_count=0,
approved=saved.approved,
login_password=saved.login_password,
)
class UserLoginPasswordRequest(BaseModel):
password: str | None = Field(
default=None,
max_length=255,
description="设为非空 → 该账号用户名直登需带此明文口令;传 null/空字符串 → 清除该口令(回到共享口令门/免口令)。",
)
@router.put("/users/{user_id}/login-password", response_model=UserListItem)
async def set_user_login_password(user_id: str, body: UserLoginPasswordRequest, request: Request) -> UserListItem:
"""设置 / 清除某个账号的「登录口令」(明文)。Admin only。
设置后,该账号的用户名直登(``/login/username``)必须带匹配的 ``?password=``;
管理员据返回的明文口令拼出登录链接 ``…/login/<用户名>?password=xxxx`` 发给用户。
不动 ``token_version``,现有会话不受影响(仅影响下次登录)。
"""
admin = await get_current_user_from_request(request)
if admin.system_role != "admin":
raise HTTPException(status_code=403, detail="Admin only")
provider = get_local_provider()
target = await provider.get_user(user_id)
if target is None:
raise HTTPException(status_code=404, detail="用户不存在")
# 空字符串 / null / 纯空白 → 清除(None);否则存去掉首尾空白后的明文。
pwd = (body.password or "").strip()
target.login_password = pwd or None
saved = await provider.update_user(target)
return UserListItem(
id=str(saved.id),
email=saved.email,
system_role=saved.system_role,
skill_count=0,
agent_count=0,
qa_count=0,
approved=saved.approved,
login_password=saved.login_password,
)
# ── OAuth Endpoints (Future/Placeholder) ─────────────────────────────────
@router.get("/oauth/{provider}")
async def oauth_login(provider: str):
"""Initiate OAuth login flow.
Redirects to the OAuth provider's authorization URL.
Currently a placeholder - requires OAuth provider implementation.
"""
if provider not in ["github", "google"]:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported OAuth provider: {provider}",
)
raise HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail="OAuth login not yet implemented",
)
@router.get("/callback/{provider}")
async def oauth_callback(provider: str, code: str, state: str):
"""OAuth callback endpoint.
Handles the OAuth provider's callback after user authorization.
Currently a placeholder.
"""
raise HTTPException(
status_code=status.HTTP_501_NOT_IMPLEMENTED,
detail="OAuth callback not yet implemented",
)