优化(模型传输): 移除强制短连接并以执行作用域复用HTTP客户端
This commit is contained in:
parent
1571abad43
commit
cd8914f1f4
@ -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传输(提供方.地址, 凭据.来源, 凭据.位置, 提供方.协议, 客户端=客户端)),
|
||||
实际计价,
|
||||
固定配置=版本,
|
||||
)
|
||||
|
||||
@ -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()
|
||||
|
||||
@ -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 步骤结果(
|
||||
|
||||
120
tests/单元/test_模型长连接复用.py
Normal file
120
tests/单元/test_模型长连接复用.py
Normal file
@ -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"
|
||||
@ -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",
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user