refactor(auth): 优化IP归属地处理逻辑,调整缓存与展示

1. 调整IP归属地缓存TTL从7天改为30天
2. 优化登录IP列展示,增加whitespace-nowrap类并加宽最小宽度
3. 新增异步补全OAuth/微信登录会话归属地逻辑
4. 重构IP归属地工具类,统一降级文案与查询决策逻辑
5. 修复部分类型检查与异常处理问题
This commit is contained in:
zhangtao
2026-08-16 02:57:19 +08:00
parent 61336936a0
commit 11160b05e4
6 changed files with 112 additions and 32 deletions
@@ -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,
+1 -1
View File
@@ -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 超时(秒)
# ================================================= # # ================================================= #
+64 -26
View File
@@ -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",
{ {