diff --git a/astrbot/core/config/default.py b/astrbot/core/config/default.py index ff3ccd4f96..9028d84909 100644 --- a/astrbot/core/config/default.py +++ b/astrbot/core/config/default.py @@ -1841,6 +1841,8 @@ "embedding_model": "", "embedding_dimensions": 1024, "embedding_dimensions_mode": "auto", + "embedding_max_requests_per_minute": 120, + "embedding_rate_limit_cooldown": 120, "timeout": 20, "proxy": "", }, @@ -1855,6 +1857,8 @@ "embedding_api_base": "", "embedding_model": "gemini-embedding-exp-03-07", "embedding_dimensions": 768, + "embedding_max_requests_per_minute": 120, + "embedding_rate_limit_cooldown": 120, "timeout": 20, "proxy": "", }, @@ -1870,6 +1874,8 @@ "embedding_model": "nvidia/llama-nemotron-embed-1b-v2", "input_type": "passage", "embedding_dimensions": 1024, + "embedding_max_requests_per_minute": 120, + "embedding_rate_limit_cooldown": 120, "timeout": 20, "proxy": "", }, @@ -1883,6 +1889,8 @@ "embedding_api_base": "http://localhost:11434", "embedding_model": "nomic-embed-text", "embedding_dimensions": 768, + "embedding_max_requests_per_minute": 0, + "embedding_rate_limit_cooldown": 120, "timeout": 60, "proxy": "", }, @@ -2267,6 +2275,16 @@ "description": "API Base URL", "type": "string", }, + "embedding_max_requests_per_minute": { + "description": "Embedding 每分钟最大请求数", + "type": "int", + "hint": "限制同一 Embedding Provider 的批量向量化请求速率。设置为 0 表示不限制;远程 API 建议按服务商限额填写。", + }, + "embedding_rate_limit_cooldown": { + "description": "Embedding 限速冷却上限(秒)", + "type": "int", + "hint": "当 Embedding API 返回 Retry-After 或限速错误时,同一 Provider 后续请求的最大等待秒数。", + }, "volcengine_cluster": { "type": "string", "description": "火山引擎集群", diff --git a/astrbot/core/db/vec_db/faiss_impl/vec_db.py b/astrbot/core/db/vec_db/faiss_impl/vec_db.py index c641dd56a4..06c54d43d1 100644 --- a/astrbot/core/db/vec_db/faiss_impl/vec_db.py +++ b/astrbot/core/db/vec_db/faiss_impl/vec_db.py @@ -46,7 +46,7 @@ async def insert( metadata = metadata or {} str_id = id or str(uuid.uuid4()) # 使用 UUID 作为原始 ID - vector = await self.embedding_provider.get_embedding(content) + vector = await self.embedding_provider.get_embedding_with_retry(content) vector = np.array(vector, dtype=np.float32) # 使用 DocumentStorage 的方法插入文档 @@ -109,13 +109,26 @@ async def insert_batch( start = time.time() logger.debug(f"Generating embeddings for {len(contents)} contents...") - vectors = await self.embedding_provider.get_embeddings_batch( - contents, - batch_size=batch_size, - tasks_limit=tasks_limit, - max_retries=max_retries, - progress_callback=progress_callback, - ) + try: + vectors = await self.embedding_provider.get_embeddings_batch( + contents, + batch_size=batch_size, + tasks_limit=tasks_limit, + max_retries=max_retries, + progress_callback=progress_callback, + ) + except KnowledgeBaseUploadError: + raise + except Exception as exc: + raise KnowledgeBaseUploadError( + stage="embedding", + user_message=( + "向量化失败:调用 Embedding API 生成向量时出错。" + "请检查 Embedding 服务是否可用、是否触发限流," + "并尝试降低批量大小或并发数后重试。" + ), + details={"cause": str(exc)}, + ) from exc end = time.time() logger.debug( f"Generated embeddings for {len(contents)} contents in {end - start:.2f} seconds.", @@ -217,7 +230,7 @@ async def retrieve( List[Result]: 查询结果 """ - embedding = await self.embedding_provider.get_embedding(query) + embedding = await self.embedding_provider.get_embedding_with_retry(query) scores, indices = await self.embedding_storage.search( vector=np.array(embedding).astype("float32"), k=fetch_k if metadata_filters else k, diff --git a/astrbot/core/provider/provider.py b/astrbot/core/provider/provider.py index 891bfdea9e..ab01586589 100644 --- a/astrbot/core/provider/provider.py +++ b/astrbot/core/provider/provider.py @@ -1,9 +1,14 @@ import abc import asyncio import os +import random +import time from collections.abc import AsyncGenerator +from datetime import datetime, timezone +from email.utils import parsedate_to_datetime from typing import Literal, TypeAlias, Union +from astrbot import logger from astrbot.core.agent.message import ContentPart, Message, is_checkpoint_message from astrbot.core.agent.tool import ToolSet from astrbot.core.provider.entities import ( @@ -14,6 +19,8 @@ ) from astrbot.core.provider.register import provider_cls_map from astrbot.core.utils.astrbot_path import get_astrbot_path +from astrbot.core.utils.config_number import coerce_int_config +from astrbot.core.utils.network_utils import is_connection_error Providers: TypeAlias = Union[ "Provider", @@ -317,11 +324,51 @@ async def test(self) -> None: pass +class EmbeddingProviderError(Exception): + """Represent an Embedding API failure with HTTP response metadata. + + Args: + message: Human-readable provider error message. + status_code: HTTP status code returned by the provider, if available. + response: Original provider response, used to read headers such as + ``Retry-After``. + """ + + def __init__( + self, + message: str, + *, + status_code: int | None = None, + response: object | None = None, + ) -> None: + super().__init__(message) + self.status_code = status_code + self.response = response + + class EmbeddingProvider(AbstractProvider): def __init__(self, provider_config: dict, provider_settings: dict) -> None: super().__init__(provider_config) self.provider_config = provider_config self.provider_settings = provider_settings + max_rpm = coerce_int_config( + provider_config.get("embedding_max_requests_per_minute", 0), + default=0, + min_value=0, + field_name="embedding_max_requests_per_minute", + source=f"Embedding provider {provider_config.get('id', 'unknown')}", + ) + cooldown_max_s = coerce_int_config( + provider_config.get("embedding_rate_limit_cooldown", 120), + default=120, + min_value=1, + field_name="embedding_rate_limit_cooldown", + source=f"Embedding provider {provider_config.get('id', 'unknown')}", + ) + self._embedding_min_request_interval_s = 60.0 / max_rpm if max_rpm else 0.0 + self._embedding_rate_limit_cooldown_max_s = float(cooldown_max_s) + self._embedding_request_lock = asyncio.Lock() + self._embedding_next_request_at = 0.0 @abc.abstractmethod async def get_embedding(self, text: str) -> list[float]: @@ -339,7 +386,69 @@ def get_dim(self) -> int: ... async def test(self) -> None: - await self.get_embedding("astrbot") + await self.get_embedding_with_retry("astrbot") + + async def get_embedding_with_retry( + self, + text: str, + *, + max_retries: int = 3, + ) -> list[float]: + """Get one embedding through the shared throttled retry path. + + Args: + text: Input text to embed. + max_retries: Maximum attempts, including the first try. + + Returns: + The embedding vector. + + Raises: + Exception: When the provider returns no embedding or retries fail. + """ + embeddings = await self.get_embeddings_batch( + [text], + batch_size=1, + tasks_limit=1, + max_retries=max_retries, + ) + if not embeddings: + raise Exception("Embedding provider returned no vectors.") + return embeddings[0] + + async def _wait_for_embedding_request_slot(self) -> None: + """Wait for the provider-level embedding request throttle. + + This limiter is shared by all knowledge-base upload tasks using the same + provider instance, so separate large imports cannot bypass per-provider + pacing or a dynamic rate-limit cooldown by running concurrently. + """ + async with self._embedding_request_lock: + while True: + now = time.monotonic() + wait_s = self._embedding_next_request_at - now + if wait_s <= 0: + break + await asyncio.sleep(wait_s) + if self._embedding_min_request_interval_s > 0: + self._embedding_next_request_at = ( + time.monotonic() + self._embedding_min_request_interval_s + ) + + def _delay_embedding_requests(self, delay_s: float) -> None: + """Delay future embedding requests for this provider instance. + + Args: + delay_s: Requested delay in seconds, typically from Retry-After or + a computed rate-limit retry backoff. + """ + if delay_s <= 0: + return + capped_delay = min(delay_s, self._embedding_rate_limit_cooldown_max_s) + self._embedding_next_request_at = max( + self._embedding_next_request_at, + time.monotonic() + capped_delay, + ) async def get_embeddings_batch( self, @@ -349,30 +458,113 @@ async def get_embeddings_batch( max_retries: int = 3, progress_callback=None, ) -> list[list[float]]: - """批量获取文本的向量,分批处理以节省内存 + """Batch-embed texts with bounded concurrency and retry. Args: - texts: 文本列表 - batch_size: 每批处理的文本数量 - tasks_limit: 并发任务数量限制 - max_retries: 失败时的最大重试次数 - progress_callback: 进度回调函数,接收参数 (current, total) + texts: Input texts to embed. + batch_size: Number of texts per provider request. + tasks_limit: Maximum number of concurrent batch requests. + max_retries: Maximum attempts per batch (including the first try). + progress_callback: Optional async callback ``(current, total)``. Returns: - 向量列表 + Embedding vectors in the same order as ``texts``. + Raises: + Exception: When one or more batches fail after all retries. """ + batch_size = max(1, int(batch_size or 1)) + tasks_limit = max(1, int(tasks_limit or 1)) + max_retries = max(1, int(max_retries or 1)) + + retry_status_codes = {408, 409, 429, 500, 502, 503, 504, 529} + retry_backoff_max_s = 30.0 + semaphore = asyncio.Semaphore(tasks_limit) batch_results: dict[int, list[list[float]]] = {} - failed_batches: list[tuple[int, list[str]]] = [] completed_count = 0 total_count = len(texts) + def _get_status_code(error: BaseException) -> int | None: + for attr in ("status_code", "status", "code"): + value = getattr(error, attr, None) + if isinstance(value, int): + return value + + response = getattr(error, "response", None) + if response is not None: + status_code = getattr(response, "status_code", None) + if isinstance(status_code, int): + return status_code + return None + + def _get_retry_after_seconds(error: BaseException) -> float | None: + candidates: list[object] = [ + getattr(error, "retry_after", None), + getattr(error, "Retry-After", None), + ] + response = getattr(error, "response", None) + headers = ( + getattr(response, "headers", None) if response is not None else None + ) + if headers is not None: + try: + candidates.append(headers.get("Retry-After")) + candidates.append(headers.get("retry-after")) + except Exception: + pass + + for value in candidates: + if value is None: + continue + try: + seconds = float(value) + except (TypeError, ValueError): + try: + retry_at = parsedate_to_datetime(str(value)) + except (TypeError, ValueError): + continue + if retry_at.tzinfo is None: + retry_at = retry_at.replace(tzinfo=timezone.utc) + seconds = ( + retry_at - datetime.now(tz=retry_at.tzinfo) + ).total_seconds() + if seconds >= 0: + return seconds + return None + + def _is_retryable_embedding_error(error: BaseException) -> bool: + if is_connection_error(error): + return True + + error_type_name = type(error).__name__ + if error_type_name in {"APIConnectionError", "APITimeoutError"}: + return True + + status_code = _get_status_code(error) + if status_code is None: + # Preserve historical forgiving behavior: retry unknown errors. + return True + return status_code in retry_status_codes or 500 <= status_code <= 599 + + def _retry_delay_seconds(attempt: int, error: BaseException) -> float: + retry_after = _get_retry_after_seconds(error) + if retry_after is not None: + return min(retry_after, self._embedding_rate_limit_cooldown_max_s) + + # Keep an exponential lower bound so rate-limit retries never happen + # immediately, then add a small jitter to avoid synchronized waves. + base = min(float(2**attempt), retry_backoff_max_s) + jitter = random.uniform(0.0, min(1.0, base)) + return min(base + jitter, retry_backoff_max_s) + async def process_batch(batch_idx: int, batch_texts: list[str]) -> None: nonlocal completed_count async with semaphore: + last_error: Exception | None = None for attempt in range(max_retries): try: + await self._wait_for_embedding_request_slot() batch_embeddings = await self.get_embeddings(batch_texts) batch_results[batch_idx] = batch_embeddings completed_count += len(batch_texts) @@ -380,25 +572,46 @@ async def process_batch(batch_idx: int, batch_texts: list[str]) -> None: await progress_callback(completed_count, total_count) return except Exception as e: - if attempt == max_retries - 1: - # 最后一次重试失败,记录失败的批次 - failed_batches.append((batch_idx, batch_texts)) - raise Exception( - f"批次 {batch_idx} 处理失败,已重试 {max_retries} 次: {e!s}", - ) - # 等待一段时间后重试,使用指数退避 - await asyncio.sleep(2**attempt) + last_error = e + if ( + attempt >= max_retries - 1 + or not _is_retryable_embedding_error(e) + ): + break + + delay = _retry_delay_seconds(attempt, e) + retry_after = _get_retry_after_seconds(e) + status_code = _get_status_code(e) + if ( + retry_after is not None + or status_code in retry_status_codes + or (status_code is not None and 500 <= status_code <= 599) + ): + self._delay_embedding_requests(delay) + logger.warning( + "Embedding batch %s failed (attempt %s/%s): %s; " + "retrying in %.2fs", + batch_idx, + attempt + 1, + max_retries, + e, + delay, + ) + await asyncio.sleep(delay) + + raise Exception( + f"批次 {batch_idx} 处理失败,共尝试 {attempt + 1} 次: " + f"{last_error!s}", + ) from last_error tasks = [] - for i in range(0, len(texts), batch_size): + for batch_idx, i in enumerate(range(0, len(texts), batch_size)): batch_texts = texts[i : i + batch_size] - batch_idx = i // batch_size - tasks.append(process_batch(batch_idx, batch_texts)) + tasks.append(asyncio.create_task(process_batch(batch_idx, batch_texts))) - # 收集所有任务的结果,包括失败的任务 + # Collect all task outcomes, including failures. results = await asyncio.gather(*tasks, return_exceptions=True) - # 检查是否有失败的任务 errors = [r for r in results if isinstance(r, Exception)] if errors: error_msg = ( diff --git a/astrbot/core/provider/sources/gemini_embedding_source.py b/astrbot/core/provider/sources/gemini_embedding_source.py index 71e9dadc9d..3554632a31 100644 --- a/astrbot/core/provider/sources/gemini_embedding_source.py +++ b/astrbot/core/provider/sources/gemini_embedding_source.py @@ -5,7 +5,7 @@ from astrbot import logger from ..entities import ProviderType -from ..provider import EmbeddingProvider +from ..provider import EmbeddingProvider, EmbeddingProviderError from ..register import register_provider_adapter @@ -54,7 +54,11 @@ async def get_embedding(self, text: str) -> list[float]: assert result.embeddings[0].values is not None return result.embeddings[0].values except APIError as e: - raise Exception(f"Gemini Embedding API请求失败: {e.message}") + raise EmbeddingProviderError( + f"Gemini Embedding API请求失败: {e.message}", + status_code=e.code, + response=e.response, + ) from e async def get_embeddings(self, text: list[str]) -> list[list[float]]: """批量获取文本的嵌入""" @@ -77,7 +81,11 @@ async def get_embeddings(self, text: list[str]) -> list[list[float]]: embeddings.append(embedding.values) return embeddings except APIError as e: - raise Exception(f"Gemini Embedding API批量请求失败: {e.message}") + raise EmbeddingProviderError( + f"Gemini Embedding API批量请求失败: {e.message}", + status_code=e.code, + response=e.response, + ) from e def get_dim(self) -> int: """获取向量的维度""" diff --git a/astrbot/core/provider/sources/nvidia_embedding_source.py b/astrbot/core/provider/sources/nvidia_embedding_source.py index b13ee0d201..9eeb7f18a9 100644 --- a/astrbot/core/provider/sources/nvidia_embedding_source.py +++ b/astrbot/core/provider/sources/nvidia_embedding_source.py @@ -3,7 +3,7 @@ from astrbot import logger from ..entities import ProviderType -from ..provider import EmbeddingProvider +from ..provider import EmbeddingProvider, EmbeddingProviderError from ..register import register_provider_adapter @@ -96,8 +96,10 @@ async def get_embeddings(self, text: list[str]) -> list[list[float]]: logger.error( f"[NVIDIA Embedding] API Error: {response.status} - {error_text}" ) - raise Exception( - f"NVIDIA Embedding API request failed: HTTP {response.status} - {error_text}" + raise EmbeddingProviderError( + f"NVIDIA Embedding API request failed: HTTP {response.status} - {error_text}", + status_code=response.status, + response=response, ) response_data = await response.json() diff --git a/astrbot/core/provider/sources/ollama_embedding_source.py b/astrbot/core/provider/sources/ollama_embedding_source.py index 8982fc51de..f4dd25af1b 100644 --- a/astrbot/core/provider/sources/ollama_embedding_source.py +++ b/astrbot/core/provider/sources/ollama_embedding_source.py @@ -3,7 +3,7 @@ from astrbot import logger from ..entities import ProviderType -from ..provider import EmbeddingProvider +from ..provider import EmbeddingProvider, EmbeddingProviderError from ..register import register_provider_adapter @@ -82,8 +82,10 @@ async def get_embeddings(self, text: list[str]) -> list[list[float]]: logger.error( f"[Ollama Embedding] API Error: {response.status} - {error_text}" ) - raise Exception( - f"Ollama Embedding API request failed: HTTP {response.status} - {error_text}" + raise EmbeddingProviderError( + f"Ollama Embedding API request failed: HTTP {response.status} - {error_text}", + status_code=response.status, + response=response, ) response_data = await response.json() diff --git a/astrbot/dashboard/services/config_service.py b/astrbot/dashboard/services/config_service.py index a35dda8610..868247a8ba 100644 --- a/astrbot/dashboard/services/config_service.py +++ b/astrbot/dashboard/services/config_service.py @@ -1527,7 +1527,7 @@ async def get_embedding_dimension(self, provider_config: dict | None) -> dict: if not isinstance(inst, EmbeddingProvider): raise ValueError("提供商不是 EmbeddingProvider 类型") - vec = await inst.get_embedding("echo") + vec = await inst.get_embedding_with_retry("echo") dim = len(vec) logger.info( f"检测到 {provider_config.get('id', 'unknown')} 的嵌入向量维度为 {dim}", diff --git a/astrbot/dashboard/services/knowledge_base_service.py b/astrbot/dashboard/services/knowledge_base_service.py index 76ad91359d..04eacab596 100644 --- a/astrbot/dashboard/services/knowledge_base_service.py +++ b/astrbot/dashboard/services/knowledge_base_service.py @@ -325,7 +325,7 @@ async def create_kb(self, data: object) -> tuple[dict[str, Any], str]: f"嵌入模型不存在或类型错误({type(provider)})" ) try: - vec = await provider.get_embedding("astrbot") + vec = await provider.get_embedding_with_retry("astrbot") if len(vec) != provider.get_dim(): raise ValueError( f"嵌入向量维度不匹配,实际是 {len(vec)},然而配置是 {provider.get_dim()}", diff --git a/astrbot/dashboard/utils.py b/astrbot/dashboard/utils.py index 2f0189cd9f..c00f120b21 100644 --- a/astrbot/dashboard/utils.py +++ b/astrbot/dashboard/utils.py @@ -86,7 +86,7 @@ async def generate_tsne_visualization( # 获取查询向量 vec_db: FaissVecDB = kb_helper.vec_db # type: ignore embedding_provider = vec_db.embedding_provider - query_embedding = await embedding_provider.get_embedding(query) + query_embedding = await embedding_provider.get_embedding_with_retry(query) query_vector = np.array([query_embedding], dtype=np.float32) # 合并所有向量和查询向量 diff --git a/dashboard/src/i18n/locales/en-US/features/config-metadata.json b/dashboard/src/i18n/locales/en-US/features/config-metadata.json index 979be4fed4..8d34184014 100644 --- a/dashboard/src/i18n/locales/en-US/features/config-metadata.json +++ b/dashboard/src/i18n/locales/en-US/features/config-metadata.json @@ -1407,6 +1407,14 @@ "embedding_api_base": { "description": "API Base URL" }, + "embedding_max_requests_per_minute": { + "description": "Max embedding requests per minute", + "hint": "Limits batch embedding request rate for the same Embedding Provider. Set to 0 to disable; for remote APIs, use the provider's documented quota." + }, + "embedding_rate_limit_cooldown": { + "description": "Embedding rate-limit cooldown cap (seconds)", + "hint": "Maximum seconds to delay later requests for the same Provider when the Embedding API returns Retry-After or a rate-limit error." + }, "openai_embedding": { "hint": "If testing fails, try adding /v1 at the end for some OpenAI API versions." }, diff --git a/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json b/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json index be0731079a..6a6e58fddd 100644 --- a/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json +++ b/dashboard/src/i18n/locales/ru-RU/features/config-metadata.json @@ -1404,6 +1404,14 @@ "embedding_api_base": { "description": "Адрес прокси-сервера" }, + "embedding_max_requests_per_minute": { + "description": "Максимум запросов Embedding в минуту", + "hint": "Ограничивает частоту пакетных запросов для одного провайдера Embedding. Значение 0 отключает ограничение; для удалённых API укажите лимит провайдера." + }, + "embedding_rate_limit_cooldown": { + "description": "Максимальная пауза при ограничении Embedding (секунды)", + "hint": "Максимальная задержка последующих запросов к тому же провайдеру, если API Embedding возвращает Retry-After или ошибку ограничения частоты." + }, "openai_embedding": { "hint": "Если тест не проходит, попробуйте добавить /v1 в конец embedding_api_base для совместимости с некоторыми версиями OpenAI API." }, diff --git a/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json b/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json index 5036a51606..b29f713d7d 100644 --- a/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json +++ b/dashboard/src/i18n/locales/zh-CN/features/config-metadata.json @@ -1409,6 +1409,14 @@ "embedding_api_base": { "description": "API Base URL" }, + "embedding_max_requests_per_minute": { + "description": "Embedding 每分钟最大请求数", + "hint": "限制同一 Embedding Provider 的批量向量化请求速率。设置为 0 表示不限制;远程 API 建议按服务商限额填写。" + }, + "embedding_rate_limit_cooldown": { + "description": "Embedding 限速冷却上限(秒)", + "hint": "当 Embedding API 返回 Retry-After 或限速错误时,同一 Provider 后续请求的最大等待秒数。" + }, "openai_embedding": { "hint": "如果测试不通过,可以尝试添加 /v1 在末尾以兼容部分 OpenAI API 版本。" }, diff --git a/tests/unit/test_embedding_provider_batch.py b/tests/unit/test_embedding_provider_batch.py new file mode 100644 index 0000000000..19d3b851f7 --- /dev/null +++ b/tests/unit/test_embedding_provider_batch.py @@ -0,0 +1,396 @@ +import asyncio +from types import SimpleNamespace +from unittest.mock import AsyncMock + +import pytest +from google.genai.errors import APIError + +from astrbot.core.db.vec_db.faiss_impl.vec_db import FaissVecDB +from astrbot.core.exceptions import KnowledgeBaseUploadError +from astrbot.core.provider.provider import EmbeddingProvider, EmbeddingProviderError +from astrbot.core.provider.sources.gemini_embedding_source import ( + GeminiEmbeddingProvider, +) +from astrbot.core.provider.sources.nvidia_embedding_source import ( + NvidiaEmbeddingProvider, +) +from astrbot.core.provider.sources.ollama_embedding_source import ( + OllamaEmbeddingProvider, +) + + +class RecordingEmbeddingProvider(EmbeddingProvider): + """Embedding provider used to assert batch ordering and concurrency.""" + + def __init__(self, provider_config: dict | None = None) -> None: + super().__init__(provider_config or {}, {}) + self.calls: list[list[str]] = [] + self._fail_counts: dict[str, int] = {} + + def set_fail_count(self, first_text: str, count: int) -> None: + """Fail the batch that starts with first_text a fixed number of times. + + Args: + first_text: First text in the batch used as the batch key. + count: Number of transient failures before success. + """ + self._fail_counts[first_text] = count + + async def get_embedding(self, text: str) -> list[float]: + return [float(text.removeprefix("chunk-"))] + + async def get_embeddings(self, text: list[str]) -> list[list[float]]: + self.calls.append(list(text)) + first = text[0] + remaining = self._fail_counts.get(first, 0) + if remaining > 0: + self._fail_counts[first] = remaining - 1 + raise EmbeddingProviderError( + "transient rate limit", + status_code=429, + ) + return [[float(item.removeprefix("chunk-"))] for item in text] + + def get_dim(self) -> int: + return 1 + + +@pytest.fixture +def no_sleep(monkeypatch: pytest.MonkeyPatch): + """Disable asyncio.sleep waits used by batch staggering and retries.""" + + async def _instant_sleep(_delay: float = 0, *_args, **_kwargs) -> None: + return None + + monkeypatch.setattr(asyncio, "sleep", _instant_sleep) + # Keep random jitter deterministic for retry-delay code paths. + monkeypatch.setattr( + "astrbot.core.provider.provider.random.uniform", + lambda _a, _b: 0.0, + ) + + +@pytest.mark.asyncio +async def test_get_embeddings_batch_retries_transient_failure(no_sleep) -> None: + provider = RecordingEmbeddingProvider() + provider.set_fail_count("chunk-0", 1) + + embeddings = await provider.get_embeddings_batch( + ["chunk-0", "chunk-1"], + batch_size=2, + tasks_limit=1, + max_retries=3, + ) + + assert embeddings == [[0.0], [1.0]] + assert len(provider.calls) == 2 + + +@pytest.mark.asyncio +async def test_get_embedding_with_retry_uses_batch_retry_path(no_sleep) -> None: + provider = RecordingEmbeddingProvider() + provider.set_fail_count("chunk-0", 1) + + embedding = await provider.get_embedding_with_retry("chunk-0") + + assert embedding == [0.0] + assert provider.calls == [["chunk-0"], ["chunk-0"]] + + +@pytest.mark.asyncio +async def test_get_embeddings_batch_respects_retry_after_cooldown_cap( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = RecordingEmbeddingProvider( + {"embedding_rate_limit_cooldown": 120}, + ) + now = 100.0 + sleep_delays: list[float] = [] + + def _monotonic() -> float: + return now + + async def _advance_sleep(delay: float = 0, *_args, **_kwargs) -> None: + nonlocal now + sleep_delays.append(delay) + now += delay + + async def fail_once(text: list[str]) -> list[list[float]]: + provider.calls.append(list(text)) + if len(provider.calls) == 1: + raise EmbeddingProviderError( + "rate limited", + status_code=429, + response=SimpleNamespace(headers={"Retry-After": "60"}), + ) + return [[0.0]] + + monkeypatch.setattr("astrbot.core.provider.provider.time.monotonic", _monotonic) + monkeypatch.setattr(asyncio, "sleep", _advance_sleep) + provider.get_embeddings = fail_once # type: ignore[method-assign] + + embeddings = await provider.get_embeddings_batch( + ["chunk-0"], + batch_size=1, + tasks_limit=1, + max_retries=2, + ) + + assert embeddings == [[0.0]] + assert sleep_delays == [60.0] + assert provider._embedding_next_request_at == 160.0 + + +@pytest.mark.asyncio +async def test_get_embeddings_batch_uses_provider_level_request_pacing( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = RecordingEmbeddingProvider( + {"embedding_max_requests_per_minute": 60}, + ) + now = 100.0 + sleep_delays: list[float] = [] + + def _monotonic() -> float: + return now + + async def _advance_sleep(delay: float = 0, *_args, **_kwargs) -> None: + nonlocal now + sleep_delays.append(delay) + now += delay + + monkeypatch.setattr("astrbot.core.provider.provider.time.monotonic", _monotonic) + monkeypatch.setattr(asyncio, "sleep", _advance_sleep) + + await asyncio.gather( + provider.get_embeddings_batch( + ["chunk-0"], + batch_size=1, + tasks_limit=1, + ), + provider.get_embeddings_batch( + ["chunk-1"], + batch_size=1, + tasks_limit=1, + ), + ) + + assert sleep_delays == [1.0] + assert provider.calls == [["chunk-0"], ["chunk-1"]] + + +@pytest.mark.asyncio +async def test_concurrent_batches_honor_cooldown_without_rpm( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = RecordingEmbeddingProvider( + {"embedding_max_requests_per_minute": 0}, + ) + now = 100.0 + sleep_delays: list[float] = [] + + def _monotonic() -> float: + return now + + async def _advance_sleep(delay: float = 0, *_args, **_kwargs) -> None: + nonlocal now + sleep_delays.append(delay) + now += delay + + monkeypatch.setattr("astrbot.core.provider.provider.time.monotonic", _monotonic) + monkeypatch.setattr(asyncio, "sleep", _advance_sleep) + provider._delay_embedding_requests(10.0) + + await asyncio.gather( + provider.get_embeddings_batch(["chunk-0"], batch_size=1, tasks_limit=1), + provider.get_embeddings_batch(["chunk-1"], batch_size=1, tasks_limit=1), + ) + + assert sleep_delays == [10.0] + assert provider.calls == [["chunk-0"], ["chunk-1"]] + + +@pytest.mark.asyncio +async def test_provider_level_pacing_keeps_later_rate_limit_cooldown( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = RecordingEmbeddingProvider( + { + "embedding_max_requests_per_minute": 60, + "embedding_rate_limit_cooldown": 120, + }, + ) + now = 100.0 + sleep_delays: list[float] = [] + inject_cooldown = True + + def _monotonic() -> float: + return now + + async def _advance_sleep(delay: float = 0, *_args, **_kwargs) -> None: + nonlocal inject_cooldown, now + sleep_delays.append(delay) + if inject_cooldown: + provider._delay_embedding_requests(10.0) + inject_cooldown = False + now += delay + + monkeypatch.setattr("astrbot.core.provider.provider.time.monotonic", _monotonic) + monkeypatch.setattr(asyncio, "sleep", _advance_sleep) + provider._embedding_next_request_at = 101.0 + + await provider._wait_for_embedding_request_slot() + + assert sleep_delays == [1.0, 9.0] + + +@pytest.mark.asyncio +async def test_get_embeddings_batch_global_cooldown_for_5xx( + monkeypatch: pytest.MonkeyPatch, +) -> None: + provider = RecordingEmbeddingProvider( + {"embedding_max_requests_per_minute": 60}, + ) + now = 100.0 + + def _monotonic() -> float: + return now + + async def _advance_sleep(delay: float = 0, *_args, **_kwargs) -> None: + nonlocal now + now += delay + + async def fail_once(text: list[str]) -> list[list[float]]: + if not provider.calls: + provider.calls.append(list(text)) + raise EmbeddingProviderError( + "service unavailable", + status_code=503, + ) + provider.calls.append(list(text)) + return [[float(item.removeprefix("chunk-"))] for item in text] + + monkeypatch.setattr("astrbot.core.provider.provider.time.monotonic", _monotonic) + monkeypatch.setattr(asyncio, "sleep", _advance_sleep) + monkeypatch.setattr( + "astrbot.core.provider.provider.random.uniform", + lambda _a, _b: 0.0, + ) + provider.get_embeddings = fail_once # type: ignore[method-assign] + + embeddings = await provider.get_embeddings_batch( + ["chunk-0"], + batch_size=1, + tasks_limit=1, + max_retries=2, + ) + + assert embeddings == [[0.0]] + assert provider._embedding_next_request_at >= 101.0 + + +@pytest.mark.asyncio +async def test_faiss_insert_batch_classifies_provider_failure_as_embedding() -> None: + vec_db = FaissVecDB.__new__(FaissVecDB) + vec_db.embedding_provider = AsyncMock() + vec_db.embedding_provider.get_embeddings_batch.side_effect = RuntimeError( + "rate limited", + ) + vec_db.document_storage = AsyncMock() + vec_db.embedding_storage = AsyncMock() + + with pytest.raises(KnowledgeBaseUploadError) as exc_info: + await FaissVecDB.insert_batch( + vec_db, + contents=["hello world"], + metadatas=[{}], + ids=["doc-1"], + ) + + error = exc_info.value + assert error.stage == "embedding" + assert "向量化失败" in error.user_message + assert "Embedding API" in error.user_message + vec_db.document_storage.insert_documents_batch.assert_not_awaited() + vec_db.embedding_storage.insert_batch.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_get_embeddings_batch_does_not_retry_permanent_4xx(no_sleep) -> None: + provider = RecordingEmbeddingProvider() + + async def unauthorized(text: list[str]) -> list[list[float]]: + provider.calls.append(list(text)) + raise EmbeddingProviderError("unauthorized", status_code=401) + + provider.get_embeddings = unauthorized # type: ignore[method-assign] + + with pytest.raises(Exception, match="共尝试 1 次"): + await provider.get_embeddings_batch( + ["chunk-0"], + batch_size=1, + tasks_limit=1, + max_retries=3, + ) + + assert provider.calls == [["chunk-0"]] + + +@pytest.mark.asyncio +async def test_gemini_embedding_preserves_api_status() -> None: + provider = GeminiEmbeddingProvider.__new__(GeminiEmbeddingProvider) + provider.client = SimpleNamespace( + models=SimpleNamespace( + embed_content=AsyncMock( + side_effect=APIError(429, {"message": "rate limited"}), + ), + ), + ) + provider.model = "embedding-model" + provider.provider_config = {"embedding_dimensions": 768} + + with pytest.raises(EmbeddingProviderError) as exc_info: + await provider.get_embeddings(["chunk-0"]) + + assert exc_info.value.status_code == 429 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "provider_class", + [NvidiaEmbeddingProvider, OllamaEmbeddingProvider], +) +async def test_aiohttp_embedding_providers_preserve_api_status( + provider_class: type[EmbeddingProvider], +) -> None: + class RateLimitedResponse: + status = 429 + headers = {"Retry-After": "10"} + + async def __aenter__(self): + return self + + async def __aexit__(self, _exc_type, _exc, _traceback) -> None: + return None + + async def text(self) -> str: + return "rate limited" + + response = RateLimitedResponse() + provider = provider_class.__new__(provider_class) + provider.client = SimpleNamespace( + closed=False, + post=lambda *_args, **_kwargs: response, + ) + provider.base_url = "https://embedding.example.com" + provider.model = "embedding-model" + provider.proxy = "" + provider.provider_config = {} + if isinstance(provider, NvidiaEmbeddingProvider): + provider.input_type = "passage" + + with pytest.raises(EmbeddingProviderError) as exc_info: + await provider.get_embeddings(["chunk-0"]) + + assert exc_info.value.status_code == 429 + assert exc_info.value.response is response