Files
PolyGateway/src/polygateway/client.py
T
iomgaa 7622eb0402 refactor: give reasoning decisions their own module
providers.py had been holding two jobs: the registry of what each
provider looks like, and the decisions made from those declarations.
Adding response-side judgement would have made it the module for
everything about reasoning, so the decisions move to thinking.py and
the registry keeps only profiles and their lookup.

Moving a module breaks any deep-path import of what moved, so the six
public symbols are promoted to the package root at the same time. The
top level is this library's stated API surface; giving downstream a
stable name to import is what makes the next reorganisation harmless.
observe_thinking stays unexported — downstream reads the verdict off
LLMResponse, and exporting it would be a permanent promise for nothing.
2026-08-25 23:48:45 -04:00

537 lines
22 KiB
Python

"""GatewayClient: 唯一组装层(D1 组装点)+ from_env/from_settings 工厂 + gather_bounded。
90% 用户三行起步: `client = GatewayClient.from_env(); resp = await client.chat(messages)`。
多逻辑角色 = 每角色调一次 `from_env(scope=...)`;共享限流/熔断状态 = 自建
后端实例(源表取并集)显式传入多次调用——共享必须显式,禁止隐式全局
(VT `evolve_llm = llm` 教训;ARCH §7.7 R5)。
"""
from __future__ import annotations
import asyncio
import hashlib
import json
import random
import time
from typing import TYPE_CHECKING, Any, Literal
from polygateway.backends.memory.breaker import InMemoryGate
from polygateway.backends.memory.cache import InMemoryCache
from polygateway.backends.memory.limiter import InMemoryLimiter
from polygateway.config import GatewaySettings
from polygateway.middleware.base import compose
from polygateway.middleware.cache import CacheMW
from polygateway.middleware.retry import RetryMW
from polygateway.middleware.structured import StructuredMW
from polygateway.middleware.telemetry import TelemetryEmitter, TelemetryMW
from polygateway.ports import TelemetryStatusProvider
from polygateway.pricing import PricingTable
from polygateway.providers import get_provider
from polygateway.sources import (
AdaptivePacer,
HealthAwareSelector,
LeastInflightSelector,
RoundRobinSelector,
SourceCooldownMemo,
)
from polygateway.thinking import get_capability, resolve_thinking
from polygateway.transports.openai_compat import OpenAICompatTransport
from polygateway.types import (
ChatRequest,
LLMResponse,
TelemetryStatus,
validate_caller_dimensions,
validate_request_overlay,
)
if TYPE_CHECKING:
from collections.abc import Awaitable, Iterable, Mapping
from pydantic import BaseModel
from polygateway.ports import (
CacheBackend,
Middleware,
ProviderGate,
RateLimiter,
SourceSelector,
StructuredOutputStrategy,
TelemetryRecorder,
Transport,
)
from polygateway.providers import ProviderProfile
from polygateway.thinking import ThinkingCapability
from polygateway.types import (
BackpressurePolicy,
RetryPolicy,
SourceConfig,
)
def _guard_thinking(
sources: list[SourceConfig],
profiles: list[ProviderProfile],
capabilities: Mapping[str, ThinkingCapability] | None,
) -> None:
"""装配期把不可满足的推理开关炸掉,而不是留到运行时(issue #5)。
与 transport 内的同一次判定不是重复: 那里兜的是"构造函数全量注入"这条路
(CLAUDE.md §4.5 的第二条装配路),而工厂路占 90% 场景,配置错误应当在装配期
就带着指路信息炸掉。`get_provider` 现在就是同一形态的双点调用。
"""
for source, profile in zip(sources, profiles, strict=True):
resolve_thinking(
profile,
get_capability(source.model, table=capabilities),
source.enable_thinking,
model=source.model,
)
def _fingerprint_mark(source: SourceConfig) -> str:
"""单源的指纹标记;`enable_thinking` 仅在**表态时**追加。
只在表态时追加不是省事: 这样只配了 `extra_body` 的存量源字面量与 issue #4
时期逐字相同,升级本版本不会给它们平白来一次全量缓存冷启动。
"""
parts: list[Any] = [source.model, dict(source.extra_body)]
if source.enable_thinking is not None:
parts.append(source.enable_thinking)
return json.dumps(parts, sort_keys=True, ensure_ascii=False)
def build_model_fingerprint(sources: Iterable[SourceConfig]) -> str:
"""缓存 key 的模型身份: 多源 scope = 排序去重的 model 合集。
配置级采样参数(`extra_body`)必须参与,否则把 temperature 从 0 改成 1
后重启仍会读到旧缓存(issue #4 设计决策 C)。`enable_thinking` 同理
(issue #5): 它一旦真正改变请求体,"关掉推理后重启"就会读到开着推理时
缓存的旧响应。全源两者皆未表态时字面量与历史实现逐字相同,不触发存量
缓存冷启动。
"""
fingerprint = ",".join(sorted({s.model for s in sources}))
# 按 (model, extra_body[, enable_thinking]) 而非源名摘要: 语义是"本 scope
# 会用哪些(模型, 请求形态)组合",改源名不该误触全量冷启动
marks = sorted(
{_fingerprint_mark(s) for s in sources if s.extra_body or s.enable_thinking is not None}
)
if marks:
digest = hashlib.sha256("".join(marks).encode("utf-8")).hexdigest()
fingerprint = f"{fingerprint}|{digest}"
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 _telemetry_status_of(telemetry: TelemetryRecorder | None) -> TelemetryStatus | None:
"""三个 client 共用的状态取值点: 不提供状态的 recorder 一律返回 None。
判定写成 `isinstance(可选端口)` 而不是裸 `getattr`: 两者运行时都是结构检查
(`@runtime_checkable` 按属性存在性判定),差别在**契约有没有名字**——端口是
写进 `ports.py` 的公开承诺,下游可以照着实现;散落的 `getattr` 不是,而
`aclose` 当年正是被复制成三份鸭子类型探测才漂移出越权关闭(设计 §3.3/§3.4)。
"""
if isinstance(telemetry, TelemetryStatusProvider):
return telemetry.telemetry_status
return None
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% 场景。"""
def __init__(
self,
*,
scope: str,
sources: list[SourceConfig],
selector: SourceSelector,
limiter: RateLimiter,
breaker: ProviderGate,
transport: Transport,
retry: RetryPolicy,
backpressure: BackpressurePolicy,
quota_full: str = "wait",
circuit_open: str = "fail_fast",
telemetry: TelemetryRecorder | None = None,
pricing: PricingTable | None = None,
text_cap: int | None = None,
cache: CacheBackend | None = None,
cache_namespace: str | None = None,
cache_ttl_s: int | None = None,
structured_strategy: StructuredOutputStrategy | None = None,
structured_escalation: StructuredOutputStrategy | None = None,
structured_max_retries: int = 1,
now: Any = time.monotonic,
sleep: Any = asyncio.sleep,
rng: Any = random.random,
) -> None:
emitter = (
TelemetryEmitter(telemetry, pricing=pricing, text_cap=text_cap)
if telemetry is not None
else None
)
terminal = RetryMW(
scope=scope,
sources=sources,
selector=selector,
limiter=limiter,
gate=breaker,
transport=transport,
retry=retry,
backpressure=backpressure,
quota_full=quota_full,
circuit_open=circuit_open,
cooldown_memo=SourceCooldownMemo(now=now),
# AIMD ceiling 尊重源级静态并发上限(独立核验 I1: 不得静默钳制大于 64 的配置)
pacer=AdaptivePacer(
ceiling=float(max([64, *(s.max_concurrency for s in sources if s.max_concurrency)]))
),
emitter=emitter,
now=now,
sleep=sleep,
rng=rng,
)
middlewares: list[Middleware] = []
if emitter is not None:
middlewares.append(TelemetryMW(emitter, now=now))
if cache is not None:
if cache_namespace is None or cache_ttl_s is None:
raise ValueError("启用缓存必须提供 cache_namespace 与 cache_ttl_s")
# 多源 scope 的 key 身份;源集合或其 extra_body 变化 → 一次性冷启动
middlewares.append(
CacheMW(
backend=cache,
model_fingerprint=build_model_fingerprint(sources),
default_namespace=cache_namespace,
ttl_s=cache_ttl_s,
strategy=structured_strategy,
)
)
if structured_strategy is not None:
middlewares.append(
StructuredMW(
strategy=structured_strategy,
max_retries=structured_max_retries,
escalation=structured_escalation,
)
)
self._structured_available = structured_strategy is not None
self._terminal = terminal # 内部引用: 装配自省/测试用
self._handler = compose(middlewares, terminal)
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
@property
def telemetry_status(self) -> TelemetryStatus | None:
"""遥测后端的可写状态;无遥测或注入的 recorder 不提供状态时为 None。
判定收敛在 `_telemetry_status_of` 一处(不是三处各自探测): 三个 client
的 `aclose` 曾各持一份逐字复制,漂移的结果就是越权关闭(设计 §3.3/§3.4)。
"""
return _telemetry_status_of(self._telemetry)
async def chat(
self,
messages: list[dict[str, Any]],
*,
session_id: str | None = None,
parent_call_id: str | None = None,
cache_salt: str | None = None,
cache_namespace: str | None = None,
structured: type[BaseModel] | Literal["json"] | None = None,
stream: bool = True,
overlay: Mapping[str, Any] | None = None,
tenant_id: str | None = None,
meta: Mapping[str, Any] | None = None,
) -> LLMResponse:
"""一次治理调用(签名冻结,ARCH §5.2;与三项目 LLMProvider 协议兼容)。
`overlay` 是采样参数覆盖层(`temperature`/`seed`/`max_tokens` 等),优先级
高于源级 `extra_body`、低于结构化输出的注入。带默认值的 keyword-only
参数不影响既有调用点(issue #4)。
`tenant_id` 与 `meta` 是调用方自定义维度,只进遥测、**不进缓存 key**
(租户隔离由 `cache_namespace` 负责,ARCH §7.5);前者享有真实列待遇
(可挂 RLS、可进复合索引),后者是任意 KV 容器(issue #11)。
"""
if structured is not None and not self._structured_available:
raise ImportError(
"结构化输出未启用: 安装 pip install 'polygateway[structured]' 后重新装配"
)
# 进洋葱之前校验并拷贝: 保护键/不可序列化值在此收口(否则会在 CacheMW
# 的降级 try 之外抛裸 TypeError);拷贝防调用方复用同一 dict 逐次改 seed
# 造成的竞态。同一份快照填 overlay 与 sampling——前者会被结构化注入,
# 后者跨层恒定,供缓存 key 与遥测读取(设计决策 A/B/E)
sampling = validate_request_overlay(overlay or {}, origin="chat(overlay=...)")
# 同理必须在洋葱之外: 洋葱内的一切失败都被遥测层降级成 warning(库铁律
# 「遥测写失败降级不冒泡」),校验放里面等于没有校验——非法维度会变成
# 静默丢失的遥测行,而调用照常发出(issue #11 §4.2)
dimension_tenant_id, dimensions = validate_caller_dimensions(
tenant_id, meta, origin="chat(tenant_id=..., meta=...)"
)
request = ChatRequest(
messages=messages,
session_id=session_id,
parent_call_id=parent_call_id,
cache_salt=cache_salt,
cache_namespace=cache_namespace,
structured=structured,
stream=stream,
overlay=sampling,
sampling=sampling,
tenant_id=dimension_tenant_id,
meta=dimensions,
)
return await self._handler(request)
async def aclose(self) -> None:
"""幂等释放**自建**资源: transport、遥测、缓存、限流/熔断后端。
注入的组件一律不碰——它们可能被别的 client 共享,关掉即越权。
"""
if self._closed:
return
self._closed = True
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
async def __aexit__(self, *exc_info: object) -> None:
await self.aclose()
# —— 工厂 ——
@classmethod
def from_settings(
cls,
settings: GatewaySettings,
*,
limiter: RateLimiter | None = None,
breaker: ProviderGate | None = None,
cache: CacheBackend | None = None,
telemetry: TelemetryRecorder | None = None,
registry: Mapping[str, ProviderProfile] | None = None,
capabilities: Mapping[str, ThinkingCapability] | None = None,
rng: Any = random.random,
) -> GatewayClient:
"""按配置装配;显式传入的后端实例即共享(None 项按配置自建私有实例)。"""
sources = list(settings.sources)
profiles = [get_provider(s.provider, registry=registry) for s in sources]
_guard_thinking(sources, profiles, capabilities)
strategy, escalation = _build_structured(profiles)
client = cls(
scope=settings.scope,
sources=sources,
selector=_build_selector(settings.selector, rng=rng),
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,
quota_full=settings.quota_full,
circuit_open=settings.circuit_open,
telemetry=telemetry if telemetry is not None else _build_telemetry(settings),
pricing=PricingTable.from_file(settings.pricing_path)
if settings.pricing_path is not None
else None,
text_cap=settings.telemetry_text_cap,
cache=cache if cache is not None else _build_cache(settings),
cache_namespace=settings.cache_namespace,
cache_ttl_s=settings.cache_ttl_s,
structured_strategy=strategy,
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(
cls,
scope: str = "LLM",
*,
limiter: RateLimiter | None = None,
breaker: ProviderGate | None = None,
cache: CacheBackend | None = None,
telemetry: TelemetryRecorder | None = None,
registry: Mapping[str, ProviderProfile] | None = None,
capabilities: Mapping[str, ThinkingCapability] | None = None,
env: Mapping[str, str] | None = None,
) -> GatewayClient:
"""从 .env/环境变量装配一个 scope 的 client(键名清单见 .env.example)。"""
return cls.from_settings(
GatewaySettings.from_env(scope, env=env),
limiter=limiter,
breaker=breaker,
cache=cache,
telemetry=telemetry,
registry=registry,
capabilities=capabilities,
)
def _build_limiter(settings: GatewaySettings, sources: list[SourceConfig]) -> RateLimiter:
if settings.limiter_backend == "redis":
from polygateway.backends.redis.limiter import RedisLimiter
assert settings.redis_url is not None # 内部不变量: _validate_backends 已保证
return RedisLimiter.from_url(
settings.redis_url,
scope=settings.scope,
sources={s.name: s for s in sources},
global_limits=settings.global_limits,
lease_ttl_s=settings.lease_ttl_s,
)
return InMemoryLimiter(
scope=settings.scope,
sources={s.name: s for s in sources},
global_limits=settings.global_limits,
lease_ttl_s=settings.lease_ttl_s,
)
def _build_breaker(settings: GatewaySettings) -> ProviderGate:
if settings.breaker_backend == "redis":
from polygateway.backends.redis.breaker import RedisGate
assert settings.redis_url is not None # 内部不变量: _validate_backends 已保证
return RedisGate.from_url(settings.redis_url, config=settings.breaker, scope=settings.scope)
return InMemoryGate(config=settings.breaker)
def _build_selector(name: str, *, rng: Any = random.random) -> SourceSelector:
if name == "round_robin":
return RoundRobinSelector()
if name == "least_inflight":
return LeastInflightSelector()
return HealthAwareSelector(rng=rng)
def _build_cache(settings: GatewaySettings) -> CacheBackend | None:
if settings.cache_backend == "none":
return None
if settings.cache_backend == "memory":
return InMemoryCache()
from polygateway.backends.redis_cache import RedisCache
assert settings.redis_url is not None # 内部不变量: _validate_backends 已保证
return RedisCache.from_url(settings.redis_url)
def _build_telemetry(settings: GatewaySettings) -> TelemetryRecorder | None:
if settings.telemetry_backend == "none":
return None
if settings.telemetry_backend == "postgres":
from polygateway.telemetry.postgres import PostgresRecorder
assert settings.telemetry_pg_dsn is not None # 内部不变量: _validate_telemetry 已保证
return PostgresRecorder(
settings.telemetry_pg_dsn,
auto_migrate=settings.telemetry_auto_migrate,
pool_max=settings.telemetry_pg_pool_max,
write_timeout_s=settings.telemetry_pg_write_timeout_s,
)
from polygateway.telemetry.sqlite import SQLiteRecorder
assert settings.telemetry_sqlite_path is not None # 内部不变量: _validate_telemetry 已保证
return SQLiteRecorder(
settings.telemetry_sqlite_path, auto_migrate=settings.telemetry_auto_migrate
)
def _build_structured(
profiles: list[ProviderProfile],
) -> tuple[StructuredOutputStrategy | None, StructuredOutputStrategy | None]:
"""按注册表能力选策略(设计 §5): 全员支持原生 schema 才用 NativeSchema。
json_repair 缺失(未装 structured extra)→ 返回 (None, None),
chat(structured=...) 时显式报缺 extra,绝不静默跳过校验。
"""
try:
from polygateway.structured.json_repair import JsonRepairStrategy
from polygateway.structured.native_schema import NativeSchemaStrategy
except ImportError:
return None, None
try:
native_all = all(p.supports_native_schema for p in profiles)
if native_all:
return NativeSchemaStrategy(), NativeSchemaStrategy()
return JsonRepairStrategy(), None
except ImportError:
return None, None
async def gather_bounded[T](aws: Iterable[Awaitable[T]], *, concurrency: int) -> list[T]:
"""有界并发 gather(D5 便利函数,替代 VT 手搓 semaphore+gather 样板)。
语义与 `asyncio.gather` 默认一致: 结果保序、首个异常上抛;仅增加并发上限。
"""
if concurrency < 1:
raise ValueError("concurrency 必须 ≥ 1")
sem = asyncio.Semaphore(concurrency)
async def _run(aw: Awaitable[T]) -> T:
async with sem:
return await aw
return await asyncio.gather(*(_run(aw) for aw in aws))