1032 lines
42 KiB
Python
1032 lines
42 KiB
Python
"""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",
|
||
)
|