mirror of
https://github.com/fastapiadmin/FastapiAdmin.git
synced 2026-10-09 11:12:02 +00:00
refactor(auth): 优化IP归属地处理逻辑,调整缓存与展示
1. 调整IP归属地缓存TTL从7天改为30天 2. 优化登录IP列展示,增加whitespace-nowrap类并加宽最小宽度 3. 新增异步补全OAuth/微信登录会话归属地逻辑 4. 重构IP归属地工具类,统一降级文案与查询决策逻辑 5. 修复部分类型检查与异常处理问题
This commit is contained in:
@@ -169,6 +169,7 @@ async def oauth_login_redirect_controller(
|
|||||||
@AuthRouter.get("/oauth/{provider}/callback", summary="第三方OAuth回调", include_in_schema=False)
|
@AuthRouter.get("/oauth/{provider}/callback", summary="第三方OAuth回调", include_in_schema=False)
|
||||||
async def oauth_callback_controller(
|
async def oauth_callback_controller(
|
||||||
request: Request,
|
request: Request,
|
||||||
|
background_tasks: BackgroundTasks,
|
||||||
redis: Annotated[Redis, Depends(redis_getter)],
|
redis: Annotated[Redis, Depends(redis_getter)],
|
||||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||||
provider: Annotated[OAuthProvider, Path(description="wechat | qq | github | gitee")],
|
provider: Annotated[OAuthProvider, Path(description="wechat | qq | github | gitee")],
|
||||||
@@ -205,6 +206,7 @@ async def oauth_callback_controller(
|
|||||||
provider=provider,
|
provider=provider,
|
||||||
code=code,
|
code=code,
|
||||||
state=state,
|
state=state,
|
||||||
|
background_tasks=background_tasks,
|
||||||
)
|
)
|
||||||
success_url = oauth_service_frontend_redirect_from_token(fe, token)
|
success_url = oauth_service_frontend_redirect_from_token(fe, token)
|
||||||
return RedirectContentResponse(url=success_url, status_code=302)
|
return RedirectContentResponse(url=success_url, status_code=302)
|
||||||
@@ -224,6 +226,7 @@ async def wx_mini_login_controller(
|
|||||||
redis: Annotated[Redis, Depends(redis_getter)],
|
redis: Annotated[Redis, Depends(redis_getter)],
|
||||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||||
body: WxLoginSchema,
|
body: WxLoginSchema,
|
||||||
|
background_tasks: BackgroundTasks,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
"""微信小程序登录(code2Session)。
|
"""微信小程序登录(code2Session)。
|
||||||
|
|
||||||
@@ -252,6 +255,7 @@ async def wx_mini_login_controller(
|
|||||||
redis=redis,
|
redis=redis,
|
||||||
user=user,
|
user=user,
|
||||||
login_type="wx_mini",
|
login_type="wx_mini",
|
||||||
|
background_tasks=background_tasks,
|
||||||
)
|
)
|
||||||
|
|
||||||
user_info = {
|
user_info = {
|
||||||
@@ -282,6 +286,7 @@ async def wx_mini_phone_login_controller(
|
|||||||
redis: Annotated[Redis, Depends(redis_getter)],
|
redis: Annotated[Redis, Depends(redis_getter)],
|
||||||
db: Annotated[AsyncSession, Depends(db_getter)],
|
db: Annotated[AsyncSession, Depends(db_getter)],
|
||||||
body: WxPhoneLoginSchema,
|
body: WxPhoneLoginSchema,
|
||||||
|
background_tasks: BackgroundTasks,
|
||||||
) -> JSONResponse:
|
) -> JSONResponse:
|
||||||
"""微信小程序手机号快速登录。
|
"""微信小程序手机号快速登录。
|
||||||
|
|
||||||
@@ -335,6 +340,7 @@ async def wx_mini_phone_login_controller(
|
|||||||
redis=redis,
|
redis=redis,
|
||||||
user=user,
|
user=user,
|
||||||
login_type="wx_mini_phone",
|
login_type="wx_mini_phone",
|
||||||
|
background_tasks=background_tasks,
|
||||||
)
|
)
|
||||||
|
|
||||||
user_info = {
|
user_info = {
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from typing import Any, Literal
|
|||||||
from urllib.parse import quote, urlencode
|
from urllib.parse import quote, urlencode
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import Request
|
from fastapi import BackgroundTasks, Request
|
||||||
from redis.asyncio.client import Redis
|
from redis.asyncio.client import Redis
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
@@ -347,6 +347,7 @@ async def complete_oauth_login(
|
|||||||
provider: OAuthProvider,
|
provider: OAuthProvider,
|
||||||
code: str,
|
code: str,
|
||||||
state: str,
|
state: str,
|
||||||
|
background_tasks: BackgroundTasks | None = None,
|
||||||
) -> tuple[JWTOutSchema, str]:
|
) -> tuple[JWTOutSchema, str]:
|
||||||
rc = RedisCURD(redis)
|
rc = RedisCURD(redis)
|
||||||
raw = await rc.get(f"{STATE_PREFIX}{state}")
|
raw = await rc.get(f"{STATE_PREFIX}{state}")
|
||||||
@@ -394,7 +395,7 @@ async def complete_oauth_login(
|
|||||||
raise CustomException(msg="用户不存在")
|
raise CustomException(msg="用户不存在")
|
||||||
|
|
||||||
login_type = f"oauth_{provider}"
|
login_type = f"oauth_{provider}"
|
||||||
token = await LoginService.create_token(request=request, redis=redis, user=user, login_type=login_type)
|
token = await LoginService.create_token(request=request, redis=redis, user=user, login_type=login_type, background_tasks=background_tasks)
|
||||||
return token, frontend
|
return token, frontend
|
||||||
finally:
|
finally:
|
||||||
await rc.delete(f"{STATE_PREFIX}{state}")
|
await rc.delete(f"{STATE_PREFIX}{state}")
|
||||||
|
|||||||
@@ -84,6 +84,30 @@ async def _async_fill_login_location(redis, login_log_id: int, ip: str | None) -
|
|||||||
logger.warning(f"异步补全登录归属地失败: {e}")
|
logger.warning(f"异步补全登录归属地失败: {e}")
|
||||||
|
|
||||||
|
|
||||||
|
async def _async_fill_session_location(redis, session_id: str, ip: str | None) -> None:
|
||||||
|
"""后台异步补全会话缓存中的归属地(微信/OAuth 登录不写登录日志,需单独更新 Redis 会话)。"""
|
||||||
|
if not ip:
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
location = await IpLocalUtil.resolve_location_async(redis, ip)
|
||||||
|
logger.info(f"异步解析IP归属地结果: ip={ip}, session_id={session_id}, location={location}")
|
||||||
|
if location == "归属地查询中" or not location:
|
||||||
|
return
|
||||||
|
key = f"{RedisInitKeyConfig.USER_SESSION.key}:{session_id}"
|
||||||
|
raw = await RedisCURD(redis).get(key)
|
||||||
|
if not raw:
|
||||||
|
return
|
||||||
|
if isinstance(raw, bytes):
|
||||||
|
raw = raw.decode("utf-8")
|
||||||
|
session_dict = json.loads(raw)
|
||||||
|
session_dict["login_location"] = location
|
||||||
|
ttl = await RedisCURD(redis).ttl(key)
|
||||||
|
await RedisCURD(redis).set(key, json.dumps(session_dict, default=str), expire=max(int(ttl), 1))
|
||||||
|
logger.info(f"会话归属地已更新: session_id={session_id}, location={location}")
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"异步补全会话归属地失败: {e}")
|
||||||
|
|
||||||
|
|
||||||
class LoginService:
|
class LoginService:
|
||||||
"""登录认证服务"""
|
"""登录认证服务"""
|
||||||
|
|
||||||
@@ -275,7 +299,14 @@ class LoginService:
|
|||||||
}
|
}
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def create_token(cls, request: Request, redis: Redis, user: UserModel, login_type: str) -> JWTOutSchema:
|
async def create_token(
|
||||||
|
cls,
|
||||||
|
request: Request,
|
||||||
|
redis: Redis,
|
||||||
|
user: UserModel,
|
||||||
|
login_type: str,
|
||||||
|
background_tasks: BackgroundTasks | None = None,
|
||||||
|
) -> JWTOutSchema:
|
||||||
"""创建访问令牌和刷新令牌"""
|
"""创建访问令牌和刷新令牌"""
|
||||||
session_id = str(uuid.uuid4())
|
session_id = str(uuid.uuid4())
|
||||||
ua_result = ua_parser.parse(request.headers.get("user-agent") or "")
|
ua_result = ua_parser.parse(request.headers.get("user-agent") or "")
|
||||||
@@ -336,6 +367,10 @@ class LoginService:
|
|||||||
expire=int(refresh_expires.total_seconds()),
|
expire=int(refresh_expires.total_seconds()),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# 归属地为待解析时后台补全会话中的 login_location(微信/OAuth 登录路径)
|
||||||
|
if background_tasks and login_location == "归属地查询中":
|
||||||
|
background_tasks.add_task(_async_fill_session_location, redis, session_id, request_ip)
|
||||||
|
|
||||||
return JWTOutSchema(
|
return JWTOutSchema(
|
||||||
access_token=access_token,
|
access_token=access_token,
|
||||||
refresh_token=refresh_token,
|
refresh_token=refresh_token,
|
||||||
|
|||||||
@@ -144,7 +144,7 @@ class Settings(BaseSettings):
|
|||||||
# ================================================= #
|
# ================================================= #
|
||||||
HTTPX_DEFAULT_TIMEOUT: float = 10.0 # 对外 HTTP 请求默认超时(秒)
|
HTTPX_DEFAULT_TIMEOUT: float = 10.0 # 对外 HTTP 请求默认超时(秒)
|
||||||
IP_LOCATION_ENABLE: bool = True # 是否启用 IP 归属地查询(登录时对外发起 HTTP 请求)
|
IP_LOCATION_ENABLE: bool = True # 是否启用 IP 归属地查询(登录时对外发起 HTTP 请求)
|
||||||
IP_LOCATION_CACHE_TTL: int = 604800 # IP 归属地缓存时间(秒,默认 7 天)
|
IP_LOCATION_CACHE_TTL: int = 2592000 # IP 归属地缓存时间(秒,默认 30 天)
|
||||||
IP_LOCATION_QUERY_TIMEOUT: float = 3.0 # IP 归属地查询单次 HTTP 超时(秒)
|
IP_LOCATION_QUERY_TIMEOUT: float = 3.0 # IP 归属地查询单次 HTTP 超时(秒)
|
||||||
|
|
||||||
# ================================================= #
|
# ================================================= #
|
||||||
|
|||||||
@@ -16,6 +16,12 @@ _IP_CACHE_TTL: int = settings.IP_LOCATION_CACHE_TTL
|
|||||||
# 硬超时(秒),避免外网查询阻塞主流程
|
# 硬超时(秒),避免外网查询阻塞主流程
|
||||||
_IP_QUERY_TIMEOUT: float = settings.IP_LOCATION_QUERY_TIMEOUT
|
_IP_QUERY_TIMEOUT: float = settings.IP_LOCATION_QUERY_TIMEOUT
|
||||||
|
|
||||||
|
# 归属地降级返回文案(登录日志/会话共用,调用方按字面量比较)
|
||||||
|
LOCATION_INTRANET = "内网IP"
|
||||||
|
LOCATION_DISABLED = "未解析(已关闭归属地查询)"
|
||||||
|
LOCATION_PENDING = "归属地查询中"
|
||||||
|
LOCATION_UNKNOWN = "未知"
|
||||||
|
|
||||||
|
|
||||||
def get_client_ip(request: Request) -> str:
|
def get_client_ip(request: Request) -> str:
|
||||||
"""从请求中提取客户端真实 IP(优先取反向代理透传的头部,返回空字串表示无法识别)。"""
|
"""从请求中提取客户端真实 IP(优先取反向代理透传的头部,返回空字串表示无法识别)。"""
|
||||||
@@ -31,10 +37,12 @@ def get_client_ip(request: Request) -> str:
|
|||||||
|
|
||||||
|
|
||||||
class IpLocalUtil:
|
class IpLocalUtil:
|
||||||
"""获取 IP 归属地工具类(带 Redis 缓存、硬超时、降级)。"""
|
"""获取 IP 归属地工具类(带 Redis 缓存、查询决策、硬超时、降级)。"""
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def is_valid_ip(cls, ip: str) -> bool:
|
def is_valid_ip(cls, ip: str | None) -> bool:
|
||||||
|
if not ip:
|
||||||
|
return False
|
||||||
try:
|
try:
|
||||||
ipaddress.ip_address(ip)
|
ipaddress.ip_address(ip)
|
||||||
return True
|
return True
|
||||||
@@ -42,11 +50,18 @@ class IpLocalUtil:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def is_private_ip(cls, ip: str) -> bool:
|
def is_private_ip(cls, ip: str | None) -> bool:
|
||||||
|
"""判断是否为非公网地址(内网/回环/链路本地/保留段)。
|
||||||
|
|
||||||
|
这类地址无法通过外网归属地 API 解析,用 ``is_global`` 取反判断最稳妥。
|
||||||
|
"""
|
||||||
|
if not ip:
|
||||||
|
return False
|
||||||
try:
|
try:
|
||||||
return ipaddress.ip_address(ip).is_private
|
addr = ipaddress.ip_address(ip)
|
||||||
except ValueError:
|
except ValueError:
|
||||||
return False
|
return False
|
||||||
|
return not addr.is_global
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def _is_location_enabled(cls, redis) -> bool:
|
async def _is_location_enabled(cls, redis) -> bool:
|
||||||
@@ -60,44 +75,67 @@ class IpLocalUtil:
|
|||||||
payload = json.loads(raw)
|
payload = json.loads(raw)
|
||||||
cv = payload.get("config_value", "off")
|
cv = payload.get("config_value", "off")
|
||||||
return cv in (True, "true", "1", "yes", "on")
|
return cv in (True, "true", "1", "yes", "on")
|
||||||
except (json.JSONDecodeError, TypeError, Exception):
|
except Exception:
|
||||||
pass
|
pass
|
||||||
return False
|
return False
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def resolve_location_for_log(cls, redis, ip: str | None) -> str | None:
|
async def _should_query(cls, redis, ip: str | None) -> bool:
|
||||||
"""登录日志写入入口:仅返回可同步获取的值(内网/缓存/降级),
|
"""判断 IP 是否需要发起外网归属地查询(前置条件 + 开关 + IP 类型)。
|
||||||
|
|
||||||
外网查询由后台任务异步执行(见 ``resolve_location_async``)。
|
不查询的情形(短路返回 False):
|
||||||
|
1. IP 为空或非法 —— 无法查询
|
||||||
|
2. Redis 不可用 —— 读不到开关与缓存,直接降级
|
||||||
|
3. 开关关闭(ip_location_enable != on)—— 功能总闸
|
||||||
|
4. 内网/回环/保留地址 —— 外网 API 无法解析
|
||||||
|
|
||||||
|
缓存命中与否由调用方处理(命中直接返回,未命中才需要查)。
|
||||||
"""
|
"""
|
||||||
if not ip:
|
if not cls.is_valid_ip(ip):
|
||||||
return None
|
return False
|
||||||
|
if not redis:
|
||||||
|
return False
|
||||||
if not await cls._is_location_enabled(redis):
|
if not await cls._is_location_enabled(redis):
|
||||||
return "内网IP" if cls.is_private_ip(ip) else "未解析(已关闭归属地查询)"
|
return False
|
||||||
if cls.is_private_ip(ip):
|
if cls.is_private_ip(ip):
|
||||||
return "内网IP"
|
return False
|
||||||
|
return True
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def resolve_location_for_log(cls, redis, ip: str | None) -> str | None:
|
||||||
|
"""登录日志写入入口:仅返回可同步获取的值(开关/内网/缓存)。
|
||||||
|
|
||||||
|
外网查询由后台任务异步执行(见 ``resolve_location_async``),
|
||||||
|
此处只做查询决策与缓存读取,不发起任何外网请求。
|
||||||
|
"""
|
||||||
|
if not cls.is_valid_ip(ip):
|
||||||
|
return None
|
||||||
|
assert ip is not None # is_valid_ip 已确保非空
|
||||||
|
if not await cls._should_query(redis, ip):
|
||||||
|
return LOCATION_INTRANET if cls.is_private_ip(ip) else LOCATION_DISABLED
|
||||||
if redis:
|
if redis:
|
||||||
cached = await cls._cache_get(redis, ip)
|
cached = await cls._cache_get(redis, ip)
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
return cached
|
return cached
|
||||||
return "归属地查询中"
|
return LOCATION_PENDING
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def resolve_location_async(cls, redis, ip: str) -> str:
|
async def resolve_location_async(cls, redis, ip: str) -> str:
|
||||||
"""异步查询归属地(含缓存、降级、硬超时)。"""
|
"""异步查询归属地(含决策、缓存、降级、硬超时)。仅供后台任务调用。
|
||||||
|
|
||||||
|
仅查询成功的结果写入缓存池(IP 池),失败(未知)不写缓存,
|
||||||
|
下次登录仍会重试,避免缓存池被无效值污染。
|
||||||
|
"""
|
||||||
if not cls.is_valid_ip(ip):
|
if not cls.is_valid_ip(ip):
|
||||||
return "未知"
|
return LOCATION_UNKNOWN
|
||||||
if not await cls._is_location_enabled(redis):
|
if not await cls._should_query(redis, ip):
|
||||||
return "未解析(已关闭归属地查询)"
|
return LOCATION_INTRANET if cls.is_private_ip(ip) else LOCATION_DISABLED
|
||||||
if cls.is_private_ip(ip):
|
|
||||||
return "内网IP"
|
|
||||||
|
|
||||||
cached = await cls._cache_get(redis, ip) if redis else None
|
|
||||||
if cached is not None:
|
|
||||||
return cached
|
|
||||||
|
|
||||||
result = await cls._query_with_timeout(ip)
|
|
||||||
if redis:
|
if redis:
|
||||||
|
cached = await cls._cache_get(redis, ip)
|
||||||
|
if cached is not None:
|
||||||
|
return cached
|
||||||
|
result = await cls._query_with_timeout(ip)
|
||||||
|
if redis and result != LOCATION_UNKNOWN:
|
||||||
await cls._cache_set(redis, ip, result)
|
await cls._cache_set(redis, ip, result)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
@@ -119,7 +157,7 @@ class IpLocalUtil:
|
|||||||
return location
|
return location
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"IP 归属地 API 失败: {url} - {e}")
|
logger.warning(f"IP 归属地 API 失败: {url} - {e}")
|
||||||
return "未知"
|
return LOCATION_UNKNOWN
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _parse_ipapi(data: dict) -> str | None:
|
def _parse_ipapi(data: dict) -> str | None:
|
||||||
|
|||||||
@@ -584,10 +584,10 @@ const {
|
|||||||
{
|
{
|
||||||
prop: "login_ip",
|
prop: "login_ip",
|
||||||
label: "登录IP",
|
label: "登录IP",
|
||||||
minWidth: 140,
|
minWidth: 190,
|
||||||
formatter: (row: LoginLogTable) =>
|
formatter: (row: LoginLogTable) =>
|
||||||
row.login_ip
|
row.login_ip
|
||||||
? h("span", { class: "inline-flex items-center gap-0.5" }, [
|
? h("span", { class: "inline-flex items-center gap-0.5 whitespace-nowrap" }, [
|
||||||
h(
|
h(
|
||||||
"span",
|
"span",
|
||||||
{
|
{
|
||||||
|
|||||||
Reference in New Issue
Block a user