refactor: 引入 safe_create_task 防止后台任务被 GC 回收,降低默认连接池和 worker 数量

- 新增 safe_create_task 统一替代裸 asyncio.create_task,通过全局集合持有 task 引用
- 默认 worker 数量从 4 降为 1,HTTP 连接池总预算从 800 降为 200
- 为 health_cache 和 affinity_manager 内存缓存增加上限淘汰机制
- MemoryCachePlugin 支持延迟启动清理任务
- gunicorn when_ready 增加 gc.collect() 并记录 post_worker_init RSS
This commit is contained in:
fawney19
2026-03-10 14:42:33 +08:00
parent 2d846b2c58
commit cfa5535f6e
22 changed files with 189 additions and 101 deletions

View File

@@ -49,6 +49,13 @@ ADMIN_PASSWORD=admin123456
# max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%) # max-requests-jitter 会自动设置为 MAX_REQUESTS/20 (5%)
# MAX_REQUESTS=4000 # MAX_REQUESTS=4000
# HTTP 连接池上限(默认总预算约 200按 worker 平分)
# 如果容器内存偏高,可继续下调;例如 2 worker 时设为 80-100
# HTTP_MAX_CONNECTIONS=100
# HTTP 保活连接数(默认约为 max_connections 的 30%
# HTTP_KEEPALIVE_CONNECTIONS=30
# API Key 前缀(默认 sk # API Key 前缀(默认 sk
# API_KEY_PREFIX=sk # API_KEY_PREFIX=sk

View File

@@ -1,7 +1,9 @@
# Gunicorn configuration file # Gunicorn configuration file
from __future__ import annotations
import gc import gc
import os import os
from typing import Any
# worker 心跳超时(秒):异步 worker 在此时间内必须向 arbiter 发送心跳 # worker 心跳超时(秒):异步 worker 在此时间内必须向 arbiter 发送心跳
# 对于 UvicornWorker事件循环偶发阻塞GC、同步 IO可能延迟心跳 # 对于 UvicornWorker事件循环偶发阻塞GC、同步 IO可能延迟心跳
@@ -13,21 +15,31 @@ timeout = int(os.getenv("GUNICORN_TIMEOUT", "300"))
graceful_timeout = int(os.getenv("GUNICORN_GRACEFUL_TIMEOUT", "120")) graceful_timeout = int(os.getenv("GUNICORN_GRACEFUL_TIMEOUT", "120"))
def when_ready(server): def _log_current_rss(log: Any, message: str) -> None:
"""
Called just after the server is started.
Freeze GC before forking workers to optimize Copy-on-Write memory sharing.
"""
gc.freeze()
server.log.info("GC frozen for Copy-on-Write optimization")
server.log.info(f"Objects in permanent generation: {gc.get_freeze_count()}")
def post_fork(server, worker):
try: try:
import resource import resource
rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss
server.log.info(f"Worker {worker.pid} RSS after fork: {rss} KB") log.info(f"{message}: {rss} KB")
except ImportError: except ImportError:
pass # Windows 不支持 resource 模块 pass
def when_ready(server: Any) -> None:
"""
Called just after the server is started.
Freeze GC before forking workers to optimize Copy-on-Write memory sharing.
"""
collected = gc.collect()
gc.freeze()
server.log.info(f"GC collected {collected} unreachable objects before freeze")
server.log.info("GC frozen for Copy-on-Write optimization")
server.log.info(f"Objects in permanent generation: {gc.get_freeze_count()}")
def post_fork(server: Any, worker: Any) -> None:
_log_current_rss(server.log, f"Worker {worker.pid} RSS after fork")
def post_worker_init(worker: Any) -> None:
_log_current_rss(worker.log, f"Worker {worker.pid} RSS after app init")

View File

@@ -43,6 +43,7 @@ from src.models.database import Provider, ProviderAPIKey, User
from src.services.provider.pool.config import parse_pool_config from src.services.provider.pool.config import parse_pool_config
from src.services.provider_keys.auth_type import OAUTH_AUTH_TYPES from src.services.provider_keys.auth_type import OAUTH_AUTH_TYPES
from src.services.scheduling.utils import release_db_connection_before_await from src.services.scheduling.utils import release_db_connection_before_await
from src.utils.async_utils import safe_create_task
from src.utils.auth_utils import require_admin from src.utils.auth_utils import require_admin
router = APIRouter(prefix="/api/admin/provider-oauth", tags=["Provider OAuth"]) router = APIRouter(prefix="/api/admin/provider-oauth", tags=["Provider OAuth"])
@@ -1300,7 +1301,7 @@ async def complete_provider_oauth(
) )
# 单个导入完成后,后台触发一次配额刷新 # 单个导入完成后,后台触发一次配额刷新
asyncio.create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)])) safe_create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
return response return response
@@ -1770,9 +1771,7 @@ async def import_refresh_token(
) )
# 单个导入完成后,后台触发一次配额刷新 # 单个导入完成后,后台触发一次配额刷新
asyncio.create_task( safe_create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)])
)
return response return response
@@ -1901,7 +1900,7 @@ async def import_refresh_token(
) )
# 单个导入完成后,后台触发一次配额刷新 # 单个导入完成后,后台触发一次配额刷新
asyncio.create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)])) safe_create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
return response return response
@@ -2315,9 +2314,7 @@ async def batch_import_oauth(
# 导入完成后,后台触发一次配额刷新 # 导入完成后,后台触发一次配额刷新
success_key_ids = _extract_success_key_ids(result) success_key_ids = _extract_success_key_ids(result)
if success_key_ids: if success_key_ids:
asyncio.create_task( safe_create_task(_refresh_quota_after_import(provider_id, provider_type, success_key_ids))
_refresh_quota_after_import(provider_id, provider_type, success_key_ids)
)
return result return result
@@ -3128,7 +3125,7 @@ async def device_poll(
# 单个导入完成后,后台触发一次配额刷新 # 单个导入完成后,后台触发一次配额刷新
provider_type = ProviderType.KIRO.value provider_type = ProviderType.KIRO.value
asyncio.create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)])) safe_create_task(_refresh_quota_after_import(provider_id, provider_type, [str(new_key.id)]))
return DevicePollResponse( return DevicePollResponse(
status="authorized", status="authorized",

View File

@@ -2044,10 +2044,11 @@ class AdminImportConfigAdapter(AdminApiAdapter):
import asyncio import asyncio
from src.services.model.fetch_scheduler import get_model_fetch_scheduler from src.services.model.fetch_scheduler import get_model_fetch_scheduler
from src.utils.async_utils import safe_create_task
scheduler = get_model_fetch_scheduler() scheduler = get_model_fetch_scheduler()
for key_id in keys_to_fetch: for key_id in keys_to_fetch:
asyncio.create_task(scheduler._fetch_models_for_key_by_id(key_id)) safe_create_task(scheduler._fetch_models_for_key_by_id(key_id))
except Exception as e: except Exception as e:
logger.error(f"触发模型获取失败: {e}") logger.error(f"触发模型获取失败: {e}")
# 不影响导入成功的返回 # 不影响导入成功的返回

View File

@@ -552,7 +552,9 @@ class BaseMessageHandler:
logger.warning(f"[{target_request_id}] 更新 Usage 状态为 streaming 失败: {e}") logger.warning(f"[{target_request_id}] 更新 Usage 状态为 streaming 失败: {e}")
# 创建后台任务,不阻塞当前流 # 创建后台任务,不阻塞当前流
asyncio.create_task(_do_update()) from src.utils.async_utils import safe_create_task
safe_create_task(_do_update())
def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None: def _update_usage_to_streaming_with_ctx(self, ctx: StreamContext) -> None:
"""更新 Usage 状态为 streaming同时更新 provider 相关信息 """更新 Usage 状态为 streaming同时更新 provider 相关信息
@@ -622,7 +624,9 @@ class BaseMessageHandler:
logger.warning(f"[{target_request_id}] 更新 Usage 状态为 streaming 失败: {e}") logger.warning(f"[{target_request_id}] 更新 Usage 状态为 streaming 失败: {e}")
# 创建后台任务,不阻塞当前流 # 创建后台任务,不阻塞当前流
asyncio.create_task(_do_update()) from src.utils.async_utils import safe_create_task
safe_create_task(_do_update())
def _log_request_error(self, message: str, error: Exception) -> None: def _log_request_error(self, message: str, error: Exception) -> None:
"""记录请求错误日志,对业务异常不打印堆栈 """记录请求错误日志,对业务异常不打印堆栈

View File

@@ -26,7 +26,7 @@ class Config:
self.port = int(os.getenv("PORT", "8084")) self.port = int(os.getenv("PORT", "8084"))
self.log_level = os.getenv("LOG_LEVEL", "INFO") self.log_level = os.getenv("LOG_LEVEL", "INFO")
self.worker_processes = int( self.worker_processes = int(
os.getenv("WEB_CONCURRENCY", os.getenv("GUNICORN_WORKERS", "4")) os.getenv("WEB_CONCURRENCY", os.getenv("GUNICORN_WORKERS", "1"))
) )
# PostgreSQL 连接池计算相关配置 # PostgreSQL 连接池计算相关配置
@@ -364,18 +364,18 @@ class Config:
- 单 Worker: 200 连接(适合开发/低负载) - 单 Worker: 200 连接(适合开发/低负载)
- 多 Worker: 按比例分配,确保总数不超过系统限制 - 多 Worker: 按比例分配,确保总数不超过系统限制
范围: 50 - 500 范围: 50 - 200
""" """
# 基础连接数:假设系统可用 socket 约 800 个用于 HTTP # 基础连接数:默认将总 HTTP 连接预算控制在 200
# (预留给 DB、Redis、内部服务等 # (预留给 DB、Redis、内部服务等
base_connections = 800 base_connections = 200
workers = max(self.worker_processes, 1) workers = max(self.worker_processes, 1)
# 每个 Worker 分配的连接数 # 每个 Worker 分配的连接数
per_worker = base_connections // workers per_worker = base_connections // workers
# 限制范围:最小 50保证基本并发最大 500避免资源耗尽 # 限制范围:最小 50保证基本并发最大 200控制内存与 socket 占用
return max(50, min(per_worker, 500)) return max(50, min(per_worker, 200))
def _auto_http_keepalive_connections(self) -> int: def _auto_http_keepalive_connections(self) -> int:
""" """

View File

@@ -36,8 +36,8 @@ class MemoryCachePlugin(CachePlugin):
try: try:
self._start_cleanup_task() self._start_cleanup_task()
except: except Exception:
pass # 忽略事件循环错误 pass # 事件循环尚未启动,将在首次 set 时延迟启动
def _start_cleanup_task(self) -> None: def _start_cleanup_task(self) -> None:
"""启动后台清理任务""" """启动后台清理任务"""
@@ -100,8 +100,17 @@ class MemoryCachePlugin(CachePlugin):
self._misses += 1 self._misses += 1
return None return None
def _ensure_cleanup_task(self) -> None:
"""确保清理任务已启动(延迟启动,在事件循环可用后调用)"""
if self._cleanup_task is None or self._cleanup_task.done():
try:
self._start_cleanup_task()
except Exception:
pass
async def set(self, key: str, value: Any, ttl: int | None = None) -> bool: async def set(self, key: str, value: Any, ttl: int | None = None) -> bool:
"""设置缓存值""" """设置缓存值"""
self._ensure_cleanup_task()
async with self._lock: async with self._lock:
# 检查大小限制 # 检查大小限制
if key not in self._cache: if key not in self._cache:

View File

@@ -189,22 +189,17 @@ class MonitorPlugin(BasePlugin):
labels["model"] = model labels["model"] = model
# 异步记录指标 # 异步记录指标
import asyncio from src.utils.async_utils import safe_create_task
try:
loop = asyncio.get_running_loop()
except RuntimeError:
return # 没有事件循环时跳过
# 请求计数 # 请求计数
loop.create_task(self.increment("http_requests_total", labels=labels)) safe_create_task(self.increment("http_requests_total", labels=labels))
# 请求延迟 # 请求延迟
loop.create_task(self.histogram("http_request_duration_seconds", duration, labels=labels)) safe_create_task(self.histogram("http_request_duration_seconds", duration, labels=labels))
# 错误计数 # 错误计数
if status_code >= 400: if status_code >= 400:
loop.create_task(self.increment("http_errors_total", labels=labels)) safe_create_task(self.increment("http_errors_total", labels=labels))
def record_token_usage( def record_token_usage(
self, self,
@@ -226,23 +221,18 @@ class MonitorPlugin(BasePlugin):
""" """
labels = {"provider": provider} labels = {"provider": provider}
import asyncio from src.utils.async_utils import safe_create_task
try:
loop = asyncio.get_running_loop()
except RuntimeError:
return # 没有事件循环时跳过
# Token计数 # Token计数
loop.create_task(self.increment("tokens_input_total", input_tokens, labels=labels)) safe_create_task(self.increment("tokens_input_total", input_tokens, labels=labels))
loop.create_task(self.increment("tokens_output_total", output_tokens, labels=labels)) safe_create_task(self.increment("tokens_output_total", output_tokens, labels=labels))
loop.create_task( safe_create_task(
self.increment("tokens_total", input_tokens + output_tokens, labels=labels) self.increment("tokens_total", input_tokens + output_tokens, labels=labels)
) )
# 费用 # 费用
if cost is not None: if cost is not None:
loop.create_task(self.increment("usage_cost_total", cost, labels=labels)) safe_create_task(self.increment("usage_cost_total", cost, labels=labels))
def configure(self, config: dict[str, Any]) -> Any: def configure(self, config: dict[str, Any]) -> Any:
""" """

View File

@@ -367,6 +367,8 @@ class EmailNotificationPlugin(NotificationPlugin):
def __del__(self) -> None: def __del__(self) -> None:
"""清理资源""" """清理资源"""
try: try:
asyncio.create_task(self.close()) from src.utils.async_utils import safe_create_task
except:
safe_create_task(self.close())
except Exception:
pass pass

View File

@@ -307,6 +307,8 @@ class WebhookNotificationPlugin(NotificationPlugin):
def __del__(self) -> None: def __del__(self) -> None:
"""清理资源""" """清理资源"""
try: try:
asyncio.create_task(self.close()) from src.utils.async_utils import safe_create_task
except:
safe_create_task(self.close())
except Exception:
pass pass

View File

@@ -4,8 +4,6 @@
from __future__ import annotations from __future__ import annotations
import asyncio
from sqlalchemy import and_, func from sqlalchemy import and_, func
from sqlalchemy.exc import IntegrityError from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -17,6 +15,7 @@ from src.models.database import Model, Provider
from src.services.cache.invalidation import get_cache_invalidation_service from src.services.cache.invalidation import get_cache_invalidation_service
from src.services.cache.model_cache import ModelCacheService from src.services.cache.model_cache import ModelCacheService
from src.services.cache.model_list_cache import invalidate_models_list_cache from src.services.cache.model_list_cache import invalidate_models_list_cache
from src.utils.async_utils import safe_create_task
class ModelService: class ModelService:
@@ -81,7 +80,7 @@ class ModelService:
# 清除 Redis 缓存(异步执行,不阻塞返回) # 清除 Redis 缓存(异步执行,不阻塞返回)
# 重要:新增模型可能需要清除 resolver 的 NOT_FOUND 负缓存global_model:resolve:* # 重要:新增模型可能需要清除 resolver 的 NOT_FOUND 负缓存global_model:resolve:*
# 否则请求链路在 TTL 内可能无法立刻解析到新模型。 # 否则请求链路在 TTL 内可能无法立刻解析到新模型。
asyncio.create_task( safe_create_task(
ModelCacheService.invalidate_model_cache( ModelCacheService.invalidate_model_cache(
model_id=model.id, model_id=model.id,
provider_id=model.provider_id, provider_id=model.provider_id,
@@ -97,7 +96,7 @@ class ModelService:
cache_service.on_model_changed(model.provider_id, model.global_model_id) cache_service.on_model_changed(model.provider_id, model.global_model_id)
# 清除 /v1/models 列表缓存 # 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache()) safe_create_task(invalidate_models_list_cache())
return model return model
@@ -199,7 +198,7 @@ class ModelService:
# 清除 Redis 缓存(异步执行,不阻塞返回) # 清除 Redis 缓存(异步执行,不阻塞返回)
# 先清除旧的映射缓存 # 先清除旧的映射缓存
asyncio.create_task( safe_create_task(
ModelCacheService.invalidate_model_cache( ModelCacheService.invalidate_model_cache(
model_id=model.id, model_id=model.id,
provider_id=model.provider_id, provider_id=model.provider_id,
@@ -214,7 +213,7 @@ class ModelService:
or model.provider_model_mappings != old_provider_model_mappings or model.provider_model_mappings != old_provider_model_mappings
or model.global_model_id != old_global_model_id or model.global_model_id != old_global_model_id
): ):
asyncio.create_task( safe_create_task(
ModelCacheService.invalidate_model_cache( ModelCacheService.invalidate_model_cache(
model_id=model.id, model_id=model.id,
provider_id=model.provider_id, provider_id=model.provider_id,
@@ -230,7 +229,7 @@ class ModelService:
cache_service.on_model_changed(model.provider_id, model.global_model_id) cache_service.on_model_changed(model.provider_id, model.global_model_id)
# 清除 /v1/models 列表缓存 # 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache()) safe_create_task(invalidate_models_list_cache())
logger.info( logger.info(
f"更新模型成功: id={model_id}, 最终 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}" f"更新模型成功: id={model_id}, 最终 supports_vision: {model.supports_vision}, supports_function_calling: {model.supports_function_calling}, supports_extended_thinking: {model.supports_extended_thinking}"
@@ -286,7 +285,7 @@ class ModelService:
db.commit() db.commit()
# 清除 Redis 缓存 # 清除 Redis 缓存
asyncio.create_task( safe_create_task(
ModelCacheService.invalidate_model_cache( ModelCacheService.invalidate_model_cache(
model_id=cache_info["model_id"], model_id=cache_info["model_id"],
provider_id=cache_info["provider_id"], provider_id=cache_info["provider_id"],
@@ -304,7 +303,7 @@ class ModelService:
) )
# 清除 /v1/models 列表缓存 # 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache()) safe_create_task(invalidate_models_list_cache())
logger.info( logger.info(
f"删除模型成功: id={model_id}, provider_model_name={cache_info['provider_model_name']}, " f"删除模型成功: id={model_id}, provider_model_name={cache_info['provider_model_name']}, "
@@ -327,7 +326,7 @@ class ModelService:
db.refresh(model) db.refresh(model)
# 清除 Redis 缓存 # 清除 Redis 缓存
asyncio.create_task( safe_create_task(
ModelCacheService.invalidate_model_cache( ModelCacheService.invalidate_model_cache(
model_id=model.id, model_id=model.id,
provider_id=model.provider_id, provider_id=model.provider_id,
@@ -343,7 +342,7 @@ class ModelService:
cache_service.on_model_changed(model.provider_id, model.global_model_id) cache_service.on_model_changed(model.provider_id, model.global_model_id)
# 清除 /v1/models 列表缓存 # 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache()) safe_create_task(invalidate_models_list_cache())
status = "可用" if is_available else "不可用" status = "可用" if is_available else "不可用"
logger.info(f"更新模型可用状态: id={model_id}, status={status}") logger.info(f"更新模型可用状态: id={model_id}, status={status}")
@@ -413,7 +412,7 @@ class ModelService:
# 清除 Redis 缓存(异步执行,不阻塞返回) # 清除 Redis 缓存(异步执行,不阻塞返回)
# 逐个清除 resolver 的映射缓存,避免 NOT_FOUND 负缓存阻塞新模型生效。 # 逐个清除 resolver 的映射缓存,避免 NOT_FOUND 负缓存阻塞新模型生效。
for model in created_models: for model in created_models:
asyncio.create_task( safe_create_task(
ModelCacheService.invalidate_model_cache( ModelCacheService.invalidate_model_cache(
model_id=model.id, model_id=model.id,
provider_id=model.provider_id, provider_id=model.provider_id,
@@ -428,7 +427,7 @@ class ModelService:
cache_service.on_model_changed(provider_id, created_models[0].global_model_id) cache_service.on_model_changed(provider_id, created_models[0].global_model_id)
# 清除 /v1/models 列表缓存 # 清除 /v1/models 列表缓存
asyncio.create_task(invalidate_models_list_cache()) safe_create_task(invalidate_models_list_cache())
except IntegrityError as e: except IntegrityError as e:
db.rollback() db.rollback()
logger.error(f"批量创建模型失败: {str(e)}") logger.error(f"批量创建模型失败: {str(e)}")

View File

@@ -305,7 +305,7 @@ def schedule_lazy_fingerprint_persist(key_id: str, fp: dict[str, str]) -> None:
_clear_pending_persist(key_id) _clear_pending_persist(key_id)
try: try:
loop = asyncio.get_running_loop() asyncio.get_running_loop()
except RuntimeError: except RuntimeError:
try: try:
_persist_fingerprint_if_missing_sync(key_id, payload) _persist_fingerprint_if_missing_sync(key_id, payload)
@@ -313,15 +313,9 @@ def schedule_lazy_fingerprint_persist(key_id: str, fp: dict[str, str]) -> None:
_clear_pending_persist(key_id) _clear_pending_persist(key_id)
return return
task = loop.create_task(_persist_async()) from src.utils.async_utils import safe_create_task
def _on_done(done_task: asyncio.Task[None]) -> None: safe_create_task(_persist_async())
try:
done_task.result()
except Exception as exc:
logger.debug("lazy fingerprint background task failed: {}", str(exc))
task.add_done_callback(_on_done)
def ensure_key_fingerprint( def ensure_key_fingerprint(

View File

@@ -12,6 +12,7 @@ import time
from typing import Any from typing import Any
_TTL_SECONDS = 30.0 _TTL_SECONDS = 30.0
_MAX_PROVIDERS = 500 # 最大缓存 provider 数量,防止无界增长
_LOCK = threading.Lock() _LOCK = threading.Lock()
_CACHE: dict[str, tuple[float, dict[str, float]]] = {} _CACHE: dict[str, tuple[float, dict[str, float]]] = {}
@@ -74,6 +75,14 @@ def get_health_scores(provider_id: str, keys: list[Any]) -> dict[str, float]:
with _LOCK: with _LOCK:
_CACHE[provider_id] = (now + _TTL_SECONDS, fresh) _CACHE[provider_id] = (now + _TTL_SECONDS, fresh)
# 超出上限时清理过期条目,仍超限则淘汰最旧条目
if len(_CACHE) > _MAX_PROVIDERS:
expired = [k for k, (exp, _) in _CACHE.items() if now >= exp]
for k in expired:
del _CACHE[k]
if len(_CACHE) > _MAX_PROVIDERS:
oldest_key = min(_CACHE, key=lambda k: _CACHE[k][0])
del _CACHE[oldest_key]
return dict(fresh) return dict(fresh)

View File

@@ -92,20 +92,14 @@ class _DeleteKeyResult:
def _run_async_with_fallback(coro: Any) -> None: def _run_async_with_fallback(coro: Any) -> None:
"""在同步上下文中执行异步任务(有事件循环则调度,无则阻塞执行)。""" """在同步上下文中执行异步任务(有事件循环则调度,无则阻塞执行)。"""
try: try:
loop = asyncio.get_running_loop() asyncio.get_running_loop()
except RuntimeError: except RuntimeError:
asyncio.run(coro) asyncio.run(coro)
return return
task = loop.create_task(coro) from src.utils.async_utils import safe_create_task
def _log_task_error(done_task: asyncio.Task[Any]) -> None: safe_create_task(coro)
try:
done_task.result()
except Exception as exc:
logger.warning("异步缓存失效任务执行失败: {}", exc)
task.add_done_callback(_log_task_error)
async def _invalidate_cache_after_clear_oauth_invalid(key_id: str) -> None: async def _invalidate_cache_after_clear_oauth_invalid(key_id: str) -> None:

View File

@@ -488,7 +488,9 @@ class ProviderOpsService:
# 有缓存,可选触发后台刷新 # 有缓存,可选触发后台刷新
if trigger_refresh: if trigger_refresh:
# 后台任务内部已处理异常并记录日志,无需额外回调 # 后台任务内部已处理异常并记录日志,无需额外回调
asyncio.create_task(self._refresh_balance_async(provider_id)) from src.utils.async_utils import safe_create_task
safe_create_task(self._refresh_balance_async(provider_id))
return cached return cached
# 没有缓存 # 没有缓存
@@ -499,7 +501,9 @@ class ProviderOpsService:
else: else:
# 仅触发异步刷新,立即返回 # 仅触发异步刷新,立即返回
logger.debug("余额缓存未命中,触发异步刷新: provider_id={}", provider_id) logger.debug("余额缓存未命中,触发异步刷新: provider_id={}", provider_id)
asyncio.create_task(self._refresh_balance_async(provider_id)) from src.utils.async_utils import safe_create_task
safe_create_task(self._refresh_balance_async(provider_id))
return ActionResult( return ActionResult(
status=ActionStatus.PENDING, status=ActionStatus.PENDING,
action_type=ProviderActionType.QUERY_BALANCE, action_type=ProviderActionType.QUERY_BALANCE,

View File

@@ -93,6 +93,7 @@ class CacheAffinityManager:
self.redis = redis_client self.redis = redis_client
self.default_ttl = default_ttl self.default_ttl = default_ttl
self._memory_store: dict[str, dict[str, Any]] = {} self._memory_store: dict[str, dict[str, Any]] = {}
self._memory_store_max_size: int = 5000 # 内存模式最大条目数
self._memory_lock: asyncio.Lock | None = None self._memory_lock: asyncio.Lock | None = None
# L1 缓存(即使使用 Redis 也启用,减少网络往返) # L1 缓存(即使使用 Redis 也启用,减少网络往返)
@@ -257,6 +258,12 @@ class CacheAffinityManager:
lock = self._get_memory_lock() lock = self._get_memory_lock()
async with lock: async with lock:
self._memory_store[cache_key] = dict(affinity_dict) self._memory_store[cache_key] = dict(affinity_dict)
# 超出上限时清理过期条目
if len(self._memory_store) > self._memory_store_max_size:
now = time.time()
expired = [k for k, v in self._memory_store.items() if now > v.get("expire_at", 0)]
for k in expired:
del self._memory_store[k]
await self._set_l1_entry(cache_key, affinity_dict) await self._set_l1_entry(cache_key, affinity_dict)
async def _delete_affinity_key(self, cache_key: str) -> None: async def _delete_affinity_key(self, cache_key: str) -> None:

View File

@@ -234,7 +234,9 @@ class MaintenanceScheduler:
# 启动时执行一次初始化任务 # 启动时执行一次初始化任务
if config.maintenance_startup_tasks_enabled: if config.maintenance_startup_tasks_enabled:
asyncio.create_task(self._run_startup_tasks()) from src.utils.async_utils import safe_create_task
safe_create_task(self._run_startup_tasks())
else: else:
logger.info("维护调度器启动任务已禁用MAINTENANCE_STARTUP_TASKS_ENABLED=false") logger.info("维护调度器启动任务已禁用MAINTENANCE_STARTUP_TASKS_ENABLED=false")

View File

@@ -4,7 +4,6 @@
from __future__ import annotations from __future__ import annotations
import asyncio
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any from typing import Any
@@ -16,6 +15,7 @@ from src.core.validators import EmailValidator, PasswordValidator, UsernameValid
from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, User, UserRole from src.models.database import ApiKey, GlobalModel, Model, Provider, Usage, User, UserRole
from src.services.cache.user_cache import UserCacheService from src.services.cache.user_cache import UserCacheService
from src.services.user.bulk_cleanup import batch_nullify_fk, pre_clean_api_key from src.services.user.bulk_cleanup import batch_nullify_fk, pre_clean_api_key
from src.utils.async_utils import safe_create_task
from src.utils.transaction_manager import retry_on_database_error, transactional from src.utils.transaction_manager import retry_on_database_error, transactional
@@ -247,7 +247,7 @@ class UserService:
db.refresh(user) db.refresh(user)
# 清除用户缓存 # 清除用户缓存
asyncio.create_task(UserCacheService.invalidate_user_cache(user.id, user.email)) safe_create_task(UserCacheService.invalidate_user_cache(user.id, user.email))
logger.debug(f"更新用户信息: {user.email} (ID: {user_id})") logger.debug(f"更新用户信息: {user.email} (ID: {user_id})")
return user return user
@@ -349,7 +349,7 @@ class UserService:
raise raise
# 清除用户缓存 # 清除用户缓存
asyncio.create_task(UserCacheService.invalidate_user_cache(user_id, email)) safe_create_task(UserCacheService.invalidate_user_cache(user_id, email))
logger.info(f"删除用户: {email} (ID: {user_id}), 同时删除 {api_key_count} 个API密钥") logger.info(f"删除用户: {email} (ID: {user_id}), 同时删除 {api_key_count} 个API密钥")
return True return True
@@ -394,7 +394,7 @@ class UserService:
user.updated_at = datetime.now(timezone.utc) user.updated_at = datetime.now(timezone.utc)
# 清除用户缓存 # 清除用户缓存
asyncio.create_task(UserCacheService.invalidate_user_cache(user.id, user.email)) safe_create_task(UserCacheService.invalidate_user_cache(user.id, user.email))
logger.info(f"密码更改成功: 用户ID {user_id}") logger.info(f"密码更改成功: 用户ID {user_id}")
return True, "密码更改成功" return True, "密码更改成功"

View File

@@ -4,6 +4,8 @@
提供在异步上下文中安全执行同步函数的工具,避免阻塞事件循环。 提供在异步上下文中安全执行同步函数的工具,避免阻塞事件循环。
""" """
from __future__ import annotations
import asyncio import asyncio
from collections.abc import Callable, Coroutine from collections.abc import Callable, Coroutine
from functools import partial, wraps from functools import partial, wraps
@@ -11,6 +13,28 @@ from typing import Any, TypeVar
T = TypeVar("T") T = TypeVar("T")
# 全局 task 引用集合,防止 fire-and-forget task 被 GC 回收
_background_tasks: set[asyncio.Task[Any]] = set()
def safe_create_task(coro: Coroutine[Any, Any, Any]) -> asyncio.Task[Any] | None:
"""创建后台 task 并持有引用,防止被 GC 回收。
用于 fire-and-forget 场景(如缓存失效、异步指标上报等),
替代裸 ``asyncio.create_task()`` 调用。
Returns:
创建的 Task 对象;若没有运行中的事件循环则返回 None。
"""
try:
loop = asyncio.get_running_loop()
except RuntimeError:
return None
task = loop.create_task(coro)
_background_tasks.add(task)
task.add_done_callback(_background_tasks.discard)
return task
async def run_in_executor(func: Callable[..., T], *args: Any, **kwargs: Any) -> T: async def run_in_executor(func: Callable[..., T], *args: Any, **kwargs: Any) -> T:
""" """

View File

@@ -1,6 +1,5 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import random import random
import time import time
from typing import Any from typing import Any
@@ -131,8 +130,6 @@ class PerfRecorder:
@staticmethod @staticmethod
def _create_task(coro: Any) -> None: def _create_task(coro: Any) -> None:
try: from src.utils.async_utils import safe_create_task
loop = asyncio.get_running_loop()
except RuntimeError: safe_create_task(coro)
return
loop.create_task(coro)

View File

@@ -141,7 +141,7 @@ def test_delete_user_precleans_large_tables_before_final_delete(
"src.services.user.service.UserCacheService.invalidate_user_cache", "src.services.user.service.UserCacheService.invalidate_user_cache",
invalidate_user_cache, invalidate_user_cache,
) )
monkeypatch.setattr("src.services.user.service.asyncio.create_task", create_task) monkeypatch.setattr("src.services.user.service.safe_create_task", create_task)
assert UserService.delete_user(db, "user-3") is True assert UserService.delete_user(db, "user-3") is True

View File

@@ -0,0 +1,34 @@
import pytest
from src.config.settings import Config
def test_config_defaults_to_single_worker_and_capped_http_pool(
monkeypatch: pytest.MonkeyPatch,
) -> None:
for key in (
"WEB_CONCURRENCY",
"GUNICORN_WORKERS",
"HTTP_MAX_CONNECTIONS",
"HTTP_KEEPALIVE_CONNECTIONS",
):
monkeypatch.delenv(key, raising=False)
cfg = Config()
assert cfg.worker_processes == 1
assert cfg.http_max_connections == 200
assert cfg.http_keepalive_connections == 60
def test_config_scales_http_pool_down_for_multi_worker(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("GUNICORN_WORKERS", "2")
monkeypatch.delenv("WEB_CONCURRENCY", raising=False)
monkeypatch.delenv("HTTP_MAX_CONNECTIONS", raising=False)
monkeypatch.delenv("HTTP_KEEPALIVE_CONNECTIONS", raising=False)
cfg = Config()
assert cfg.worker_processes == 2
assert cfg.http_max_connections == 100
assert cfg.http_keepalive_connections == 30