diff --git a/.dockerignore b/.dockerignore new file mode 100644 index 0000000000..0ddd6d5e11 --- /dev/null +++ b/.dockerignore @@ -0,0 +1,19 @@ +.git +.venv +.coverage +.pytest_cache +__pycache__ +*.py[cod] +*.egg-info +test_data +data + +node_modules +web-studio/node_modules +web-studio/dist +openviking/web_studio/dist + +.mypy_cache +.ruff_cache +.pytest_cache +.DS_Store diff --git a/docs/en/guides/01-configuration.md b/docs/en/guides/01-configuration.md index 67a4080a38..5d2e6e491b 100644 --- a/docs/en/guides/01-configuration.md +++ b/docs/en/guides/01-configuration.md @@ -809,7 +809,7 @@ The `PermissionDeniedError` message names the exact key to add for the blocked h ### rerank -Reranking model for search result refinement. Supports VikingDB (Volcengine), Cohere, and OpenAI-compatible APIs. +Reranking model for search result refinement. Supports VikingDB (Volcengine), Cohere, OpenAI-compatible APIs, LiteLLM, and Hugging Face Text Embeddings Inference (TEI). **Volcengine (VikingDB):** @@ -840,25 +840,45 @@ Reranking model for search result refinement. Supports VikingDB (Volcengine), Co } ``` +**Hugging Face Text Embeddings Inference (TEI):** + +```json +{ + "rerank": { + "provider": "tei", + "api_base": "http://localhost:8080", + "api_key": "optional-tei-api-key", + "model": "BAAI/bge-reranker-v2-m3", + "batch_size": 32, + "threshold": 0.05 + } +} +``` + +For TEI, `api_base` may be either the server base URL (`http://localhost:8080`) or the full rerank endpoint (`http://localhost:8080/rerank`). `api_key` is optional and is sent as a Bearer token when configured. TEI is auto-detected when only `api_base` is set; if your TEI deployment also uses `api_key`, set `"provider": "tei"` explicitly so it is not treated as an OpenAI-compatible rerank endpoint. + **Parameters** | Parameter | Type | Description | |-----------|------|-------------| -| `provider` | str | `"vikingdb"`, `"cohere"`, or `"openai"`. Auto-detected if omitted. | +| `provider` | str | `"vikingdb"`, `"cohere"`, `"openai"`, `"litellm"`, or `"tei"`. Auto-detected if omitted. | | `ak` | str | VikingDB Access Key (vikingdb provider only) | | `sk` | str | VikingDB Secret Key (vikingdb provider only) | | `model_name` | str | Model name (vikingdb provider only, default: `doubao-seed-rerank`) | -| `api_key` | str | API key (for `openai` or `cohere` providers) | -| `api_base` | str | Endpoint URL (for `openai` provider) | -| `model` | str | Model name (for `openai` providers) | +| `api_key` | str | API key (for `openai` or `cohere` providers, optional for `tei`) | +| `api_base` | str | Endpoint URL (for `openai` provider) or TEI base/rerank URL (for `tei`) | +| `model` | str | Model name (for `openai` and `litellm`; optional label for TEI usage tracking) | | `timeout` | float | HTTP request timeout in seconds for OpenAI-compatible providers. Increase for slow or cold-starting local rerank servers. Default: `30.0` | +| `batch_size` | int | Maximum number of documents sent in a single rerank provider call. TEI deployments commonly cap this at `32`; larger candidate sets are chunked. Default: `32` | | `threshold` | float | Score threshold between `0.0` and `1.0`; results below this are filtered out. Default: `0.1` | -| `extra_headers` | object | Custom HTTP headers (for OpenAI-compatible providers, optional) | +| `extra_headers` | object | Custom HTTP headers (for OpenAI-compatible or TEI providers, optional) | **Supported providers:** - `vikingdb`: Volcengine VikingDB Rerank API (uses AK/SK) - `cohere`: Cohere Rerank API - `openai`: OpenAI-compatible Rerank API +- `litellm`: LiteLLM rerank API +- `tei`: Hugging Face Text Embeddings Inference rerank API If rerank is not configured, search uses vector similarity only. diff --git a/openviking/models/rerank/__init__.py b/openviking/models/rerank/__init__.py index 2cf5ae3a3e..17958b8b97 100644 --- a/openviking/models/rerank/__init__.py +++ b/openviking/models/rerank/__init__.py @@ -8,12 +8,14 @@ - cohere: Cohere Rerank v3.5 API - litellm: LiteLLM rerank (supports multiple providers) - openai: OpenAI-compatible rerank API +- tei: Hugging Face Text Embeddings Inference rerank API """ from openviking.models.rerank.base import RerankBase from openviking.models.rerank.cohere_rerank import CohereRerankClient from openviking.models.rerank.litellm_rerank import LiteLLMRerankClient from openviking.models.rerank.openai_rerank import OpenAIRerankClient +from openviking.models.rerank.tei_rerank import TEIRerankClient from openviking.models.rerank.volcengine_rerank import RerankClient __all__ = [ @@ -22,4 +24,5 @@ "CohereRerankClient", "LiteLLMRerankClient", "OpenAIRerankClient", + "TEIRerankClient", ] diff --git a/openviking/models/rerank/tei_rerank.py b/openviking/models/rerank/tei_rerank.py new file mode 100644 index 0000000000..bd1a2a0562 --- /dev/null +++ b/openviking/models/rerank/tei_rerank.py @@ -0,0 +1,178 @@ +# Copyright (c) 2026 Beijing Volcano Engine Technology Co., Ltd. +# SPDX-License-Identifier: AGPL-3.0 +""" +Text Embeddings Inference rerank API client. + +Hugging Face Text Embeddings Inference (TEI) exposes rerank models through a +provider-specific `/rerank` endpoint. Its request/response shape differs from +OpenAI-compatible rerank APIs, so it needs a dedicated adapter. +""" + +import time +from typing import Dict, List, Optional + +import requests + +from openviking.models.rerank.base import RerankBase +from openviking_cli.utils import get_logger + +logger = get_logger(__name__) + + +class TEIRerankClient(RerankBase): + """ + TEI rerank API client. + + TEI accepts `texts` and returns a list of `{index, score}` items: + https://huggingface.co/docs/text-embeddings-inference + """ + + def __init__( + self, + api_base: str, + api_key: Optional[str] = None, + model_name: Optional[str] = None, + extra_headers: Optional[Dict[str, str]] = None, + batch_size: int = 32, + ) -> None: + """ + Initialize TEI rerank client. + + Args: + api_base: TEI base URL (`http://host:port`) or full rerank endpoint. + api_key: Optional Bearer token for TEI deployments that enforce auth. + model_name: Optional model name used for usage tracking. + extra_headers: Optional extra headers for API requests. + batch_size: Maximum number of documents to send per TEI request. + """ + super().__init__() + self.api_base = api_base + self.api_key = api_key + self.model_name = model_name + self.extra_headers = extra_headers or {} + self.batch_size = max(1, int(batch_size)) + self.provider = "tei" + + @property + def rerank_url(self) -> str: + """Return the full TEI rerank URL while accepting base or endpoint config.""" + base = self.api_base.rstrip("/") + if base.endswith("/rerank"): + return base + return f"{base}/rerank" + + def rerank_batch(self, query: str, documents: List[str]) -> Optional[List[float]]: + """ + Batch rerank documents against a query. + + Args: + query: Query text + documents: List of document texts to rank + + Returns: + List of rerank scores in the same order as input documents, or None + when rerank fails and the caller should fall back. + """ + if not documents: + return [] + + scores = [0.0] * len(documents) + for start in range(0, len(documents), self.batch_size): + chunk = documents[start : start + self.batch_size] + chunk_scores = self._rerank_chunk(query, chunk) + if chunk_scores is None: + return None + scores[start : start + len(chunk_scores)] = chunk_scores + + logger.debug( + "[TEIRerankClient] Reranked %s documents in %s request(s)", + len(documents), + (len(documents) + self.batch_size - 1) // self.batch_size, + ) + return scores + + def _rerank_chunk(self, query: str, documents: List[str]) -> Optional[List[float]]: + """Rerank one TEI-sized chunk and return scores in chunk-local order.""" + + req_body = { + "query": query, + "texts": documents, + "raw_scores": False, + } + + try: + headers = {"Content-Type": "application/json"} + if self.api_key: + headers["Authorization"] = f"Bearer {self.api_key}" + if self.extra_headers: + headers.update(self.extra_headers) + + started = time.monotonic() + response = requests.post( + url=self.rerank_url, + headers=headers, + json=req_body, + timeout=30, + ) + response.raise_for_status() + result = response.json() + + self._extract_and_update_token_usage( + {"results": result} if isinstance(result, list) else result, + query, + documents, + duration_seconds=time.monotonic() - started, + ) + + results = self._extract_results(result) + if not results: + logger.warning(f"[TEIRerankClient] Unexpected response format: {result}") + return None + + scores = [0.0] * len(documents) + for item in results: + idx = item.get("index") + if idx is None or not (0 <= idx < len(documents)): + logger.warning( + "[TEIRerankClient] Out-of-bounds or missing index in result: %s", item + ) + return None + scores[idx] = float(item.get("score", item.get("relevance_score", 0.0))) + + return scores + + except Exception as e: + logger.error(f"[TEIRerankClient] Rerank failed: {e}") + return None + + @staticmethod + def _extract_results(result) -> Optional[List[dict]]: + """Extract TEI rerank rows from supported response shapes.""" + if isinstance(result, list): + return result + if isinstance(result, dict): + rows = result.get("results") + if isinstance(rows, list): + return rows + return None + + @classmethod + def from_config(cls, config) -> Optional["TEIRerankClient"]: + """ + Create TEIRerankClient from RerankConfig. + + Args: + config: RerankConfig instance with provider='tei' + + Returns: + TEIRerankClient instance or None if config is not available + """ + if not config or not config.is_available(): + return None + return cls( + api_base=config.api_base, + api_key=config.api_key, + model_name=config.model, + extra_headers=config.extra_headers, + batch_size=config.batch_size, + ) diff --git a/openviking/models/rerank/volcengine_rerank.py b/openviking/models/rerank/volcengine_rerank.py index 92a5a55ab8..af6e9feb26 100644 --- a/openviking/models/rerank/volcengine_rerank.py +++ b/openviking/models/rerank/volcengine_rerank.py @@ -219,6 +219,11 @@ def from_config(cls, config) -> Optional["RerankClient"]: return OpenAIRerankClient.from_config(config) + if provider == "tei": + from openviking.models.rerank.tei_rerank import TEIRerankClient + + return TEIRerankClient.from_config(config) + return cls( ak=config.ak, sk=config.sk, diff --git a/openviking/server/routers/admin.py b/openviking/server/routers/admin.py index 0e08ed3a86..e6efa08554 100644 --- a/openviking/server/routers/admin.py +++ b/openviking/server/routers/admin.py @@ -66,6 +66,25 @@ class MigrateLegacyDataRequest(BaseModel): action: str = "migrate" +class _LegacyCleanupUserIdentifier: + """Minimal user identity for deleting pre-validation legacy accounts.""" + + def __init__(self, account_id: str, user_id: str): + self._account_id = account_id + self._user_id = user_id + + @property + def account_id(self) -> str: + return self._account_id + + @property + def user_id(self) -> str: + return self._user_id + + def user_space_name(self) -> str: + return self._user_id + + def _get_api_key_manager(request: Request): """Get APIKeyManager from app state.""" return get_api_key_manager_or_raise(request) @@ -133,6 +152,17 @@ def _validate_register_user_role(ctx: RequestContext, role: str) -> Role: return resolved_role +def _cleanup_user_identifier(account_id: str) -> UserIdentifier | _LegacyCleanupUserIdentifier: + """Build a cleanup identity, allowing already-existing legacy account ids.""" + try: + return UserIdentifier(account_id, "system") + except ValueError: + logger.warning( + "Using legacy cleanup identity for non-conforming account_id=%s", account_id + ) + return _LegacyCleanupUserIdentifier(account_id, "system") + + async def _run_legacy_migration_task( task_id: str, migration: LegacyDataMigration, @@ -264,7 +294,7 @@ async def delete_account( # Build a ROOT-level context scoped to the target account for cleanup cleanup_ctx = RequestContext( - user=UserIdentifier(account_id, "system"), + user=_cleanup_user_identifier(account_id), role=Role.ROOT, ) diff --git a/openviking_cli/utils/config/rerank_config.py b/openviking_cli/utils/config/rerank_config.py index 6cc5803ad1..90287a4001 100644 --- a/openviking_cli/utils/config/rerank_config.py +++ b/openviking_cli/utils/config/rerank_config.py @@ -6,11 +6,11 @@ class RerankConfig(BaseModel): - """Configuration for rerank API. Supports VikingDB, Cohere, OpenAI-compatible, and LiteLLM providers.""" + """Configuration for rerank API. Supports VikingDB, Cohere, OpenAI-compatible, LiteLLM, and TEI providers.""" provider: Optional[str] = Field( default=None, - description="Rerank provider: 'vikingdb', 'cohere', 'openai', or 'litellm'. Auto-detected from config if omitted.", + description="Rerank provider: 'vikingdb', 'cohere', 'openai', 'litellm', or 'tei'. Auto-detected from config if omitted.", ) # VikingDB fields @@ -22,17 +22,18 @@ class RerankConfig(BaseModel): model_name: str = Field(default="doubao-seed-rerank", description="Rerank model name") model_version: str = Field(default="251028", description="Rerank model version") - # Shared / OpenAI-compatible / Cohere fields + # Shared / OpenAI-compatible / Cohere / TEI fields api_key: Optional[str] = Field( - default=None, description="API key (Cohere Bearer token or OpenAI-compatible providers)" + default=None, + description="API key (Cohere Bearer token, OpenAI-compatible providers, or optional TEI auth)", ) api_base: Optional[str] = Field(default=None, description="Custom endpoint URL") model: Optional[str] = Field( - default=None, description="Model name for OpenAI-compatible or LiteLLM providers" + default=None, description="Model name for OpenAI-compatible, LiteLLM, or TEI providers" ) extra_headers: Optional[Dict[str, str]] = Field( - default=None, description="Extra HTTP headers for OpenAI-compatible providers" + default=None, description="Extra HTTP headers for OpenAI-compatible or TEI providers" ) timeout: float = Field( @@ -43,6 +44,15 @@ class RerankConfig(BaseModel): ), ) + batch_size: int = Field( + default=32, + ge=1, + description=( + "Maximum number of documents to send in a single rerank provider call. " + "TEI deployments commonly cap this at 32; larger candidate sets are chunked." + ), + ) + threshold: float = Field( default=0.1, description="Relevance threshold (score > threshold is relevant)" ) @@ -59,14 +69,16 @@ def _effective_provider(self) -> Optional[str]: return "cohere" if self.ak and self.sk: return "vikingdb" + if self.api_base: + return "tei" return None @model_validator(mode="after") def validate_provider_fields(self) -> "RerankConfig": provider = self._effective_provider() - if provider and provider not in ["vikingdb", "cohere", "openai", "litellm"]: + if provider and provider not in ["vikingdb", "cohere", "openai", "litellm", "tei"]: raise ValueError( - f"Rerank provider must be one of ['vikingdb', 'cohere', 'openai', 'litellm'], got '{provider}'" + f"Rerank provider must be one of ['vikingdb', 'cohere', 'openai', 'litellm', 'tei'], got '{provider}'" ) if provider == "openai": if not self.api_key or not self.api_base: @@ -76,6 +88,9 @@ def validate_provider_fields(self) -> "RerankConfig": if provider == "litellm": if not self.model: raise ValueError("LiteLLM rerank provider requires 'model'") + if provider == "tei": + if not self.api_base: + raise ValueError("TEI rerank provider requires 'api_base'") return self def is_available(self) -> bool: @@ -87,6 +102,8 @@ def is_available(self) -> bool: return self.api_key is not None and self.api_base is not None if p == "litellm": return self.model is not None + if p == "tei": + return self.api_base is not None if p == "vikingdb": return self.ak is not None and self.sk is not None return False diff --git a/tests/misc/test_rerank_openai.py b/tests/misc/test_rerank_openai.py index b835950ef7..31f068c90b 100644 --- a/tests/misc/test_rerank_openai.py +++ b/tests/misc/test_rerank_openai.py @@ -251,10 +251,11 @@ def test_openai_requires_api_key_and_api_base(self): with pytest.raises(ValidationError): RerankConfig(provider="openai", api_base="https://example.com/rerank") - def test_default_provider_is_vikingdb(self): + def test_default_provider_is_auto_detected(self): config = RerankConfig() - assert config.provider == "vikingdb" + assert config.provider is None + assert config._effective_provider() is None def test_unknown_provider_raises_value_error(self): - with pytest.raises(ValueError, match="provider"): - RerankConfig(provider="cohere", ak="ak", sk="sk") + with pytest.raises(ValidationError, match="provider"): + RerankConfig(provider="unknown") diff --git a/tests/server/test_admin_api.py b/tests/server/test_admin_api.py index ac4699d159..65aef74b1c 100644 --- a/tests/server/test_admin_api.py +++ b/tests/server/test_admin_api.py @@ -17,6 +17,7 @@ from openviking.pyagfs.exceptions import AGFSNotFoundError from openviking.server.api_keys import APIKeyManager +from openviking.server.api_keys.models import AccountInfo from openviking.server.app import create_app from openviking.server.config import ServerConfig from openviking.server.dependencies import set_service @@ -350,6 +351,31 @@ async def test_delete_account(admin_client: httpx.AsyncClient): assert resp.status_code == 401 +async def test_delete_legacy_nonconforming_account( + admin_app: FastAPI, admin_client: httpx.AsyncClient +): + """ROOT can delete legacy account ids that predate account id validation.""" + acct = "chatwoot:1" + manager = admin_app.state.api_key_manager + manager._legacy._accounts[acct] = AccountInfo( + created_at="2026-03-30T08:54:32.006457+00:00", + users={}, + ) + await manager._legacy._save_accounts_json() + await manager._legacy._save_users_json(acct) + + resp = await admin_client.delete( + f"/api/v1/admin/accounts/{acct}", headers=root_headers() + ) + assert resp.status_code == 200 + assert resp.json()["result"]["deleted"] is True + + resp = await admin_client.get("/api/v1/admin/accounts", headers=root_headers()) + accounts = resp.json()["result"] + account_ids = {a["account_id"] for a in accounts} + assert acct not in account_ids + + async def test_create_duplicate_account_fails(admin_client: httpx.AsyncClient): """Creating duplicate account should fail.""" acct = _uid() diff --git a/tests/unit/test_tei_rerank.py b/tests/unit/test_tei_rerank.py new file mode 100644 index 0000000000..3ef55679fd --- /dev/null +++ b/tests/unit/test_tei_rerank.py @@ -0,0 +1,173 @@ +# Copyright (c) 2026 Beijing Volcano Engine Technology Co., Ltd. +# SPDX-License-Identifier: AGPL-3.0 +"""Tests for Hugging Face Text Embeddings Inference rerank client.""" + +from unittest.mock import MagicMock, patch + +from openviking.models.rerank import RerankClient, TEIRerankClient +from openviking_cli.utils.config.rerank_config import RerankConfig + + +class TestTEIRerankClient: + """Test cases for TEIRerankClient.""" + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_basic(self, mock_post): + mock_response = MagicMock() + mock_response.json.return_value = [ + {"index": 1, "score": 0.95}, + {"index": 0, "score": 0.42}, + {"index": 2, "score": 0.10}, + ] + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + client = TEIRerankClient( + api_base="http://tei.local:8080", + api_key="test-key", + model_name="BAAI/bge-reranker-v2-m3", + ) + scores = client.rerank_batch("What is UCW?", ["doc A", "doc B", "doc C"]) + + assert scores == [0.42, 0.95, 0.10] + kwargs = mock_post.call_args.kwargs + assert kwargs["url"] == "http://tei.local:8080/rerank" + assert kwargs["headers"]["Authorization"] == "Bearer test-key" + assert kwargs["json"] == { + "query": "What is UCW?", + "texts": ["doc A", "doc B", "doc C"], + "raw_scores": False, + } + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_accepts_full_endpoint(self, mock_post): + mock_response = MagicMock() + mock_response.json.return_value = [{"index": 0, "score": 0.9}] + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + client = TEIRerankClient(api_base="http://tei.local:8080/rerank") + assert client.rerank_batch("q", ["doc"]) == [0.9] + assert mock_post.call_args.kwargs["url"] == "http://tei.local:8080/rerank" + assert "Authorization" not in mock_post.call_args.kwargs["headers"] + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_fills_missing_top_n_results(self, mock_post): + mock_response = MagicMock() + mock_response.json.return_value = [{"index": 2, "score": 0.99}] + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + client = TEIRerankClient(api_base="http://tei.local:8080") + assert client.rerank_batch("q", ["first", "second", "third"]) == [0.0, 0.0, 0.99] + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_supports_results_wrapper(self, mock_post): + mock_response = MagicMock() + mock_response.json.return_value = { + "results": [ + {"index": 0, "relevance_score": 0.7}, + {"index": 1, "relevance_score": 0.2}, + ] + } + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + client = TEIRerankClient(api_base="http://tei.local:8080") + assert client.rerank_batch("q", ["first", "second"]) == [0.7, 0.2] + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_invalid_index_returns_none(self, mock_post): + mock_response = MagicMock() + mock_response.json.return_value = [{"index": 10, "score": 0.99}] + mock_response.raise_for_status = MagicMock() + mock_post.return_value = mock_response + + client = TEIRerankClient(api_base="http://tei.local:8080") + assert client.rerank_batch("q", ["doc"]) is None + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_api_error_returns_none(self, mock_post): + mock_response = MagicMock() + mock_response.raise_for_status.side_effect = RuntimeError("boom") + mock_post.return_value = mock_response + + client = TEIRerankClient(api_base="http://tei.local:8080") + assert client.rerank_batch("q", ["doc"]) is None + + def test_rerank_batch_empty(self): + client = TEIRerankClient(api_base="http://tei.local:8080") + assert client.rerank_batch("query", []) == [] + + @patch("openviking.models.rerank.tei_rerank.requests.post") + def test_rerank_batch_chunks_documents(self, mock_post): + first_response = MagicMock() + first_response.json.return_value = [ + {"index": 1, "score": 0.2}, + {"index": 0, "score": 0.1}, + ] + first_response.raise_for_status = MagicMock() + second_response = MagicMock() + second_response.json.return_value = [{"index": 0, "score": 0.3}] + second_response.raise_for_status = MagicMock() + mock_post.side_effect = [first_response, second_response] + + client = TEIRerankClient(api_base="http://tei.local:8080", batch_size=2) + + assert client.rerank_batch("q", ["a", "b", "c"]) == [0.1, 0.2, 0.3] + assert mock_post.call_count == 2 + assert mock_post.call_args_list[0].kwargs["json"]["texts"] == ["a", "b"] + assert mock_post.call_args_list[1].kwargs["json"]["texts"] == ["c"] + + +class TestTEIRerankConfig: + """Test TEI rerank config parsing and dispatch.""" + + def test_config_requires_api_base(self): + import pytest + from pydantic import ValidationError + + with pytest.raises(ValidationError): + RerankConfig(provider="tei") + + def test_config_available_without_api_key(self): + config = RerankConfig(provider="tei", api_base="http://tei.local:8080") + + assert config._effective_provider() == "tei" + assert config.is_available() is True + + def test_config_auto_detects_tei_without_api_key(self): + config = RerankConfig(api_base="http://tei.local:8080") + + assert config._effective_provider() == "tei" + assert config.is_available() is True + + def test_config_api_key_and_api_base_auto_detects_openai(self): + config = RerankConfig(api_key="key", api_base="https://example.com/rerank") + + assert config._effective_provider() == "openai" + assert config.is_available() is True + + def test_from_config_creates_tei_client(self): + config = RerankConfig( + provider="tei", + api_base="http://tei.local:8080", + api_key="key", + model="BAAI/bge-reranker-v2-m3", + extra_headers={"X-Test": "1"}, + batch_size=16, + ) + + client = RerankClient.from_config(config) + + assert isinstance(client, TEIRerankClient) + assert client.api_base == "http://tei.local:8080" + assert client.api_key == "key" + assert client.model_name == "BAAI/bge-reranker-v2-m3" + assert client.extra_headers == {"X-Test": "1"} + assert client.batch_size == 16 + + def test_config_default_batch_size_is_tei_safe(self): + config = RerankConfig(provider="tei", api_base="http://tei.local:8080") + + assert config.batch_size == 32 diff --git a/uv.lock b/uv.lock index f888c0a186..0f96930b6c 100644 --- a/uv.lock +++ b/uv.lock @@ -3679,34 +3679,6 @@ benchmark = [ { name = "tiktoken" }, ] bot = [ - { name = "beautifulsoup4" }, - { name = "croniter" }, - { name = "ddgs" }, - { name = "gradio" }, - { name = "html2text" }, - { name = "httpx", extra = ["socks"] }, - { name = "mcp" }, - { name = "msgpack" }, - { name = "prompt-toolkit" }, - { name = "py-machineid" }, - { name = "pydantic-settings" }, - { name = "pygments" }, - { name = "python-socketio" }, - { name = "python-socks", extra = ["asyncio"] }, - { name = "readability-lxml" }, - { name = "rich" }, - { name = "socksio" }, - { name = "tavily-python" }, - { name = "websocket-client" }, - { name = "websockets" }, -] -bot-dingtalk = [ - { name = "dingtalk-stream" }, -] -bot-feishu = [ - { name = "lark-oapi" }, -] -bot-full = [ { name = "agent-sandbox" }, { name = "beautifulsoup4" }, { name = "croniter" }, @@ -3739,29 +3711,6 @@ bot-full = [ { name = "websocket-client" }, { name = "websockets" }, ] -bot-fuse = [ - { name = "fusepy" }, -] -bot-langfuse = [ - { name = "langfuse" }, -] -bot-opencode = [ - { name = "opencode-ai" }, -] -bot-qq = [ - { name = "qq-botpy" }, -] -bot-sandbox = [ - { name = "agent-sandbox" }, - { name = "opensandbox" }, - { name = "opensandbox-server" }, -] -bot-slack = [ - { name = "slack-sdk" }, -] -bot-telegram = [ - { name = "python-telegram-bot", extra = ["socks"] }, -] build = [ { name = "build" }, { name = "cmake" }, @@ -3833,7 +3782,7 @@ dev = [ [package.metadata] requires-dist = [ - { name = "agent-sandbox", marker = "extra == 'bot-sandbox'", specifier = ">=0.0.23" }, + { name = "agent-sandbox", marker = "extra == 'bot'", specifier = ">=0.0.23" }, { name = "anyio", marker = "extra == 'gemini-async'", specifier = ">=4.0.0" }, { name = "apscheduler", specifier = ">=3.11.0" }, { name = "argon2-cffi", specifier = ">=23.0.0" }, @@ -3850,11 +3799,11 @@ requires-dist = [ { name = "ddgs", marker = "extra == 'bot'", specifier = ">=9.0.0" }, { name = "defusedxml", specifier = ">=0.7.1" }, { name = "diff-match-patch", marker = "extra == 'test'", specifier = ">=20200713" }, - { name = "dingtalk-stream", marker = "extra == 'bot-dingtalk'", specifier = ">=0.4.0" }, + { name = "dingtalk-stream", marker = "extra == 'bot'", specifier = ">=0.4.0" }, { name = "ebooklib", specifier = ">=0.18.0" }, { name = "fastapi", specifier = ">=0.128.0" }, { name = "feedparser", specifier = ">=6.0.0" }, - { name = "fusepy", marker = "extra == 'bot-fuse'", specifier = ">=3.0.1" }, + { name = "fusepy", marker = "extra == 'bot'", specifier = ">=3.0.1" }, { name = "google-genai", marker = "extra == 'gemini'", specifier = ">=1.0.0" }, { name = "google-genai", marker = "extra == 'gemini-async'", specifier = ">=1.0.0" }, { name = "gradio", marker = "extra == 'bot'", specifier = ">=6.6.0" }, @@ -3870,11 +3819,11 @@ requires-dist = [ { name = "langchain-core", marker = "extra == 'langchain'", specifier = ">=1.0.0,<2.0.0" }, { name = "langchain-core", marker = "extra == 'langgraph'", specifier = ">=1.0.0,<2.0.0" }, { name = "langchain-openai", marker = "extra == 'benchmark'", specifier = ">=1.0.0" }, - { name = "langfuse", marker = "extra == 'bot-langfuse'", specifier = ">=3.0.0" }, + { name = "langfuse", marker = "extra == 'bot'", specifier = ">=3.0.0" }, { name = "langgraph", marker = "extra == 'langgraph'", specifier = ">=1.0.0,<2.0.0" }, { name = "lark-oapi", specifier = ">=1.5.3" }, - { name = "lark-oapi", marker = "extra == 'bot-feishu'", specifier = ">=1.0.0" }, - { name = "litellm", specifier = ">=1.83.7,<1.89.3" }, + { name = "lark-oapi", marker = "extra == 'bot'", specifier = ">=1.0.0" }, + { name = "litellm", specifier = ">=1.83.7,<1.90.3" }, { name = "llama-cpp-python", marker = "extra == 'local-embed'", specifier = ">=0.3.0" }, { name = "loguru", specifier = ">=0.7.3" }, { name = "mcp", specifier = ">=1.27.0" }, @@ -3884,16 +3833,15 @@ requires-dist = [ { name = "myst-parser", marker = "extra == 'doc'", specifier = ">=2.0.0" }, { name = "olefile", specifier = ">=0.47" }, { name = "openai", specifier = ">=1.0.0" }, - { name = "opencode-ai", marker = "extra == 'bot-opencode'", specifier = ">=0.1.0a0" }, + { name = "opencode-ai", marker = "extra == 'bot'", specifier = ">=0.1.0a0" }, { name = "openpyxl", specifier = ">=3.0.0" }, - { name = "opensandbox", marker = "extra == 'bot-sandbox'", specifier = ">=0.1.0" }, - { name = "opensandbox-server", marker = "extra == 'bot-sandbox'", specifier = ">=0.1.0" }, + { name = "opensandbox", marker = "extra == 'bot'", specifier = ">=0.1.0" }, + { name = "opensandbox-server", marker = "extra == 'bot'", specifier = ">=0.1.0" }, { name = "opentelemetry-api", specifier = ">=1.14" }, { name = "opentelemetry-exporter-otlp-proto-grpc", specifier = ">=1.14" }, { name = "opentelemetry-exporter-otlp-proto-http", specifier = ">=1.14" }, { name = "opentelemetry-instrumentation-asyncio", specifier = ">=0.61b0" }, { name = "opentelemetry-sdk", specifier = ">=1.14" }, - { name = "openviking", extras = ["bot", "bot-dingtalk", "bot-feishu", "bot-fuse", "bot-langfuse", "bot-opencode", "bot-qq", "bot-sandbox", "bot-slack", "bot-telegram"], marker = "extra == 'bot-full'" }, { name = "openviking-sdk", specifier = ">=0.1.1" }, { name = "pandas", marker = "extra == 'benchmark'", specifier = ">=2.0.0" }, { name = "pandas", marker = "extra == 'eval'", specifier = ">=2.0.0" }, @@ -3918,9 +3866,9 @@ requires-dist = [ { name = "python-pptx", specifier = ">=1.0.0" }, { name = "python-socketio", marker = "extra == 'bot'", specifier = ">=5.11.0" }, { name = "python-socks", extras = ["asyncio"], marker = "extra == 'bot'", specifier = ">=2.4.0" }, - { name = "python-telegram-bot", extras = ["socks"], marker = "extra == 'bot-telegram'", specifier = ">=21.0" }, + { name = "python-telegram-bot", extras = ["socks"], marker = "extra == 'bot'", specifier = ">=21.0" }, { name = "pyyaml", specifier = ">=6.0" }, - { name = "qq-botpy", marker = "extra == 'bot-qq'", specifier = ">=1.0.0" }, + { name = "qq-botpy", marker = "extra == 'bot'", specifier = ">=1.0.0" }, { name = "ragas", marker = "extra == 'eval'", specifier = ">=0.1.0" }, { name = "ragas", marker = "extra == 'test'", specifier = ">=0.1.0" }, { name = "readability-lxml", marker = "extra == 'bot'", specifier = ">=0.8.0" }, @@ -3931,7 +3879,7 @@ requires-dist = [ { name = "setuptools", marker = "extra == 'build'", specifier = ">=61.0" }, { name = "setuptools-scm", marker = "extra == 'build'", specifier = ">=8.0" }, { name = "setuptools-scm", marker = "extra == 'dev'", specifier = ">=10.0.0" }, - { name = "slack-sdk", marker = "extra == 'bot-slack'", specifier = ">=3.26.0" }, + { name = "slack-sdk", marker = "extra == 'bot'", specifier = ">=3.26.0" }, { name = "socksio", marker = "extra == 'bot'", specifier = ">=1.0.0" }, { name = "sphinx", marker = "extra == 'doc'", specifier = ">=7.0.0" }, { name = "sphinx-rtd-theme", marker = "extra == 'doc'", specifier = ">=1.3.0" }, @@ -3962,7 +3910,7 @@ requires-dist = [ { name = "xlrd", specifier = ">=2.0.1" }, { name = "xxhash", specifier = ">=3.0.0" }, ] -provides-extras = ["test", "opengauss", "dev", "doc", "eval", "gemini", "gemini-async", "ocr", "build", "bot", "bot-langfuse", "bot-telegram", "bot-feishu", "bot-dingtalk", "bot-slack", "bot-qq", "bot-sandbox", "bot-fuse", "bot-opencode", "bot-full", "benchmark", "langchain", "langgraph", "local-embed"] +provides-extras = ["test", "opengauss", "dev", "doc", "eval", "gemini", "gemini-async", "ocr", "build", "bot", "benchmark", "langchain", "langgraph", "local-embed"] [package.metadata.requires-dev] dev = [{ name = "pytest", specifier = ">=9.0.2" }]