82f4ec4910
enable_thinking=False was a no-op for minimax and openai sources: both profiles had empty dicts on each side, so the payload update injected nothing while the caller believed reasoning had been turned off. A downstream project was blocked on exactly this. The root cause is that an empty dict meant two different things -- "no injection needed" and "we do not know how this provider spells it" -- and that a provider-level table cannot express what turned out to be a per-model property. Live testing showed MiniMax-M3 can disable reasoning via reasoning_effort while M2.7 and M2.5 cannot be disabled at all, which two external registries independently confirm. So the shape stays at provider level and a capability table joins it at model level. Unknown, unsupported and no-opinion are now three distinct values, and resolve_thinking is the single place they meet: it raises at assembly time when a model cannot honour the request, warns and injects for unregistered models, and injects silently otherwise. Every registered capability carries the evidence it was derived from. enable_thinking also joins the cache fingerprint, since it now really does change the request body.
427 lines
17 KiB
Python
427 lines
17 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, 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))
|