diff --git a/src/muse/启动.py b/src/muse/启动.py index 9391e1c..259ee58 100644 --- a/src/muse/启动.py +++ b/src/muse/启动.py @@ -197,7 +197,9 @@ class 应用装配: ) return self.数据库 - def 要求模型执行器(self, 上下文: 执行上下文, 计价: 调用计价 | None) -> 模型执行器: + def 要求模型执行器( + self, 上下文: 执行上下文, 计价: 调用计价 | None, *, 客户端: Any | None = None + ) -> 模型执行器: 数据库 = self.要求数据库() 版本 = 配置版本管理(数据库).读取任务绑定(上下文.领取.任务ID) 内容 = 版本.内容 @@ -221,7 +223,7 @@ class 应用装配: self.要求原文(), self.要求证据(), 策略, - 直接宿主(HTTP传输(提供方.地址, 凭据.来源, 凭据.位置, 提供方.协议)), + 直接宿主(HTTP传输(提供方.地址, 凭据.来源, 凭据.位置, 提供方.协议, 客户端=客户端)), 实际计价, 固定配置=版本, ) diff --git a/src/muse/基础设施/模型/HTTP传输.py b/src/muse/基础设施/模型/HTTP传输.py index d5d7f85..4967778 100644 --- a/src/muse/基础设施/模型/HTTP传输.py +++ b/src/muse/基础设施/模型/HTTP传输.py @@ -5,6 +5,7 @@ from __future__ import annotations import asyncio import json from collections.abc import AsyncIterable, AsyncIterator +from contextlib import asynccontextmanager from dataclasses import dataclass, field import httpx @@ -39,12 +40,35 @@ async def 解码SSE(字节流: AsyncIterable[bytes]) -> AsyncIterator[dict]: raise 模型协议错误("SSE 在完整事件分隔符之前断开") +@asynccontextmanager +async def 作用域HTTP客户端( + *, + 连接池限制: httpx.Limits | None = None, +) -> AsyncIterator[httpx.AsyncClient]: + """在当前事件循环作用域内管理长连接客户端生命周期;退出时确保完成 aclose 释放。""" + limits = 连接池限制 or httpx.Limits( + max_keepalive_connections=20, + max_connections=50, + keepalive_expiry=60.0, + ) + 客户端 = httpx.AsyncClient( + limits=limits, + follow_redirects=False, + trust_env=False, + ) + try: + yield 客户端 + finally: + await 客户端.aclose() + + @dataclass(frozen=True) class HTTP传输: 地址: str 凭据方式: str 凭据位置: str 协议: str + 客户端: httpx.AsyncClient | None = field(default=None, repr=False) def 准备(self, 请求: dict, 总期限秒: float) -> 已准备HTTP请求: try: @@ -68,13 +92,12 @@ class HTTP传输: 头 = { "Content-Type": "application/json", "Accept-Encoding": "identity", - "Connection": "close", } if self.协议 == "anthropic": 头.update({"x-api-key": 密钥, "anthropic-version": "2023-06-01"}) else: 头["Authorization"] = "Bearer " + 密钥 - return 已准备HTTP请求(str(地址), tuple(头.items()), 数据, 总期限秒) + return 已准备HTTP请求(str(地址), tuple(头.items()), 数据, 总期限秒, 客户端=self.客户端) async def 流(self, 请求: dict, 总期限秒: float) -> AsyncIterator[dict]: async for 事件 in self.准备(请求, 总期限秒).流(): @@ -87,23 +110,37 @@ class 已准备HTTP请求: 头: tuple[tuple[str, str], ...] = field(repr=False) 数据: bytes = field(repr=False) 总期限秒: float + 客户端: httpx.AsyncClient | None = field(default=None, repr=False) async def 流(self) -> AsyncIterator[dict]: + # 作用域客户端由调用方在当前事件循环内创建并负责关闭;未注入或已关闭时退回一次性客户端。 + 作用域 = self.客户端 + if 作用域 is not None and not 作用域.is_closed: + 客户端, 需关闭 = 作用域, False + else: + 客户端 = httpx.AsyncClient( + timeout=self.总期限秒, follow_redirects=False, trust_env=False + ) + 需关闭 = True try: async with asyncio.timeout(self.总期限秒): - async with httpx.AsyncClient( - timeout=self.总期限秒, follow_redirects=False, trust_env=False - ) as 客户端: - async with 客户端.stream( - "POST", self.地址, content=self.数据, headers=dict(self.头) - ) as 响应: - if 响应.is_error: - raise 模型协议错误( - "模型 HTTP 请求失败", 上下文={"status": 响应.status_code} - ) - async for 事件 in 解码SSE(响应.aiter_bytes()): - yield 事件 + async with 客户端.stream( + "POST", + self.地址, + content=self.数据, + headers=dict(self.头), + timeout=self.总期限秒, + ) as 响应: + if 响应.is_error: + raise 模型协议错误( + "模型 HTTP 请求失败", 上下文={"status": 响应.status_code} + ) + async for 事件 in 解码SSE(响应.aiter_bytes()): + yield 事件 except (TimeoutError, httpx.TimeoutException): raise 模型协议错误("模型调用超过总期限") from None except httpx.HTTPError: raise 模型协议错误("模型传输中断,调用结果可能未知") from None + finally: + if 需关闭: + await 客户端.aclose() diff --git a/src/muse/编排/生成正文.py b/src/muse/编排/生成正文.py index 7307f7b..9fc5700 100644 --- a/src/muse/编排/生成正文.py +++ b/src/muse/编排/生成正文.py @@ -44,6 +44,7 @@ from muse.任务运行.接口 import ( from muse.共享.调用身份 import 内容用途 from muse.共享.错误 import Muse错误 from muse.基础设施.数据库.连接 import 数据库工厂 +from muse.基础设施.模型.HTTP传输 import 作用域HTTP客户端 from muse.审校修订.接口 import 审校服务 from muse.正文写作.接口 import ( 写作输出合同, @@ -202,9 +203,6 @@ def 登记生成正文(登记: 流程登记, 装配, 计价) -> None: def 受限探索(上下文: 执行上下文) -> 步骤结果: 任务 = 上下文.任务 工具 = 装配.要求只读工具(上下文.领取) - 清单: tuple[探索条目, ...] = () - 绑定集: dict[str, 来源绑定] = {} - 历史: list[回合记录] = [] 轮次上限 = int( 任务.冻结输入["冻结上下文"].get("stop_conditions", {}).get("探索轮数上限", 6) ) @@ -213,44 +211,53 @@ def 登记生成正文(登记: 流程登记, 装配, 计价) -> None: for d in 工具.工具.values() ) 探索提示, 探索模板哈希 = 受限探索模板() - for _ in range(轮次上限): - 请求 = _写手请求( - 装配, - 上下文, - 阶段="探索", - 系统提示=探索提示, - 用户输入=json.dumps( - { - "粒度": 任务.冻结输入["输入"]["写作任务"]["粒度"], - "用途说明": 任务.冻结输入["输入"]["写作任务"].get("用途说明", ""), - "材料需求": ["目标章细纲", "时点事实", "历史正文"], - }, - ensure_ascii=False, - ), - 工具=工具描述, - 历史=tuple(历史), - ) - 执行器 = 装配.要求模型执行器(上下文, 计价) - 授权 = _原文授权(装配, 任务, 请求, 方式="persistent") - 回合 = asyncio.run(执行器.执行回合(上下文, 请求, 阶段="探索", 原文授权ID=授权)) - if 回合.结果.状态 == "tool_calls": - 回传: list[工具回传] = [] - for 调用 in 回合.结果.工具调用: - 结果 = 工具.调用(调用) - 清单 = 记录探索(清单, 调用.调用ID, 调用.名称, 调用.参数, 结果) - for 源 in 结果.来源: - 绑定集.setdefault( - 源.来源ID, - 来源绑定(源.来源ID, 源.数据版本, 源.结构哈希, 源.投影版本), - ) - 回传.append(工具回传(调用, json.dumps(结果.内容, ensure_ascii=False))) - 历史.append(回合记录(回合.结果, tuple(回传))) - continue - if 回合.结果.状态 != "completed": - raise 生成正文错误(f"探索回合未完成:{回合.结果.状态}") - break - else: - raise 生成正文错误("探索轮数超出停止条件上限") + + async def _运行探索循环(): + 清单: tuple[探索条目, ...] = () + 绑定集: dict[str, 来源绑定] = {} + 历史: list[回合记录] = [] + async with 作用域HTTP客户端() as 客户端: + for _ in range(轮次上限): + 请求 = _写手请求( + 装配, + 上下文, + 阶段="探索", + 系统提示=探索提示, + 用户输入=json.dumps( + { + "粒度": 任务.冻结输入["输入"]["写作任务"]["粒度"], + "用途说明": 任务.冻结输入["输入"]["写作任务"].get("用途说明", ""), + "材料需求": ["目标章细纲", "时点事实", "历史正文"], + }, + ensure_ascii=False, + ), + 工具=工具描述, + 历史=tuple(历史), + ) + 执行器 = 装配.要求模型执行器(上下文, 计价, 客户端=客户端) + 授权 = _原文授权(装配, 任务, 请求, 方式="persistent") + 回合 = await 执行器.执行回合(上下文, 请求, 阶段="探索", 原文授权ID=授权) + if 回合.结果.状态 == "tool_calls": + 回传: list[工具回传] = [] + for 调用 in 回合.结果.工具调用: + 结果 = 工具.调用(调用) + 清单 = 记录探索(清单, 调用.调用ID, 调用.名称, 调用.参数, 结果) + for 源 in 结果.来源: + 绑定集.setdefault( + 源.来源ID, + 来源绑定(源.来源ID, 源.数据版本, 源.结构哈希, 源.投影版本), + ) + 回传.append(工具回传(调用, json.dumps(结果.内容, ensure_ascii=False))) + 历史.append(回合记录(回合.结果, tuple(回传))) + continue + if 回合.结果.状态 != "completed": + raise 生成正文错误(f"探索回合未完成:{回合.结果.状态}") + break + else: + raise 生成正文错误("探索轮数超出停止条件上限") + return 清单, 绑定集 + + 清单, 绑定集 = asyncio.run(_运行探索循环()) if not 清单: raise 生成正文错误("写手没有进行任何受限探索读取") return 步骤结果( diff --git a/tests/单元/test_模型长连接复用.py b/tests/单元/test_模型长连接复用.py new file mode 100644 index 0000000..c51a7ef --- /dev/null +++ b/tests/单元/test_模型长连接复用.py @@ -0,0 +1,120 @@ +"""模型 HTTP 传输执行作用域长连接复用与生命周期机械验证。""" + +import asyncio +import threading +from http.server import BaseHTTPRequestHandler, HTTPServer +from socketserver import ThreadingMixIn + +import pytest + +from muse.基础设施.模型.HTTP传输 import HTTP传输, 作用域HTTP客户端 + + +class _可复用HTTP服务(ThreadingMixIn, HTTPServer): + daemon_threads = True + + def __init__(self, server_address, handler_cls): + super().__init__(server_address, handler_cls) + self.tcp_accept_count = 0 + self.request_count = 0 + self.lock = threading.Lock() + + def get_request(self): + sock, addr = super().get_request() + with self.lock: + self.tcp_accept_count += 1 + return sock, addr + + +class _SSE处理器(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self): + length = int(self.headers.get("content-length", 0)) + _ = self.rfile.read(length) + with self.server.lock: + self.server.request_count += 1 + + body = b'data: {"type": "response.output_text.delta", "delta": "ok"}\n\ndata: [DONE]\n\n' + self.send_response(200) + self.send_header("Content-Type", "text/event-stream") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + self.wfile.flush() + + def log_message(self, format, *args): + pass + + +@pytest.fixture +def 本地HTTP服务(): + server = _可复用HTTP服务(("127.0.0.1", 0), _SSE处理器) + thread = threading.Thread(target=server.serve_forever) + thread.daemon = True + thread.start() + port = server.server_port + try: + yield f"http://127.0.0.1:{port}", server + finally: + server.shutdown() + server.server_close() + + +@pytest.mark.网络 +@pytest.mark.case_id("TC-HTTP-KEEP-ALIVE-REUSE") +def test_作用域HTTP客户端复用单条TCP连接(本地HTTP服务, tmp_path): + base_url, server = 本地HTTP服务 + key_file = tmp_path / "mock-key.txt" + key_file.write_text("test-key") + + async def _运行三次请求(): + async with 作用域HTTP客户端() as 客户端: + 传输 = HTTP传输(base_url, "受控存储", str(key_file), "responses", 客户端=客户端) + for _ in range(3): + 事件 = [] + async for 项 in 传输.流({"input": "hi"}, 总期限秒=5.0): + 事件.append(项) + assert len(事件) == 1 + assert 事件[0]["delta"] == "ok" + assert not 客户端.is_closed + assert 客户端.is_closed + + asyncio.run(_运行三次请求()) + + # 3 次 HTTP 请求在同一作用域内共享同一个客户端,只建立 1 次真实 TCP 连接 + assert server.request_count == 3 + assert server.tcp_accept_count == 1 + + +@pytest.mark.网络 +@pytest.mark.case_id("TC-HTTP-SHORT-CONNECTION-CONTRAST") +def test_未指定作用域时每次请求独立建连(本地HTTP服务, tmp_path): + base_url, server = 本地HTTP服务 + key_file = tmp_path / "mock-key.txt" + key_file.write_text("test-key") + + async def _运行三次独立请求(): + # 无外部客户端,每次请求使用临时客户端 + 传输 = HTTP传输(base_url, "受控存储", str(key_file), "responses") + for _ in range(3): + 事件 = [] + async for 项 in 传输.流({"input": "hi"}, 总期限秒=5.0): + 事件.append(项) + assert len(事件) == 1 + + asyncio.run(_运行三次独立请求()) + + assert server.request_count == 3 + assert server.tcp_accept_count == 3 + + +@pytest.mark.case_id("TC-HTTP-HEADER-NO-CONNECTION-CLOSE") +def test_请求头移除强制关闭并保留恒等编码(tmp_path): + key_file = tmp_path / "mock-key.txt" + key_file.write_text("test-key") + 传输 = HTTP传输("http://127.0.0.1:8000", "受控存储", str(key_file), "responses") + 准备 = 传输.准备({"input": "test"}, 总期限秒=10.0) + 头字典 = dict(准备.头) + assert "Connection" not in 头字典 or 头字典["Connection"] != "close" + assert 头字典.get("Accept-Encoding") == "identity" diff --git a/tests/用例清单.json b/tests/用例清单.json index 1cc1b4d..8ba7665 100644 --- a/tests/用例清单.json +++ b/tests/用例清单.json @@ -35381,6 +35381,66 @@ ], "markers": [] }, + { + "case_id": "TC-HTTP-HEADER-NO-CONNECTION-CLOSE", + "file": "tests/单元/test_模型长连接复用.py", + "symbol": "test_请求头移除强制关闭并保留恒等编码", + "parameter_ids": [], + "node_ids": [ + "tests/单元/test_模型长连接复用.py::test_请求头移除强制关闭并保留恒等编码" + ], + "fixtures": [ + "request", + "tmp_path", + "tmp_path_factory", + "测试资源接缝", + "源码资源", + "离线防护" + ], + "markers": [] + }, + { + "case_id": "TC-HTTP-KEEP-ALIVE-REUSE", + "file": "tests/单元/test_模型长连接复用.py", + "symbol": "test_作用域HTTP客户端复用单条TCP连接", + "parameter_ids": [], + "node_ids": [ + "tests/单元/test_模型长连接复用.py::test_作用域HTTP客户端复用单条TCP连接" + ], + "fixtures": [ + "request", + "tmp_path", + "tmp_path_factory", + "本地HTTP服务", + "测试资源接缝", + "源码资源", + "离线防护" + ], + "markers": [ + "网络" + ] + }, + { + "case_id": "TC-HTTP-SHORT-CONNECTION-CONTRAST", + "file": "tests/单元/test_模型长连接复用.py", + "symbol": "test_未指定作用域时每次请求独立建连", + "parameter_ids": [], + "node_ids": [ + "tests/单元/test_模型长连接复用.py::test_未指定作用域时每次请求独立建连" + ], + "fixtures": [ + "request", + "tmp_path", + "tmp_path_factory", + "本地HTTP服务", + "测试资源接缝", + "源码资源", + "离线防护" + ], + "markers": [ + "网络" + ] + }, { "case_id": "TC-O08-SSE-001", "file": "tests/单元/test_任务事件收尾.py",