From e69ca4c82c44fbbe1bdc390391d2295274c1455d Mon Sep 17 00:00:00 2001 From: iomgaa Date: Mon, 24 Aug 2026 08:34:06 -0400 Subject: [PATCH] fix: make every client close what it built and nothing else A client used to close whatever transport, recorder or cache it happened to hold, injected or not, so the first client to shut down killed the backend its siblings were still using. That is why the explicit-sharing path the architecture prescribes was unusable in practice and downstream projects fell back to one private instance per client. The mirror image of the same gap: the redis clients the factories build for the limiter and the breaker were never closed at all, because nobody kept a reference to them once they were handed to the retry middleware. Ownership is now stated once, the way RedisLimiter already stated it: whoever builds a resource closes it, injected ones are left alone. The constructor is the full-injection path, so it owns nothing by default and only the factories mark what they built. RedisCache gains the same rule for its own client, and the three copies of the "probe for aclose, fall back to close" dance collapse into a single helper so the next correction cannot land in only one of them. --- src/polygateway/backends/redis_cache.py | 13 +- src/polygateway/client.py | 85 +++++-- src/polygateway/embedding.py | 39 +-- src/polygateway/ocr.py | 39 +-- tests/unit/test_client.py | 308 ++++++++++++++++++++++-- 5 files changed, 423 insertions(+), 61 deletions(-) diff --git a/src/polygateway/backends/redis_cache.py b/src/polygateway/backends/redis_cache.py index a755b34..67064e2 100644 --- a/src/polygateway/backends/redis_cache.py +++ b/src/polygateway/backends/redis_cache.py @@ -25,14 +25,20 @@ class RedisCache: "Redis 缓存后端需要 redis 包: pip install 'polygateway[redis]'" ) from _IMPORT_ERROR self._client = client + # 注入的客户端归注入方管理: 关掉它会弄死共享同一连接的其他组件 + # (与 RedisLimiter/RedisGate 同一纪律) + self._owns_client = False @classmethod def from_url(cls, url: str) -> RedisCache: + """自建并持有 Redis 客户端(aclose 时代关);共享后端请直接注入 client。""" if aioredis is None: raise ImportError( "Redis 缓存后端需要 redis 包: pip install 'polygateway[redis]'" ) from _IMPORT_ERROR - return cls(aioredis.from_url(url, decode_responses=True)) + cache = cls(aioredis.from_url(url, decode_responses=True)) + cache._owns_client = True + return cache async def get(self, key: str) -> str | None: return await self._client.get(key) @@ -41,4 +47,7 @@ class RedisCache: await self._client.set(key, value, ex=ttl_s) async def aclose(self) -> None: - await self._client.aclose() + """幂等释放自建客户端;注入的客户端归注入方管理。""" + if self._owns_client: + self._owns_client = False + await self._client.aclose() diff --git a/src/polygateway/client.py b/src/polygateway/client.py index 94c07db..d23defd 100644 --- a/src/polygateway/client.py +++ b/src/polygateway/client.py @@ -117,6 +117,42 @@ def build_model_fingerprint(sources: Iterable[SourceConfig]) -> str: return fingerprint +async def _aclose_component(component: object | None) -> None: + """关闭一个**自建**组件: 优先 `aclose`,退到同步 `close`,两者皆无则跳过。 + + 退到 `close` 是给 SQLiteRecorder 的(它只有同步收尾);内存后端两者皆无, + 探测后静默跳过。三个 client 曾各持一份逐字复制的探测代码,收敛为一处是 + 所有权纪律能被维持的前提——复制即是下一个 bug 的种子(设计 §3.4)。 + """ + if component is None: + return + aclose = getattr(component, "aclose", None) + if aclose is not None: + await aclose() + return + close = getattr(component, "close", None) + if close is not None: + close() + + +def _mark_owned_components( + client: Any, + *, + limiter: RateLimiter | None, + breaker: ProviderGate | None, + telemetry: TelemetryRecorder | None, +) -> None: + """工厂置位所有权(三个 client 共用): 传进来的是 None,就说明这一件是工厂自建的。 + + 与 `RedisLimiter.from_url` 逐字同款——私有属性由工厂标记,公共 API 面不变。 + transport 单列: 三处工厂都没有 transport 注入入口,它恒是自建的。 + """ + client._owns_transport = True + client._owns_limiter = limiter is None + client._owns_breaker = breaker is None + client._owns_telemetry = telemetry is None + + class GatewayClient: """统一治理入口;构造函数全量注入(测试/高级),工厂覆盖 90% 场景。""" @@ -202,6 +238,18 @@ class GatewayClient: self._transport = transport self._telemetry = telemetry self._cache = cache + # limiter/breaker 交给 RetryMW 之后仍须自持引用,否则 aclose 触达不到 + # 自建的 redis 客户端(设计 §3.4 记录的现存泄漏) + self._limiter_backend = limiter + self._breaker_backend = breaker + # 所有权默认"不拥有": `__init__` 是全量注入路径,经它传入的一切都是 + # 外部资源,关掉别人的连接会弄死共享同一后端的其他 client(ARCH §7.7 R5)。 + # 只有工厂在真正自建时才置 True + self._owns_transport = False + self._owns_telemetry = False + self._owns_cache = False + self._owns_limiter = False + self._owns_breaker = False self._closed = False async def chat( @@ -259,23 +307,23 @@ class GatewayClient: return await self._handler(request) async def aclose(self) -> None: - """幂等释放: transport 连接池、遥测连接、缓存客户端。""" + """幂等释放**自建**资源: transport、遥测、缓存、限流/熔断后端。 + + 注入的组件一律不碰——它们可能被别的 client 共享,关掉即越权。 + """ if self._closed: return self._closed = True - transport_aclose = getattr(self._transport, "aclose", None) - if transport_aclose is not None: - await transport_aclose() - telemetry_aclose = getattr(self._telemetry, "aclose", None) - if telemetry_aclose is not None: - await telemetry_aclose() # Postgres 等异步后端 - else: - telemetry_close = getattr(self._telemetry, "close", None) - if telemetry_close is not None: - telemetry_close() - cache_aclose = getattr(self._cache, "aclose", None) - if cache_aclose is not None: - await cache_aclose() + if self._owns_transport: + await _aclose_component(self._transport) + if self._owns_telemetry: + await _aclose_component(self._telemetry) + if self._owns_cache: + await _aclose_component(self._cache) + if self._owns_limiter: + await _aclose_component(self._limiter_backend) + if self._owns_breaker: + await _aclose_component(self._breaker_backend) async def __aenter__(self) -> GatewayClient: return self @@ -303,12 +351,12 @@ class GatewayClient: profiles = [get_provider(s.provider, registry=registry) for s in sources] _guard_thinking(sources, profiles, capabilities) strategy, escalation = _build_structured(profiles) - return cls( + client = cls( scope=settings.scope, sources=sources, selector=_build_selector(settings.selector, rng=rng), - limiter=limiter or _build_limiter(settings, sources), - breaker=breaker or _build_breaker(settings), + limiter=limiter if limiter is not None else _build_limiter(settings, sources), + breaker=breaker if breaker is not None else _build_breaker(settings), transport=OpenAICompatTransport(registry=registry, capabilities=capabilities), retry=settings.retry, backpressure=settings.backpressure, @@ -326,6 +374,9 @@ class GatewayClient: structured_escalation=escalation, structured_max_retries=settings.structured_max_retries, ) + _mark_owned_components(client, limiter=limiter, breaker=breaker, telemetry=telemetry) + client._owns_cache = cache is None # 缓存后端可以是 None(backend=none),helper 会跳过 + return client @classmethod def from_env( diff --git a/src/polygateway/embedding.py b/src/polygateway/embedding.py index f06b97c..4a7ba3e 100644 --- a/src/polygateway/embedding.py +++ b/src/polygateway/embedding.py @@ -25,6 +25,7 @@ from typing import TYPE_CHECKING, Any from loguru import logger +from polygateway.client import _aclose_component from polygateway.config import EmbeddingSettings from polygateway.errors import ( AllSourcesExhausted, @@ -128,6 +129,15 @@ class EmbeddingClient: TelemetryEmitter(telemetry, pricing=pricing, text_cap=text_cap) if telemetry else None ) self._telemetry = telemetry + # 限流/熔断后端在此之外只以 QuotaGate/BreakerGate 的形态存在,自持一份 + # 引用才关得到自建的 redis 客户端(设计 §3.4) + self._limiter_backend = limiter + self._breaker_backend = breaker + # 所有权默认"不拥有": `__init__` 是全量注入路径,只有工厂自建时才置 True + self._owns_transport = False + self._owns_telemetry = False + self._owns_limiter = False + self._owns_breaker = False self._pricing = pricing self._batch_size = batch_size self._normalize = normalize @@ -445,20 +455,18 @@ class EmbeddingClient: return sum(known) if known else None async def aclose(self) -> None: - """幂等释放 transport 连接池与遥测连接(与 GatewayClient 对称)。""" + """幂等释放**自建**资源(与 GatewayClient 对称);注入的组件一律不碰。""" if self._closed: return self._closed = True - transport_aclose = getattr(self._transport, "aclose", None) - if transport_aclose is not None: - await transport_aclose() - telemetry_aclose = getattr(self._telemetry, "aclose", None) - if telemetry_aclose is not None: - await telemetry_aclose() - else: - telemetry_close = getattr(self._telemetry, "close", None) - if telemetry_close is not None: - telemetry_close() + if self._owns_transport: + await _aclose_component(self._transport) + if self._owns_telemetry: + await _aclose_component(self._telemetry) + if self._owns_limiter: + await _aclose_component(self._limiter_backend) + if self._owns_breaker: + await _aclose_component(self._breaker_backend) async def __aenter__(self) -> EmbeddingClient: return self @@ -484,18 +492,19 @@ class EmbeddingClient: _build_limiter, _build_selector, _build_telemetry, + _mark_owned_components, ) from polygateway.pricing import PricingTable from polygateway.transports.openai_compat import OpenAICompatTransport gw = settings.gateway sources = list(gw.sources) - return cls( + client = cls( scope=gw.scope, sources=sources, selector=_build_selector(gw.selector), - limiter=limiter or _build_limiter(gw, sources), - breaker=breaker or _build_breaker(gw), + limiter=limiter if limiter is not None else _build_limiter(gw, sources), + breaker=breaker if breaker is not None else _build_breaker(gw), transport=OpenAICompatTransport(registry=registry), retry=gw.retry, backpressure=gw.backpressure, @@ -512,6 +521,8 @@ class EmbeddingClient: normalize=settings.normalize, expected_dim=settings.expected_dim, ) + _mark_owned_components(client, limiter=limiter, breaker=breaker, telemetry=telemetry) + return client @classmethod def from_env( diff --git a/src/polygateway/ocr.py b/src/polygateway/ocr.py index 4f6da20..a39c98f 100644 --- a/src/polygateway/ocr.py +++ b/src/polygateway/ocr.py @@ -22,6 +22,7 @@ from typing import TYPE_CHECKING, Any, Literal from loguru import logger +from polygateway.client import _aclose_component from polygateway.errors import ( AllSourcesExhausted, GovernanceBackendError, @@ -126,6 +127,15 @@ class OcrClient: self._retry = retry self._emitter = TelemetryEmitter(telemetry, text_cap=text_cap) if telemetry else None self._telemetry = telemetry + # 限流/熔断后端在此之外只以 QuotaGate/BreakerGate 的形态存在,自持一份 + # 引用才关得到自建的 redis 客户端(设计 §3.4) + self._limiter_backend = limiter + self._breaker_backend = breaker + # 所有权默认"不拥有": `__init__` 是全量注入路径,只有工厂自建时才置 True + self._owns_transport = False + self._owns_telemetry = False + self._owns_limiter = False + self._owns_breaker = False self._now = now self._sleep = sleep self._rng = rng @@ -455,20 +465,18 @@ class OcrClient: # —— 生命周期 —— async def aclose(self) -> None: - """幂等释放 transport 连接池与遥测连接(与 EmbeddingClient 对称)。""" + """幂等释放**自建**资源(与 EmbeddingClient 对称);注入的组件一律不碰。""" if self._closed: return self._closed = True - transport_aclose = getattr(self._transport, "aclose", None) - if transport_aclose is not None: - await transport_aclose() - telemetry_aclose = getattr(self._telemetry, "aclose", None) - if telemetry_aclose is not None: - await telemetry_aclose() - else: - telemetry_close = getattr(self._telemetry, "close", None) - if telemetry_close is not None: - telemetry_close() + if self._owns_transport: + await _aclose_component(self._transport) + if self._owns_telemetry: + await _aclose_component(self._telemetry) + if self._owns_limiter: + await _aclose_component(self._limiter_backend) + if self._owns_breaker: + await _aclose_component(self._breaker_backend) async def __aenter__(self) -> OcrClient: return self @@ -493,6 +501,7 @@ class OcrClient: _build_limiter, _build_selector, _build_telemetry, + _mark_owned_components, ) from polygateway.transports.monkey_ocr import MonkeyOcrTransport @@ -503,12 +512,12 @@ class OcrClient: alien = sorted({s.provider for s in sources if s.provider != "monkey"}) if alien: raise ValueError(f"OCR 装配仅支持 provider=monkey(D9 其余后端预留未实现): 发现 {alien}") - return cls( + client = cls( scope=gw.scope, sources=sources, selector=_build_selector(gw.selector), - limiter=limiter or _build_limiter(gw, sources), - breaker=breaker or _build_breaker(gw), + limiter=limiter if limiter is not None else _build_limiter(gw, sources), + breaker=breaker if breaker is not None else _build_breaker(gw), transport=MonkeyOcrTransport(), retry=gw.retry, backpressure=gw.backpressure, @@ -519,6 +528,8 @@ class OcrClient: # 一半不受控(issue #12) text_cap=gw.telemetry_text_cap, ) + _mark_owned_components(client, limiter=limiter, breaker=breaker, telemetry=telemetry) + return client @classmethod def from_env( diff --git a/tests/unit/test_client.py b/tests/unit/test_client.py index 302e644..b49fb1b 100644 --- a/tests/unit/test_client.py +++ b/tests/unit/test_client.py @@ -45,6 +45,20 @@ _ENV = { "PGW_TELEMETRY_BACKEND": "none", } +_OCR_ENV = { + "OCR__MONKEY__1__BASE_URL": "http://10.77.0.20:7866", + "OCR__MONKEY__1__API_KEY": "none", + "OCR__MONKEY__1__MODEL": "monkey-ocr", + "OCR__MONKEY__1__TIMEOUT_S": "120", + "LLM_MAX_RETRIES": "3", + "LLM_RETRY_BASE_DELAY": "2.0", + "LLM_RETRY_MAX_DELAY": "30.0", + "LLM_CIRCUIT_BREAKER_THRESHOLD": "5", + "LLM_CIRCUIT_BREAKER_COOLDOWN": "60", + "PGW_CACHE_BACKEND": "none", + "PGW_TELEMETRY_BACKEND": "none", +} + def _sse(content='{"answer": 1}'): chunk = json.dumps({"choices": [{"delta": {"content": content}}]}) @@ -359,20 +373,7 @@ class TestTelemetryTextCapWiring: """ _CAP_ENV = dict(_ENV, PGW_TELEMETRY_TEXT_CAP="8") - _OCR_CAP_ENV = { - "OCR__MONKEY__1__BASE_URL": "http://10.77.0.20:7866", - "OCR__MONKEY__1__API_KEY": "none", - "OCR__MONKEY__1__MODEL": "monkey-ocr", - "OCR__MONKEY__1__TIMEOUT_S": "120", - "LLM_MAX_RETRIES": "3", - "LLM_RETRY_BASE_DELAY": "2.0", - "LLM_RETRY_MAX_DELAY": "30.0", - "LLM_CIRCUIT_BREAKER_THRESHOLD": "5", - "LLM_CIRCUIT_BREAKER_COOLDOWN": "60", - "PGW_CACHE_BACKEND": "none", - "PGW_TELEMETRY_BACKEND": "none", - "PGW_TELEMETRY_TEXT_CAP": "8", - } + _OCR_CAP_ENV = dict(_OCR_ENV, PGW_TELEMETRY_TEXT_CAP="8") def test_gateway_from_settings_wires_the_cap(self): settings = GatewaySettings.from_env("LLM", env=self._CAP_ENV) @@ -555,3 +556,282 @@ class TestReferenceProtocolCompat: ): ... assert isinstance(_client(), proto) + + +# —— 资源所有权纪律(issue #15 D 组): 谁建的谁关,注入的一律不碰 —— + +_CACHE_ENV = dict( + _ENV, PGW_CACHE_BACKEND="memory", PGW_CACHE_NAMESPACE="proj", PGW_CACHE_TTL_S="3600" +) + + +class _Closable: + """记 close 次数的假组件;所有权纪律的唯一观测点。""" + + def __init__(self): + self.closed = 0 + + async def aclose(self): + self.closed += 1 + + +class _SyncClosable: + """只有同步 close 的假 recorder(SQLiteRecorder 形态,收敛后的 helper 须探测到)。""" + + def __init__(self): + self.closed = 0 + + def close(self): + self.closed += 1 + + +def _parts(*names): + return {name: _Closable() for name in names} + + +def _patch_builders(monkeypatch, built, *, transport_path): + """把工厂的自建点换成可计数假件;transport 无注入入口,故恒自建。""" + monkeypatch.setattr(transport_path, lambda **kwargs: built["transport"]) + monkeypatch.setattr("polygateway.client._build_limiter", lambda s, src: built["limiter"]) + monkeypatch.setattr("polygateway.client._build_breaker", lambda s: built["breaker"]) + monkeypatch.setattr("polygateway.client._build_telemetry", lambda s: built["telemetry"]) + if "cache" in built: + monkeypatch.setattr("polygateway.client._build_cache", lambda s: built["cache"]) + + +class TestGatewayClientOwnership: + """`__init__` 是全量注入路径,经它传入的一切都归调用方(设计 §3.4)。""" + + _GATEWAY_TRANSPORT = "polygateway.client.OpenAICompatTransport" + + async def test_injected_components_are_never_closed(self): + """共享 recorder/transport 被第一个关闭的 client 弄死,正是 R5 显式共享走不通的原因。""" + injected = _parts(*("transport", "telemetry", "cache", "limiter", "breaker")) + client = _client( + transport=injected["transport"], + telemetry=injected["telemetry"], + cache=injected["cache"], + cache_namespace="proj", + cache_ttl_s=3600, + limiter=injected["limiter"], + breaker=injected["breaker"], + ) + await client.aclose() + assert {name: part.closed for name, part in injected.items()} == { + "transport": 0, + "telemetry": 0, + "cache": 0, + "limiter": 0, + "breaker": 0, + } + + async def test_factory_closes_every_component_it_built(self, monkeypatch): + """泄漏钉子: 自建的 redis limiter/breaker 今天没人关,连引用都没留。""" + built = _parts("transport", "telemetry", "cache", "limiter", "breaker") + _patch_builders(monkeypatch, built, transport_path=self._GATEWAY_TRANSPORT) + client = GatewayClient.from_settings(GatewaySettings.from_env("LLM", env=_CACHE_ENV)) + await client.aclose() + assert {name: part.closed for name, part in built.items()} == { + "transport": 1, + "telemetry": 1, + "cache": 1, + "limiter": 1, + "breaker": 1, + } + + async def test_factory_keeps_hands_off_injected_components(self, monkeypatch): + built = _parts("transport", "telemetry", "cache", "limiter", "breaker") + _patch_builders(monkeypatch, built, transport_path=self._GATEWAY_TRANSPORT) + injected = _parts("telemetry", "cache", "limiter", "breaker") + client = GatewayClient.from_settings( + GatewaySettings.from_env("LLM", env=_CACHE_ENV), + limiter=injected["limiter"], + breaker=injected["breaker"], + cache=injected["cache"], + telemetry=injected["telemetry"], + ) + await client.aclose() + assert all(part.closed == 0 for part in injected.values()) + assert built["transport"].closed == 1 # 工厂恒自建 transport,归 client + + async def test_aclose_is_idempotent(self, monkeypatch): + built = _parts("transport", "telemetry", "cache", "limiter", "breaker") + _patch_builders(monkeypatch, built, transport_path=self._GATEWAY_TRANSPORT) + client = GatewayClient.from_settings(GatewaySettings.from_env("LLM", env=_CACHE_ENV)) + await client.aclose() + await client.aclose() + assert all(part.closed == 1 for part in built.values()) + + async def test_sync_only_recorder_is_closed(self, monkeypatch): + """SQLiteRecorder 只有同步 `close()`;收敛成 helper 之后这条分支不得丢。""" + built = _parts("transport", "cache", "limiter", "breaker") + recorder = _SyncClosable() + built["telemetry"] = recorder + _patch_builders(monkeypatch, built, transport_path=self._GATEWAY_TRANSPORT) + client = GatewayClient.from_settings(GatewaySettings.from_env("LLM", env=_CACHE_ENV)) + await client.aclose() + assert recorder.closed == 1 + + +def _embedding_client(**overrides): + from polygateway.embedding import EmbeddingClient + + defaults = { + "scope": "embed", + "sources": [_source()], + "selector": RoundRobinSelector(), + "limiter": _Closable(), + "breaker": _Closable(), + "transport": _Closable(), + "retry": RetryPolicy(3, 2.0, 30.0), + "backpressure": BackpressurePolicy(300.0, 0.01), + "batch_size": 2, + } + defaults.update(overrides) + return EmbeddingClient(**defaults) + + +class TestEmbeddingClientOwnership: + """三处必须各钉一次: 收敛成 helper 之后,有人把逻辑复制回去也得当场被发现。""" + + async def test_injected_components_are_never_closed(self): + injected = _parts("transport", "telemetry", "limiter", "breaker") + client = _embedding_client( + transport=injected["transport"], + telemetry=injected["telemetry"], + limiter=injected["limiter"], + breaker=injected["breaker"], + ) + await client.aclose() + assert all(part.closed == 0 for part in injected.values()) + + async def test_factory_closes_every_component_it_built(self, monkeypatch): + from polygateway.config import EmbeddingSettings + from polygateway.embedding import EmbeddingClient + + built = _parts("transport", "telemetry", "limiter", "breaker") + _patch_builders( + monkeypatch, + built, + transport_path="polygateway.transports.openai_compat.OpenAICompatTransport", + ) + settings = EmbeddingSettings( + gateway=GatewaySettings.from_env("LLM", env=_ENV), batch_size=2 + ) + client = EmbeddingClient.from_settings(settings) + await client.aclose() + assert all(part.closed == 1 for part in built.values()) + + async def test_factory_keeps_hands_off_injected_components(self, monkeypatch): + from polygateway.config import EmbeddingSettings + from polygateway.embedding import EmbeddingClient + + built = _parts("transport", "telemetry", "limiter", "breaker") + _patch_builders( + monkeypatch, + built, + transport_path="polygateway.transports.openai_compat.OpenAICompatTransport", + ) + injected = _parts("telemetry", "limiter", "breaker") + settings = EmbeddingSettings( + gateway=GatewaySettings.from_env("LLM", env=_ENV), batch_size=2 + ) + client = EmbeddingClient.from_settings( + settings, + limiter=injected["limiter"], + breaker=injected["breaker"], + telemetry=injected["telemetry"], + ) + await client.aclose() + assert all(part.closed == 0 for part in injected.values()) + assert built["transport"].closed == 1 + + +def _ocr_client(**overrides): + from polygateway.ocr import OcrClient + + defaults = { + "scope": "ocr", + "sources": [_source(name="m1", provider="monkey", model="monkey-ocr")], + "selector": RoundRobinSelector(), + "limiter": _Closable(), + "breaker": _Closable(), + "transport": _Closable(), + "retry": RetryPolicy(3, 2.0, 30.0), + "backpressure": BackpressurePolicy(300.0, 0.01), + } + defaults.update(overrides) + return OcrClient(**defaults) + + +class TestOcrClientOwnership: + async def test_injected_components_are_never_closed(self): + injected = _parts("transport", "telemetry", "limiter", "breaker") + client = _ocr_client( + transport=injected["transport"], + telemetry=injected["telemetry"], + limiter=injected["limiter"], + breaker=injected["breaker"], + ) + await client.aclose() + assert all(part.closed == 0 for part in injected.values()) + + async def test_factory_closes_every_component_it_built(self, monkeypatch): + from polygateway.config import OcrSettings + from polygateway.ocr import OcrClient + + built = _parts("transport", "telemetry", "limiter", "breaker") + _patch_builders( + monkeypatch, + built, + transport_path="polygateway.transports.monkey_ocr.MonkeyOcrTransport", + ) + client = OcrClient.from_settings(OcrSettings.from_env("OCR", env=dict(_OCR_ENV))) + await client.aclose() + assert all(part.closed == 1 for part in built.values()) + + async def test_factory_keeps_hands_off_injected_components(self, monkeypatch): + from polygateway.config import OcrSettings + from polygateway.ocr import OcrClient + + built = _parts("transport", "telemetry", "limiter", "breaker") + _patch_builders( + monkeypatch, + built, + transport_path="polygateway.transports.monkey_ocr.MonkeyOcrTransport", + ) + injected = _parts("telemetry", "limiter", "breaker") + client = OcrClient.from_settings( + OcrSettings.from_env("OCR", env=dict(_OCR_ENV)), + limiter=injected["limiter"], + breaker=injected["breaker"], + telemetry=injected["telemetry"], + ) + await client.aclose() + assert all(part.closed == 0 for part in injected.values()) + assert built["transport"].closed == 1 + + +class TestRedisCacheOwnership: + """组件内部自建的连接归组件自己;照抄 RedisLimiter._owns_client 的正确先例。""" + + async def test_injected_client_is_not_closed(self): + from polygateway.backends.redis_cache import RedisCache + + client = _Closable() + await RedisCache(client).aclose() + assert client.closed == 0 + + async def test_self_built_client_is_closed_once(self, monkeypatch): + from types import SimpleNamespace + + from polygateway.backends import redis_cache + + built = _Closable() + monkeypatch.setattr( + redis_cache, "aioredis", SimpleNamespace(from_url=lambda url, **kwargs: built) + ) + cache = redis_cache.RedisCache.from_url("redis://localhost:6379/0") + await cache.aclose() + await cache.aclose() # 幂等: 不重复关 + assert built.closed == 1