fix: 增加 embedding 请求退避重试

This commit is contained in:
Wenjie Zhang 2026-06-06 17:45:49 +08:00
parent f657edddf2
commit 410dbf6d1d
2 changed files with 287 additions and 19 deletions

View File

@ -1,4 +1,6 @@
import asyncio
import os
import time
from abc import ABC, abstractmethod
import httpx
@ -8,6 +10,11 @@ import requests
from yuxi.models.providers.cache import model_cache
from yuxi.utils import get_docker_safe_url, hashstr, logger
EMBEDDING_RATE_LIMIT_MAX_RETRIES = 10
EMBEDDING_TRANSIENT_MAX_RETRIES = 2
EMBEDDING_RETRY_MAX_DELAY_SECONDS = 10.0
EMBEDDING_RETRYABLE_STATUS_CODES = {429, 500, 502, 503, 504}
def sigmoid(x):
return 1 / (1 + np.exp(-x))
@ -113,31 +120,112 @@ class OtherEmbedding(BaseEmbeddingModel):
def build_payload(self, message: list[str] | str) -> dict:
return {"model": self.model, "input": message}
@staticmethod
def _retry_delay_seconds(retry_index: int, retry_after: str | None = None) -> float:
if retry_after:
try:
return min(float(retry_after), EMBEDDING_RETRY_MAX_DELAY_SECONDS)
except ValueError:
pass
return min(float(2 ** (retry_index - 1)), EMBEDDING_RETRY_MAX_DELAY_SECONDS)
def _prepare_retry(
self,
message: list[str] | str,
*,
retry_index: int,
response=None,
error: Exception | None = None,
) -> tuple[int, float] | None:
status_code = getattr(response, "status_code", None)
response_text = str(getattr(response, "text", "") or "")
messages = [message] if isinstance(message, str) else message
if status_code == 400 and response is not None:
logger.warning(
"Embedding request returned 400 Bad Request: "
f"model={self.model}, base_url={self.base_url}, input_count={len(messages)}, "
f"input_lengths={[len(item) for item in messages]}, body={response_text[:2000]}"
)
if status_code == 429:
max_retries = EMBEDDING_RATE_LIMIT_MAX_RETRIES
elif status_code in EMBEDDING_RETRYABLE_STATUS_CODES or status_code is None:
max_retries = EMBEDDING_TRANSIENT_MAX_RETRIES
else:
max_retries = 0
if retry_index >= max_retries:
return None
next_retry_index = retry_index + 1
retry_after = response.headers.get("Retry-After") if response is not None else None
delay = self._retry_delay_seconds(next_retry_index, retry_after)
reason = f"status={status_code}" if status_code is not None else f"error={type(error).__name__}"
logger.warning(
"Retrying embedding request: "
f"{reason}, model={self.model}, base_url={self.base_url}, "
f"retry={next_retry_index}/{max_retries}, delay={delay:.1f}s, "
f"input_count={len(messages)}, body={response_text[:1000]}"
)
return next_retry_index, delay
@staticmethod
def _extract_embeddings(result: dict) -> list[list[float]]:
if not isinstance(result, dict) or "data" not in result:
raise ValueError(f"Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]]
def encode(self, message: list[str] | str) -> list[list[float]]:
payload = self.build_payload(message)
try:
response = requests.post(self.base_url, json=payload, headers=self.headers, timeout=60)
response.raise_for_status()
result = response.json()
if not isinstance(result, dict) or "data" not in result:
raise ValueError(f"Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]]
except requests.RequestException as e:
logger.error(f"Embedding request failed: {e}, {payload}")
raise ValueError(f"Embedding request failed: {e}")
retry_index = 0
while True:
try:
response = requests.post(self.base_url, json=payload, headers=self.headers, timeout=60)
response.raise_for_status()
return self._extract_embeddings(response.json())
except requests.RequestException as e:
retry = self._prepare_retry(
message,
retry_index=retry_index,
response=getattr(e, "response", None),
error=e,
)
if retry:
retry_index, delay = retry
time.sleep(delay)
continue
logger.error(f"Embedding request failed: {e}, {payload}")
raise ValueError(f"Embedding request failed: {e}")
async def aencode(self, message: list[str] | str) -> list[list[float]]:
payload = self.build_payload(message)
async with httpx.AsyncClient() as client:
try:
response = await client.post(self.base_url, json=payload, headers=self.headers, timeout=60)
response.raise_for_status()
result = response.json()
if not isinstance(result, dict) or "data" not in result:
raise ValueError(f"Embedding failed: Invalid response format {result}")
return [item["embedding"] for item in result["data"]]
except httpx.RequestError as e:
raise ValueError(f"Embedding async request failed: {e}, {payload}, {self.base_url=}")
retry_index = 0
while True:
try:
response = await client.post(self.base_url, json=payload, headers=self.headers, timeout=60)
response.raise_for_status()
return self._extract_embeddings(response.json())
except httpx.HTTPStatusError as e:
retry = self._prepare_retry(
message,
retry_index=retry_index,
response=e.response,
error=e,
)
if retry:
retry_index, delay = retry
await asyncio.sleep(delay)
continue
raise
except httpx.RequestError as e:
retry = self._prepare_retry(message, retry_index=retry_index, error=e)
if retry:
retry_index, delay = retry
await asyncio.sleep(delay)
continue
raise ValueError(f"Embedding async request failed: {e}, {payload}, {self.base_url=}")
def get_embedding_model_info_by_id(model_id: str) -> dict:

View File

@ -1,4 +1,8 @@
from types import SimpleNamespace
import httpx
import pytest
import requests
from yuxi.agents.models import load_chat_model
from yuxi.models.chat import select_model
@ -20,6 +24,32 @@ def _model_info(model_type: str) -> ModelInfo:
)
def _capture_embed_warnings(monkeypatch: pytest.MonkeyPatch) -> list[str]:
warnings = []
monkeypatch.setattr(
"yuxi.models.embed.logger",
SimpleNamespace(
warning=warnings.append,
error=lambda *_args, **_kwargs: None,
info=lambda *_args, **_kwargs: None,
),
)
return warnings
def _requests_embedding_response(status_code: int, content: bytes | None = None) -> requests.Response:
response = requests.Response()
response.status_code = status_code
response.url = "https://example.com/v1/embeddings"
response._content = content or b'{"error":"temporary error"}'
return response
def _httpx_embedding_response(status_code: int, content: str | None = None) -> httpx.Response:
request = httpx.Request("POST", "https://example.com/v1/embeddings")
return httpx.Response(status_code, request=request, text=content or '{"error":"temporary error"}')
@pytest.mark.parametrize(
"selector,args",
[
@ -81,6 +111,156 @@ async def test_embedding_connection_reports_dimension_mismatch(monkeypatch):
assert await model.test_connection() == (False, "Embedding 维度不一致:配置 4实际 3")
def test_embedding_sync_400_logs_warning(monkeypatch):
warnings = _capture_embed_warnings(monkeypatch)
model = OtherEmbedding(
model="namespace/embedding-model",
base_url="https://example.com/v1/embeddings",
api_key="test-key",
)
response = _requests_embedding_response(400, b'{"error":"bad embedding input"}')
calls = []
def fake_post(*_args, **_kwargs):
calls.append(1)
return response
monkeypatch.setattr("yuxi.models.embed.requests.post", fake_post)
with pytest.raises(ValueError, match="400 Client Error"):
model.encode(["hello", "test"])
assert len(calls) == 1
assert len(warnings) == 1
warning = warnings[0]
assert "400 Bad Request" in warning
assert "model=namespace/embedding-model" in warning
assert "input_count=2" in warning
assert "input_lengths=[5, 4]" in warning
assert "bad embedding input" in warning
def test_embedding_sync_429_retries_ten_times_before_success(monkeypatch):
warnings = _capture_embed_warnings(monkeypatch)
sleeps = []
monkeypatch.setattr("yuxi.models.embed.time.sleep", sleeps.append)
model = OtherEmbedding(
model="namespace/embedding-model",
base_url="https://example.com/v1/embeddings",
api_key="test-key",
)
success = _requests_embedding_response(200, b'{"data":[{"embedding":[0.1,0.2]}]}')
responses = [_requests_embedding_response(429) for _ in range(10)] + [success]
monkeypatch.setattr("yuxi.models.embed.requests.post", lambda *_args, **_kwargs: responses.pop(0))
assert model.encode(["hello"]) == [[0.1, 0.2]]
assert len(sleeps) == 10
assert sleeps == [1.0, 2.0, 4.0, 8.0, 10.0, 10.0, 10.0, 10.0, 10.0, 10.0]
assert len(warnings) == 10
assert "status=429" in warnings[-1]
assert "retry=10/10" in warnings[-1]
def test_embedding_sync_5xx_uses_short_retry_budget(monkeypatch):
warnings = _capture_embed_warnings(monkeypatch)
sleeps = []
calls = []
monkeypatch.setattr("yuxi.models.embed.time.sleep", sleeps.append)
model = OtherEmbedding(
model="namespace/embedding-model",
base_url="https://example.com/v1/embeddings",
api_key="test-key",
)
def fake_post(*_args, **_kwargs):
calls.append(1)
return _requests_embedding_response(503)
monkeypatch.setattr("yuxi.models.embed.requests.post", fake_post)
with pytest.raises(ValueError, match="503 Server Error"):
model.encode(["hello"])
assert len(calls) == 3
assert sleeps == [1.0, 2.0]
assert len(warnings) == 2
assert "retry=2/2" in warnings[-1]
@pytest.mark.asyncio
async def test_embedding_async_400_logs_warning(monkeypatch):
warnings = _capture_embed_warnings(monkeypatch)
model = OtherEmbedding(
model="namespace/embedding-model",
base_url="https://example.com/v1/embeddings",
api_key="test-key",
)
class FakeAsyncClient:
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
return False
async def post(self, url, **_kwargs):
request = httpx.Request("POST", url)
return httpx.Response(400, request=request, text='{"error":"bad embedding input"}')
monkeypatch.setattr("yuxi.models.embed.httpx.AsyncClient", FakeAsyncClient)
with pytest.raises(httpx.HTTPStatusError, match="400 Bad Request"):
await model.aencode(["hello", "test"])
assert len(warnings) == 1
warning = warnings[0]
assert "400 Bad Request" in warning
assert "model=namespace/embedding-model" in warning
assert "input_count=2" in warning
assert "input_lengths=[5, 4]" in warning
assert "bad embedding input" in warning
@pytest.mark.asyncio
async def test_embedding_async_429_retries_ten_times_before_success(monkeypatch):
warnings = _capture_embed_warnings(monkeypatch)
sleeps = []
async def fake_sleep(delay):
sleeps.append(delay)
monkeypatch.setattr("yuxi.models.embed.asyncio.sleep", fake_sleep)
model = OtherEmbedding(
model="namespace/embedding-model",
base_url="https://example.com/v1/embeddings",
api_key="test-key",
)
success = _httpx_embedding_response(200, '{"data":[{"embedding":[0.1,0.2]}]}')
responses = [_httpx_embedding_response(429) for _ in range(10)] + [success]
class FakeAsyncClient:
async def __aenter__(self):
return self
async def __aexit__(self, exc_type, exc_val, exc_tb):
return False
async def post(self, *_args, **_kwargs):
return responses.pop(0)
monkeypatch.setattr("yuxi.models.embed.httpx.AsyncClient", FakeAsyncClient)
assert await model.aencode(["hello"]) == [[0.1, 0.2]]
assert sleeps == [1.0, 2.0, 4.0, 8.0, 10.0, 10.0, 10.0, 10.0, 10.0, 10.0]
assert len(warnings) == 10
assert "status=429" in warnings[-1]
assert "retry=10/10" in warnings[-1]
def test_get_reranker_loads_model_from_cache(monkeypatch):
monkeypatch.setattr(
"yuxi.models.rerank.model_cache.get_model_info",