refactor: 限制流式文本收集内存增长,降低默认连接池和缓存上限

- StreamContext.append_text 增加 16KB 上限,超出后仅计数不存储,
  避免长流式响应导致内存持续增长;token 估算改用 collected_text_length
- 降低 DB 连接池上限 (30->15) 和 HTTP 连接池上限 (200->100)
- tiktoken 编码器缓存从 32 缩减到 4(实际编码种类只有几种)
- dev.sh 添加开发环境低配连接池默认值,uvicorn 热重载仅监视 src 目录
This commit is contained in:
fawney19
2026-03-10 15:33:46 +08:00
parent cfa5535f6e
commit 6ec8df97e8
8 changed files with 68 additions and 24 deletions

16
dev.sh
View File

@@ -11,10 +11,16 @@ set +a
export DATABASE_URL="postgresql://${DB_USER:-postgres}:${DB_PASSWORD}@${DB_HOST:-localhost}:${DB_PORT:-5432}/${DB_NAME:-aether}" export DATABASE_URL="postgresql://${DB_USER:-postgres}:${DB_PASSWORD}@${DB_HOST:-localhost}:${DB_PORT:-5432}/${DB_NAME:-aether}"
export REDIS_URL=redis://:${REDIS_PASSWORD}@${REDIS_HOST:-localhost}:${REDIS_PORT:-6379}/0 export REDIS_URL=redis://:${REDIS_PASSWORD}@${REDIS_HOST:-localhost}:${REDIS_PORT:-6379}/0
# 启动 uvicorn热重载模式 # 开发环境连接池低配(节省内存
echo "🚀 启动本地开发服务器..." export DB_POOL_SIZE=${DB_POOL_SIZE:-5}
echo "📍 后端地址: http://localhost:8084" export DB_MAX_OVERFLOW=${DB_MAX_OVERFLOW:-5}
echo "📊 数据库: ${DATABASE_URL}" export HTTP_MAX_CONNECTIONS=${HTTP_MAX_CONNECTIONS:-20}
export HTTP_KEEPALIVE_CONNECTIONS=${HTTP_KEEPALIVE_CONNECTIONS:-5}
# 启动 uvicorn热重载模式只监视 src 目录)
echo "=> 启动本地开发服务器..."
echo "=> 后端地址: http://localhost:8084"
echo "=> 数据库: ${DATABASE_URL}"
echo "" echo ""
uv run uvicorn src.main:app --reload --port 8084 uv run uvicorn src.main:app --reload --reload-dir src --port 8084

View File

@@ -91,9 +91,8 @@ class CliPrefetchMixin:
估算采用 ~4 字符/token 的保守比例。 估算采用 ~4 字符/token 的保守比例。
""" """
# 输出 tokens从已收集的文本估算 # 输出 tokens从已收集的文本估算
collected = ctx.collected_text if ctx.collected_text_length > 0:
if collected: ctx.output_tokens = max(1, ctx.collected_text_length // 4)
ctx.output_tokens = max(1, len(collected) // 4)
# 输入 tokens从请求体文本内容估算 # 输入 tokens从请求体文本内容估算
try: try:

View File

@@ -46,6 +46,9 @@ def is_format_converted(
) )
_MAX_COLLECTED_TEXT_CHARS = 16 * 1024
@dataclass @dataclass
class StreamContext: class StreamContext:
""" """
@@ -91,6 +94,8 @@ class StreamContext:
# 响应内容 # 响应内容
_collected_text_parts: list[str] = field(default_factory=list, repr=False) _collected_text_parts: list[str] = field(default_factory=list, repr=False)
_collected_text_chars: int = field(default=0, repr=False)
_stored_collected_text_chars: int = field(default=0, repr=False)
response_id: str | None = None response_id: str | None = None
final_usage: dict[str, Any] | None = None final_usage: dict[str, Any] | None = None
final_response: dict[str, Any] | None = None final_response: dict[str, Any] | None = None
@@ -161,6 +166,8 @@ class StreamContext:
self.data_count = 0 self.data_count = 0
self.has_completion = False self.has_completion = False
self._collected_text_parts = [] self._collected_text_parts = []
self._collected_text_chars = 0
self._stored_collected_text_chars = 0
self.input_tokens = 0 self.input_tokens = 0
self.output_tokens = 0 self.output_tokens = 0
self.cached_tokens = 0 self.cached_tokens = 0
@@ -192,10 +199,30 @@ class StreamContext:
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)""" """已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""
return "".join(self._collected_text_parts) return "".join(self._collected_text_parts)
@property
def collected_text_length(self) -> int:
"""已收集文本的总字符数(包含未保留到内存的截断部分)"""
return self._collected_text_chars
def append_text(self, text: str) -> None: def append_text(self, text: str) -> None:
"""追加文本内容(仅在需要收集文本时调用""" """追加文本内容(仅保留有限前缀,避免长流导致内存增长"""
if text: if not text:
return
text_len = len(text)
self._collected_text_chars += text_len
remaining = _MAX_COLLECTED_TEXT_CHARS - self._stored_collected_text_chars
if remaining <= 0:
return
if text_len <= remaining:
self._collected_text_parts.append(text) self._collected_text_parts.append(text)
self._stored_collected_text_chars += text_len
return
self._collected_text_parts.append(text[:remaining])
self._stored_collected_text_chars += remaining
def update_provider_info( def update_provider_info(
self, self,

View File

@@ -626,9 +626,8 @@ class StreamTelemetryRecorder:
估算采用 ~4 字符/token 的保守比例。 估算采用 ~4 字符/token 的保守比例。
""" """
# 输出 tokens从已收集的文本估算 # 输出 tokens从已收集的文本估算
collected = ctx.collected_text if ctx.collected_text_length > 0:
if collected: ctx.output_tokens = max(1, ctx.collected_text_length // 4)
ctx.output_tokens = max(1, len(collected) // 4)
# 输入 tokens从请求体文本内容估算 # 输入 tokens从请求体文本内容估算
try: try:

View File

@@ -345,7 +345,7 @@ class Config:
per_worker_total = available_connections // max(self.worker_processes, 1) per_worker_total = available_connections // max(self.worker_processes, 1)
# pool_size 取总数的一半,另一半留给 overflow # pool_size 取总数的一半,另一半留给 overflow
pool_size = max(per_worker_total // 2, 5) # 最小 5 个连接 pool_size = max(per_worker_total // 2, 5) # 最小 5 个连接
return min(pool_size, 30) # 最大 30 个连接 return min(pool_size, 15) # 最大 15 个连接
def _auto_max_overflow(self) -> int: def _auto_max_overflow(self) -> int:
"""智能计算最大溢出连接数 - 与 pool_size 相同""" """智能计算最大溢出连接数 - 与 pool_size 相同"""
@@ -361,21 +361,17 @@ class Config:
3. 需要为数据库连接、Redis 连接等预留资源 3. 需要为数据库连接、Redis 连接等预留资源
公式: base_connections / workers 公式: base_connections / workers
- 单 Worker: 200 连接(适合开发/低负载) - 单 Worker: 100 连接
- 多 Worker: 按比例分配,确保总数不超过系统限制 - 多 Worker: 按比例分配,确保总数不超过系统限制
范围: 50 - 200 范围: 30 - 100
""" """
# 基础连接数:默认将总 HTTP 连接预算控制在 200 base_connections = 100
# (预留给 DB、Redis、内部服务等
base_connections = 200
workers = max(self.worker_processes, 1) workers = max(self.worker_processes, 1)
# 每个 Worker 分配的连接数
per_worker = base_connections // workers per_worker = base_connections // workers
# 限制范围:最小 50保证基本并发最大 200控制内存与 socket 占用) return max(30, min(per_worker, 100))
return max(50, min(per_worker, 200))
def _auto_http_keepalive_connections(self) -> int: def _auto_http_keepalive_connections(self) -> int:
""" """

View File

@@ -640,12 +640,14 @@ def main() -> Any:
# Start server # Start server
# 根据环境设置热重载 # 根据环境设置热重载
is_dev = config.environment == "development"
uvicorn.run( uvicorn.run(
"src.main:app", "src.main:app",
host=config.host, host=config.host,
port=config.port, port=config.port,
log_level=log_level, log_level=log_level,
reload=config.environment == "development", # 只在开发环境启用热重载 reload=is_dev,
reload_dirs=["src"] if is_dev else None,
access_log=False, # 禁用 uvicorn 访问日志,使用自定义中间件 access_log=False, # 禁用 uvicorn 访问日志,使用自定义中间件
log_config=uvicorn_log_config, # 使用自定义日志配置 log_config=uvicorn_log_config, # 使用自定义日志配置
) )

View File

@@ -22,11 +22,12 @@ except ImportError: # pragma: no cover
tiktoken = None tiktoken = None
@lru_cache(maxsize=32) @lru_cache(maxsize=4)
def _get_encoder_cached(model: str) -> Any: def _get_encoder_cached(model: str) -> Any:
"""全局编码器缓存。 """全局编码器缓存。
目的:避免在多实例/多请求场景下重复初始化 tiktoken 编码器。 目的:避免在多实例/多请求场景下重复初始化 tiktoken 编码器。
实际只有 cl100k_base / o200k_base / p50k_base 等少数几种编码4 个足够。
""" """
if not TIKTOKEN_AVAILABLE: if not TIKTOKEN_AVAILABLE:
raise RuntimeError("tiktoken not installed") raise RuntimeError("tiktoken not installed")

View File

@@ -5,11 +5,25 @@ from src.api.handlers.base.stream_context import StreamContext
def test_collected_text_append_and_property() -> None: def test_collected_text_append_and_property() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat") ctx = StreamContext(model="test-model", api_format="openai:chat")
assert ctx.collected_text == "" assert ctx.collected_text == ""
assert ctx.collected_text_length == 0
ctx.append_text("hello") ctx.append_text("hello")
ctx.append_text(" ") ctx.append_text(" ")
ctx.append_text("world") ctx.append_text("world")
assert ctx.collected_text == "hello world" assert ctx.collected_text == "hello world"
assert ctx.collected_text_length == len("hello world")
def test_collected_text_is_capped_but_total_length_is_preserved() -> None:
ctx = StreamContext(model="test-model", api_format="openai:chat")
cap = stream_context._MAX_COLLECTED_TEXT_CHARS
ctx.append_text("a" * (cap - 4))
ctx.append_text("b" * 10)
assert len(ctx.collected_text) == cap
assert ctx.collected_text == ("a" * (cap - 4)) + ("b" * 4)
assert ctx.collected_text_length == cap + 6
def test_reset_for_retry_clears_state() -> None: def test_reset_for_retry_clears_state() -> None: