diff --git a/src/coding/proxy/vendors/zhipu.py b/src/coding/proxy/vendors/zhipu.py index 528cabf..e6b4680 100644 --- a/src/coding/proxy/vendors/zhipu.py +++ b/src/coding/proxy/vendors/zhipu.py @@ -4,20 +4,51 @@ Anthropic Messages API 协议,本模块仅做两项最小适配: 1. 模型名映射(Claude -> GLM) 2. 认证头替换(x-api-key) + +额外提供 429 Rate Limit 专用重试挽回机制: + - max_attempt = 5(1 初始 + 4 重试) + - 指数退避 + Full Jitter(1s → 2s → 4s → 8s) + - 优先尊重 server retry-after header """ from __future__ import annotations +import asyncio +import logging +from collections.abc import AsyncIterator +from typing import Any + +import httpx + from ..config.schema import FailoverConfig, ZhipuConfig from ..routing.model_mapper import ModelMapper +from ..routing.rate_limit import ( + compute_effective_retry_seconds, + parse_rate_limit_headers, +) +from ..routing.retry import RetryConfig, calculate_delay +from .base import VendorResponse from .native_anthropic import NativeAnthropicVendor +logger = logging.getLogger(__name__) + +# 429 Rate Limit 重试默认配置 +_RATE_LIMIT_RETRY = RetryConfig( + max_retries=4, # 4 次重试 + 1 次初始 = 5 总尝试 + initial_delay_ms=1000, + max_delay_ms=30000, + backoff_multiplier=2.0, + jitter=True, +) + class ZhipuVendor(NativeAnthropicVendor): - """智谱 GLM 原生 Anthropic 兼容端点供应商(薄透传). + """智谱 GLM 原生 Anthropic 兼容端点供应商(薄透传 + 429 重试挽回). 通过官方 /api/anthropic 端点转发请求, 仅替换模型名和认证头,其余原样透传。 + + 429 Rate Limit 时自动重试(指数退避),降低 failover 频率。 """ _vendor_name = "zhipu" @@ -30,6 +61,121 @@ def __init__( failover_config: FailoverConfig | None = None, ) -> None: super().__init__(config, model_mapper, failover_config) + self._rl_retry = _RATE_LIMIT_RETRY + + # ── 非流式:429 重试 ──────────────────────────────────── + + async def send_message( + self, + request_body: dict[str, Any], + headers: dict[str, str], + ) -> VendorResponse: + """非流式请求,429 时自动重试.""" + max_attempts = self._rl_retry.max_attempts + + for attempt in range(max_attempts): + resp = await super().send_message(request_body, headers) + if resp.status_code != 429: + return resp + + if attempt == max_attempts - 1: + logger.warning( + "Zhipu 429 rate limit exhausted after %d attempts", + max_attempts, + ) + return resp + + delay = self._compute_retry_delay_from_headers( + resp.response_headers, attempt + ) + logger.info( + "Zhipu 429 rate limit, retry %d/%d in %.1fms", + attempt + 1, + max_attempts - 1, + delay, + ) + await asyncio.sleep(delay / 1000.0) + + return resp # pragma: no cover + + # ── 流式:429 重试 ────────────────────────────────────── + + async def send_message_stream( + self, + request_body: dict[str, Any], + headers: dict[str, str], + ) -> AsyncIterator[bytes]: + """流式请求,429 时自动重试. + + 安全性:429 在 BaseVendor.send_message_stream 中于 + status code 检查阶段即 raise(在任何 chunk yield 之前), + 因此重试不会导致已发出数据不一致。 + """ + max_attempts = self._rl_retry.max_attempts + + for attempt in range(max_attempts): + try: + # 429 在 status code 检查阶段即 raise(在任何 chunk 之前), + # 因此 __anext__ 安全:要么拿到首个 chunk,要么抛异常。 + ait = super().send_message_stream(request_body, headers) + head = await ait.__anext__() + except StopAsyncIteration: + return + except httpx.HTTPStatusError as exc: + if exc.response is None or exc.response.status_code != 429: + raise + if attempt == max_attempts - 1: + logger.warning( + "Zhipu 429 stream rate limit exhausted after %d attempts", + max_attempts, + ) + raise + + delay = self._compute_retry_delay_from_response(exc.response, attempt) + logger.info( + "Zhipu 429 stream rate limit, retry %d/%d in %.1fms", + attempt + 1, + max_attempts - 1, + delay, + ) + await asyncio.sleep(delay / 1000.0) + continue + + # yield 在 try/except 之外,避免捕获外部 athrow 的异常 + yield head + async for chunk in ait: + yield chunk + return + + # ── 延迟计算 ──────────────────────────────────────────── + + def _compute_retry_delay_from_headers( + self, + headers: dict[str, str] | None, + attempt: int, + ) -> float: + """计算重试延迟(毫秒),优先使用 server retry-after.""" + rl_info = parse_rate_limit_headers(headers, 429, None) + server_delay_s = compute_effective_retry_seconds(rl_info) + if server_delay_s is not None: + return min(server_delay_s * 1000, self._rl_retry.max_delay_ms) + return calculate_delay(attempt, self._rl_retry) + + def _compute_retry_delay_from_response( + self, + response: httpx.Response, + attempt: int, + ) -> float: + """计算重试延迟(毫秒),从 httpx.Response 提取 header.""" + rl_info = parse_rate_limit_headers( + response.headers, + response.status_code, + response.text[:500] if response.text else None, + ) + server_delay_s = compute_effective_retry_seconds(rl_info) + if server_delay_s is not None: + return min(server_delay_s * 1000, self._rl_retry.max_delay_ms) + return calculate_delay(attempt, self._rl_retry) # 向后兼容别名 diff --git a/tests/test_zhipu.py b/tests/test_zhipu.py index 2eceb41..8d010b3 100644 --- a/tests/test_zhipu.py +++ b/tests/test_zhipu.py @@ -5,20 +5,23 @@ - 其余请求体/响应原样透传 - 401 错误归一化 - 能力声明全部为 NATIVE + - 429 Rate Limit 重试挽回 """ import json +from unittest.mock import AsyncMock, patch +import httpx import pytest from coding.proxy.compat.canonical import CompatibilityStatus from coding.proxy.config.schema import ModelMappingRule, ZhipuConfig from coding.proxy.routing.model_mapper import ModelMapper +from coding.proxy.vendors.native_anthropic import NativeAnthropicVendor from coding.proxy.vendors.zhipu import ZhipuVendor -@pytest.fixture -def zhipu_vendor(): +def _make_zhipu_vendor(api_key: str = "test-zhipu-key") -> ZhipuVendor: """创建使用默认配置的 ZhipuVendor 实例.""" mapper = ModelMapper( [ @@ -42,7 +45,13 @@ def zhipu_vendor(): ), ] ) - return ZhipuVendor(ZhipuConfig(api_key="test-zhipu-key"), mapper) + return ZhipuVendor(ZhipuConfig(api_key=api_key), mapper) + + +@pytest.fixture +def zhipu_vendor(): + """创建使用默认配置的 ZhipuVendor 实例.""" + return _make_zhipu_vendor() # ── 模型映射 ────────────────────────────────────────────── @@ -292,3 +301,332 @@ def test_never_triggers_failover(self, zhipu_vendor): async def test_health_check_always_true(self, zhipu_vendor): result = await zhipu_vendor.check_health() assert result is True + + +# ── 429 Rate Limit 重试挽回 ───────────────────────────────── + + +def _make_429_response( + headers: dict[str, str] | None = None, +) -> httpx.Response: + """构造 429 HTTP 响应.""" + return httpx.Response( + status_code=429, + content=b'{"error":{"type":"rate_limit_error","message":"Too many requests"}}', + headers=headers or {}, + request=httpx.Request( + "POST", "https://open.bigmodel.cn/api/anthropic/v1/messages" + ), + ) + + +def _make_200_response() -> httpx.Response: + """构造 200 HTTP 响应.""" + body = json.dumps( + { + "id": "msg_test", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "hello"}], + "model": "glm-5.1", + "usage": {"input_tokens": 10, "output_tokens": 5}, + } + ).encode() + return httpx.Response( + status_code=200, + content=body, + headers={"content-type": "application/json"}, + request=httpx.Request( + "POST", "https://open.bigmodel.cn/api/anthropic/v1/messages" + ), + ) + + +class TestRateLimitRetry: + """429 Rate Limit 重试挽回机制.""" + + # ── 非流式 ───────────────────────────────────────────── + + @pytest.mark.asyncio + async def test_nonstream_429_retries_and_succeeds(self): + """429 两次后 200,重试成功.""" + vendor = _make_zhipu_vendor() + call_count = 0 + + async def mock_post(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count <= 2: + return _make_429_response() + return _make_200_response() + + with patch.object(vendor, "_get_client") as mock_client: + client = AsyncMock() + client.post = mock_post + mock_client.return_value = client + + resp = await vendor.send_message( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ) + + assert resp.status_code == 200 + assert call_count == 3 + + @pytest.mark.asyncio + async def test_nonstream_429_exhausted_retries(self): + """连续 5 次 429,耗尽重试后返回 429.""" + vendor = _make_zhipu_vendor() + call_count = 0 + + async def mock_post(*args, **kwargs): + nonlocal call_count + call_count += 1 + return _make_429_response() + + with patch.object(vendor, "_get_client") as mock_client: + client = AsyncMock() + client.post = mock_post + mock_client.return_value = client + + with patch("asyncio.sleep", new_callable=AsyncMock): + resp = await vendor.send_message( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ) + + assert resp.status_code == 429 + assert call_count == 5 + + @pytest.mark.asyncio + async def test_nonstream_non_429_no_retry(self): + """500 不触发重试.""" + vendor = _make_zhipu_vendor() + call_count = 0 + + async def mock_post(*args, **kwargs): + nonlocal call_count + call_count += 1 + return httpx.Response( + status_code=500, + content=b'{"error":{"type":"api_error","message":"Internal error"}}', + request=httpx.Request("POST", "https://example.com"), + ) + + with patch.object(vendor, "_get_client") as mock_client: + client = AsyncMock() + client.post = mock_post + mock_client.return_value = client + + resp = await vendor.send_message( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ) + + assert resp.status_code == 500 + assert call_count == 1 + + # ── 流式 ─────────────────────────────────────────────── + + @pytest.mark.asyncio + async def test_stream_429_retries_and_succeeds(self): + """流式 429 两次后成功.""" + call_count = 0 + + async def fake_stream(self, body, headers): + nonlocal call_count + call_count += 1 + if call_count <= 2: + resp = _make_429_response() + raise httpx.HTTPStatusError( + "429", + request=resp.request, + response=resp, + ) + yield b'data: {"type":"content_block_start"}\n\n' + yield b'data: {"type":"content_block_delta"}\n\n' + + vendor = _make_zhipu_vendor() + chunks = [] + with ( + patch.object(NativeAnthropicVendor, "send_message_stream", fake_stream), + patch("asyncio.sleep", new_callable=AsyncMock), + ): + async for chunk in vendor.send_message_stream( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ): + chunks.append(chunk) + + assert len(chunks) == 2 + assert call_count == 3 + + @pytest.mark.asyncio + async def test_stream_429_exhausted_retries_raises(self): + """流式连续 429,耗尽重试后 raise.""" + call_count = 0 + + async def fake_stream(self, body, headers): + nonlocal call_count + call_count += 1 + resp = _make_429_response() + raise httpx.HTTPStatusError( + "429", + request=resp.request, + response=resp, + ) + yield # 使函数成为 async generator(不可达,仅影响类型) + + vendor = _make_zhipu_vendor() + with ( + patch.object(NativeAnthropicVendor, "send_message_stream", fake_stream), + patch("asyncio.sleep", new_callable=AsyncMock), + pytest.raises(httpx.HTTPStatusError) as exc_info, + ): + async for _ in vendor.send_message_stream( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ): + pass + + assert exc_info.value.response.status_code == 429 + assert call_count == 5 + + @pytest.mark.asyncio + async def test_stream_500_no_retry_raises(self): + """流式 500 不触发重试,直接 raise.""" + call_count = 0 + + async def fake_stream(self, body, headers): + nonlocal call_count + call_count += 1 + resp = httpx.Response( + status_code=500, + content=b'{"error":{"type":"api_error"}}', + request=httpx.Request("POST", "https://example.com"), + ) + raise httpx.HTTPStatusError( + "500", + request=resp.request, + response=resp, + ) + yield # 使函数成为 async generator + + vendor = _make_zhipu_vendor() + with ( + patch.object(NativeAnthropicVendor, "send_message_stream", fake_stream), + pytest.raises(httpx.HTTPStatusError) as exc_info, + ): + async for _ in vendor.send_message_stream( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ): + pass + + assert exc_info.value.response.status_code == 500 + assert call_count == 1 + + # ── retry-after header ───────────────────────────────── + + @pytest.mark.asyncio + async def test_respects_retry_after_header(self): + """响应含 retry-after 时使用 server 建议延迟.""" + vendor = _make_zhipu_vendor() + call_count = 0 + sleep_delays = [] + + async def mock_post(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + return _make_429_response(headers={"retry-after": "2"}) + return _make_200_response() + + async def mock_sleep(delay): + sleep_delays.append(delay) + + with ( + patch.object(vendor, "_get_client") as mock_client, + patch("asyncio.sleep", side_effect=mock_sleep), + ): + client = AsyncMock() + client.post = mock_post + mock_client.return_value = client + + resp = await vendor.send_message( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ) + + assert resp.status_code == 200 + assert len(sleep_delays) == 1 + # retry-after=2 → 2 * 1.1 = 2.2s → 2200ms → sleep(2.2) + assert 2.0 <= sleep_delays[0] <= 2.2 + + # ── 退避延迟增长 ─────────────────────────────────────── + + @pytest.mark.asyncio + async def test_backoff_delays_increase(self): + """无 retry-after 时延迟按指数增长.""" + vendor = _make_zhipu_vendor() + sleep_delays = [] + + async def mock_sleep(delay): + sleep_delays.append(delay) + + # 禁用 jitter 以精确验证延迟 + import dataclasses + + original_jitter = vendor._rl_retry.jitter + vendor._rl_retry = dataclasses.replace(vendor._rl_retry, jitter=False) + + call_count = 0 + + async def mock_post(*args, **kwargs): + nonlocal call_count + call_count += 1 + if call_count <= 4: + return _make_429_response() + return _make_200_response() + + try: + with ( + patch.object(vendor, "_get_client") as mock_client, + patch("asyncio.sleep", side_effect=mock_sleep), + ): + client = AsyncMock() + client.post = mock_post + mock_client.return_value = client + + resp = await vendor.send_message( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ) + + assert resp.status_code == 200 + assert len(sleep_delays) == 4 + # initial=1000ms, multiplier=2.0 + # attempt 0: 1000 * 2^0 = 1000ms → sleep(1.0) + # attempt 1: 1000 * 2^1 = 2000ms → sleep(2.0) + # attempt 2: 1000 * 2^2 = 4000ms → sleep(4.0) + # attempt 3: 1000 * 2^3 = 8000ms → sleep(8.0) + assert sleep_delays[0] == pytest.approx(1.0) + assert sleep_delays[1] == pytest.approx(2.0) + assert sleep_delays[2] == pytest.approx(4.0) + assert sleep_delays[3] == pytest.approx(8.0) + finally: + vendor._rl_retry = dataclasses.replace( + vendor._rl_retry, jitter=original_jitter + ) + + # ── API key 缺失 ────────────────────────────────────── + + @pytest.mark.asyncio + async def test_missing_api_key_skips_retry(self): + """API key 缺失时 401 快速失败,不触发 429 重试.""" + vendor = _make_zhipu_vendor(api_key="") + resp = await vendor.send_message( + {"model": "claude-sonnet-4-20250514", "messages": []}, + {}, + ) + assert resp.status_code == 401