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 REDIS_URL=redis://:${REDIS_PASSWORD}@${REDIS_HOST:-localhost}:${REDIS_PORT:-6379}/0
# 启动 uvicorn热重载模式
echo "🚀 启动本地开发服务器..."
echo "📍 后端地址: http://localhost:8084"
echo "📊 数据库: ${DATABASE_URL}"
# 开发环境连接池低配(节省内存
export DB_POOL_SIZE=${DB_POOL_SIZE:-5}
export DB_MAX_OVERFLOW=${DB_MAX_OVERFLOW:-5}
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 ""
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 的保守比例。
"""
# 输出 tokens从已收集的文本估算
collected = ctx.collected_text
if collected:
ctx.output_tokens = max(1, len(collected) // 4)
if ctx.collected_text_length > 0:
ctx.output_tokens = max(1, ctx.collected_text_length // 4)
# 输入 tokens从请求体文本内容估算
try:

View File

@@ -46,6 +46,9 @@ def is_format_converted(
)
_MAX_COLLECTED_TEXT_CHARS = 16 * 1024
@dataclass
class StreamContext:
"""
@@ -91,6 +94,8 @@ class StreamContext:
# 响应内容
_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
final_usage: dict[str, Any] | None = None
final_response: dict[str, Any] | None = None
@@ -161,6 +166,8 @@ class StreamContext:
self.data_count = 0
self.has_completion = False
self._collected_text_parts = []
self._collected_text_chars = 0
self._stored_collected_text_chars = 0
self.input_tokens = 0
self.output_tokens = 0
self.cached_tokens = 0
@@ -192,10 +199,30 @@ class StreamContext:
"""已收集的文本内容(按需拼接,避免在流式过程中频繁做字符串拷贝)"""
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:
"""追加文本内容(仅在需要收集文本时调用"""
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._stored_collected_text_chars += text_len
return
self._collected_text_parts.append(text[:remaining])
self._stored_collected_text_chars += remaining
def update_provider_info(
self,

View File

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

View File

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

View File

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

View File

@@ -22,11 +22,12 @@ except ImportError: # pragma: no cover
tiktoken = None
@lru_cache(maxsize=32)
@lru_cache(maxsize=4)
def _get_encoder_cached(model: str) -> Any:
"""全局编码器缓存。
目的:避免在多实例/多请求场景下重复初始化 tiktoken 编码器。
实际只有 cl100k_base / o200k_base / p50k_base 等少数几种编码4 个足够。
"""
if not TIKTOKEN_AVAILABLE:
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:
ctx = StreamContext(model="test-model", api_format="openai:chat")
assert ctx.collected_text == ""
assert ctx.collected_text_length == 0
ctx.append_text("hello")
ctx.append_text(" ")
ctx.append_text("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: