"""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, TypeVar 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.pricing import PricingTable from polygateway.providers import get_capability, get_provider, resolve_thinking from polygateway.sources import ( AdaptivePacer, HealthAwareSelector, LeastInflightSelector, RoundRobinSelector, SourceCooldownMemo, ) from polygateway.transports.openai_compat import OpenAICompatTransport from polygateway.types import ChatRequest, LLMResponse, 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, ThinkingCapability from polygateway.types import ( BackpressurePolicy, RetryPolicy, SourceConfig, ) _T = TypeVar("_T") 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 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", telemetry: TelemetryRecorder | None = None, pricing: PricingTable | 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) 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, 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 self._closed = False 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, ) -> LLMResponse: """一次治理调用(签名冻结,ARCH §5.2;与三项目 LLMProvider 协议兼容)。 `overlay` 是采样参数覆盖层(`temperature`/`seed`/`max_tokens` 等),优先级 高于源级 `extra_body`、低于结构化输出的注入。带默认值的 keyword-only 参数不影响既有调用点(issue #4)。 """ 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=...)") 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, ) return await self._handler(request) async def aclose(self) -> None: """幂等释放: transport 连接池、遥测连接、缓存客户端。""" 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() 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) return 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), transport=OpenAICompatTransport(registry=registry, capabilities=capabilities), retry=settings.retry, backpressure=settings.backpressure, quota_full=settings.quota_full, 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, 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, ) @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) from polygateway.telemetry.sqlite import SQLiteRecorder assert settings.telemetry_sqlite_path is not None # 内部不变量: _validate_telemetry 已保证 return SQLiteRecorder(settings.telemetry_sqlite_path) 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(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))