mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-02 17:30:23 +08:00
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:
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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",
|
||||||
|
|||||||
@@ -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}")
|
||||||
# 不影响导入成功的返回
|
# 不影响导入成功的返回
|
||||||
|
|||||||
@@ -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:
|
||||||
"""记录请求错误日志,对业务异常不打印堆栈
|
"""记录请求错误日志,对业务异常不打印堆栈
|
||||||
|
|||||||
@@ -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:
|
||||||
"""
|
"""
|
||||||
|
|||||||
13
src/plugins/cache/memory.py
vendored
13
src/plugins/cache/memory.py
vendored
@@ -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:
|
||||||
|
|||||||
@@ -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:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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)}")
|
||||||
|
|||||||
@@ -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(
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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)")
|
||||||
|
|
||||||
|
|||||||
@@ -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, "密码更改成功"
|
||||||
|
|||||||
@@ -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:
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -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)
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
34
tests/unit/test_runtime_http_pool_defaults.py
Normal file
34
tests/unit/test_runtime_http_pool_defaults.py
Normal 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
|
||||||
Reference in New Issue
Block a user