Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions deepsearcher/config.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,12 @@ provide_settings:
# model: "deepseek/deepseek-v3-0324"
## api_key: "sk_xxxxxx" # Uncomment to override the `NOVITA_API_KEY` set in the environment variable

# provider: "LiteLLM"
# config:
# model: "anthropic/claude-sonnet-4-20250514" # Any LiteLLM-supported model
## api_key: "sk-xxxx" # Uncomment to override the `LITELLM_API_KEY` set in the environment variable
## api_base: "http://localhost:4000" # Uncomment to use a LiteLLM proxy server

embedding:
provider: "OpenAIEmbedding"
config:
Expand Down Expand Up @@ -105,6 +111,13 @@ provide_settings:
# config:
# model: "BAAI/bge-large-zh-v1.5"

# provider: "LiteLLMEmbedding"
# config:
# model: "text-embedding-ada-002" # Any LiteLLM-supported embedding model
## api_key: "sk-xxxx" # Uncomment to override the `LITELLM_API_KEY` set in the environment variable
## api_base: "http://localhost:4000" # Uncomment to use a LiteLLM proxy server
## dimension: 1536

file_loader:
provider: "PDFLoader"
config: {}
Expand Down
2 changes: 2 additions & 0 deletions deepsearcher/embedding/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from .gemini_embedding import GeminiEmbedding
from .glm_embedding import GLMEmbedding
from .jiekouai_embedding import JiekouAIEmbedding
from .litellm_embedding import LiteLLMEmbedding
from .milvus_embedding import MilvusEmbedding
from .novita_embedding import NovitaEmbedding
from .ollama_embedding import OllamaEmbedding
Expand Down Expand Up @@ -30,4 +31,5 @@
"SentenceTransformerEmbedding",
"WatsonXEmbedding",
"JiekouAIEmbedding",
"LiteLLMEmbedding",
]
109 changes: 109 additions & 0 deletions deepsearcher/embedding/litellm_embedding.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
import os
from typing import List

from deepsearcher.embedding.base import BaseEmbedding


class LiteLLMEmbedding(BaseEmbedding):
"""
LiteLLM embedding model implementation.

This class provides a unified interface to embedding models from 100+ providers
(OpenAI, Cohere, Bedrock, Vertex AI, and more) through the LiteLLM AI gateway.

API Documentation: https://docs.litellm.ai/docs/embedding/supported_embedding
"""

def __init__(self, model: str = "text-embedding-ada-002", **kwargs):
"""
Initialize the LiteLLM embedding model.

Args:
model (str): The model identifier to use for embeddings. Follows LiteLLM naming
conventions (e.g., "text-embedding-ada-002", "cohere/embed-english-v3.0",
"bedrock/amazon.titan-embed-text-v2:0"). Defaults to "text-embedding-ada-002".
**kwargs: Additional keyword arguments.
- api_key (str, optional): API key for the provider or LiteLLM proxy.
If not provided, uses LITELLM_API_KEY environment variable.
When not set, LiteLLM reads provider-specific env vars automatically.
- api_base (str, optional): Base URL for a LiteLLM proxy server.
If not provided, uses LITELLM_API_BASE environment variable.
- dimension (int, optional): The dimension of the embedding vectors.
Defaults to 1536.
"""
self.model = model
if "api_key" in kwargs:
self.api_key = kwargs.pop("api_key")
else:
self.api_key = os.getenv("LITELLM_API_KEY")
if "api_base" in kwargs:
self.api_base = kwargs.pop("api_base")
else:
self.api_base = os.getenv("LITELLM_API_BASE")
if "dimension" in kwargs:
self.dim = kwargs.pop("dimension")
else:
self.dim = 1536
self.kwargs = kwargs

def embed_query(self, text: str) -> List[float]:
"""
Embed a single query text.

Args:
text (str): The query text to embed.

Returns:
List[float]: A list of floats representing the embedding vector.
"""
import litellm

kwargs = {**self.kwargs}
if self.api_key:
kwargs["api_key"] = self.api_key
if self.api_base:
kwargs["api_base"] = self.api_base

response = litellm.embedding(
model=self.model,
input=[text],
drop_params=True,
**kwargs,
)
return response.data[0]["embedding"]

def embed_documents(self, texts: List[str]) -> List[List[float]]:
"""
Embed a list of document texts.

Args:
texts (List[str]): A list of document texts to embed.

Returns:
List[List[float]]: A list of embedding vectors, one for each input text.
"""
import litellm

kwargs = {**self.kwargs}
if self.api_key:
kwargs["api_key"] = self.api_key
if self.api_base:
kwargs["api_base"] = self.api_base

response = litellm.embedding(
model=self.model,
input=texts,
drop_params=True,
**kwargs,
)
return [r["embedding"] for r in response.data]

@property
def dimension(self) -> int:
"""
Get the dimensionality of the embeddings.

Returns:
int: The number of dimensions in the embedding vectors.
"""
return self.dim
2 changes: 2 additions & 0 deletions deepsearcher/llm/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from .gemini import Gemini
from .glm import GLM
from .jiekouai import JiekouAI
from .litellm_llm import LiteLLM
from .novita import Novita
from .ollama import Ollama
from .openai_llm import OpenAI
Expand Down Expand Up @@ -34,4 +35,5 @@
"Aliyun",
"WatsonX",
"JiekouAI",
"LiteLLM",
]
78 changes: 78 additions & 0 deletions deepsearcher/llm/litellm_llm.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
import os
from typing import Dict, List

from deepsearcher.llm.base import BaseLLM, ChatResponse


class LiteLLM(BaseLLM):
"""
LiteLLM language model implementation.

This class provides a unified interface to 100+ LLM providers (OpenAI, Anthropic,
Google, Azure, Bedrock, Ollama, and more) through the LiteLLM AI gateway.

API Documentation: https://docs.litellm.ai/docs/providers

Attributes:
model (str): The LiteLLM model identifier (e.g., "gpt-4o", "anthropic/claude-sonnet-4-20250514").
api_key (str): Optional API key. When not set, LiteLLM uses provider-specific
environment variables (e.g., OPENAI_API_KEY, ANTHROPIC_API_KEY).
api_base (str): Optional base URL for LiteLLM proxy or custom endpoints.
"""

def __init__(self, model: str = "gpt-4o-mini", **kwargs):
"""
Initialize a LiteLLM language model client.

Args:
model (str, optional): The model identifier to use. Follows LiteLLM naming
conventions (e.g., "gpt-4o", "anthropic/claude-sonnet-4-20250514",
"bedrock/anthropic.claude-v2"). Defaults to "gpt-4o-mini".
**kwargs: Additional keyword arguments passed to litellm.completion().
- api_key: API key for the provider or LiteLLM proxy.
If not provided, uses LITELLM_API_KEY environment variable.
When not set, LiteLLM reads provider-specific env vars automatically.
- api_base: Base URL for a LiteLLM proxy server.
If not provided, uses LITELLM_API_BASE environment variable.
"""
self.model = model
if "api_key" in kwargs:
self.api_key = kwargs.pop("api_key")
else:
self.api_key = os.getenv("LITELLM_API_KEY")
if "api_base" in kwargs:
self.api_base = kwargs.pop("api_base")
else:
self.api_base = os.getenv("LITELLM_API_BASE")
self.kwargs = kwargs

def chat(self, messages: List[Dict]) -> ChatResponse:
"""
Send a chat message to the language model and get a response.

Args:
messages (List[Dict]): A list of message dictionaries, typically in the format
[{"role": "system", "content": "..."},
{"role": "user", "content": "..."}]

Returns:
ChatResponse: An object containing the model's response and token usage information.
"""
import litellm

kwargs = {**self.kwargs}
if self.api_key:
kwargs["api_key"] = self.api_key
if self.api_base:
kwargs["api_base"] = self.api_base

response = litellm.completion(
model=self.model,
messages=messages,
drop_params=True,
**kwargs,
)
return ChatResponse(
content=response.choices[0].message.content,
total_tokens=response.usage.total_tokens,
)
7 changes: 6 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -60,7 +60,8 @@ all = [
"docling-core>=2.30.0",
"crawl4ai>=0.6.2",
"sentence-transformers>=4.1.0",
"ibm-watsonx-ai>=1.3.0"
"ibm-watsonx-ai>=1.3.0",
"litellm>=1.80.0,<1.87.0"
]

voyageai = [
Expand Down Expand Up @@ -114,6 +115,10 @@ ibm-watsonx = [
"ibm-watsonx-ai>=1.3.0"
]

litellm = [
"litellm>=1.80.0,<1.87.0",
]

[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
Expand Down
134 changes: 134 additions & 0 deletions tests/embedding/test_litellm_embedding.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,134 @@
import logging
import os
import unittest
from unittest.mock import MagicMock, patch

logging.disable(logging.CRITICAL)

from deepsearcher.embedding import LiteLLMEmbedding


class TestLiteLLMEmbedding(unittest.TestCase):
"""Tests for the LiteLLMEmbedding class."""

def setUp(self):
"""Set up test fixtures."""
self.mock_litellm = MagicMock()

mock_data_item = {"embedding": [0.1] * 1536}
self.mock_response = MagicMock()
self.mock_response.data = [mock_data_item]
self.mock_litellm.embedding.return_value = self.mock_response

self.module_patcher = patch.dict("sys.modules", {"litellm": self.mock_litellm})
self.module_patcher.start()

def tearDown(self):
"""Clean up test fixtures."""
self.module_patcher.stop()

def test_init_default(self):
"""Test initialization with default parameters."""
with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding()
self.assertEqual(embedding.model, "text-embedding-ada-002")
self.assertEqual(embedding.dim, 1536)
self.assertIsNone(embedding.api_key)
self.assertIsNone(embedding.api_base)

def test_init_with_env_vars(self):
"""Test initialization with environment variables."""
with patch.dict(
os.environ,
{"LITELLM_API_KEY": "test-key", "LITELLM_API_BASE": "http://localhost:4000"},
):
embedding = LiteLLMEmbedding()
self.assertEqual(embedding.api_key, "test-key")
self.assertEqual(embedding.api_base, "http://localhost:4000")

def test_init_with_parameters(self):
"""Test initialization with parameters."""
with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding(
model="cohere/embed-english-v3.0",
api_key="param-key",
api_base="http://proxy:4000",
dimension=1024,
)
self.assertEqual(embedding.model, "cohere/embed-english-v3.0")
self.assertEqual(embedding.api_key, "param-key")
self.assertEqual(embedding.api_base, "http://proxy:4000")
self.assertEqual(embedding.dim, 1024)

def test_embed_query(self):
"""Test embedding a single query."""
with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding()

result = embedding.embed_query("test query")

self.mock_litellm.embedding.assert_called_once_with(
model="text-embedding-ada-002",
input=["test query"],
drop_params=True,
)
self.assertEqual(result, [0.1] * 1536)

def test_embed_documents(self):
"""Test embedding multiple documents."""
mock_data_items = [
{"embedding": [0.1 * (i + 1)] * 1536} for i in range(3)
]
self.mock_response.data = mock_data_items

with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding()

texts = ["text 1", "text 2", "text 3"]
results = embedding.embed_documents(texts)

self.mock_litellm.embedding.assert_called_once_with(
model="text-embedding-ada-002",
input=texts,
drop_params=True,
)
self.assertEqual(len(results), 3)

def test_embed_query_with_credentials(self):
"""Test embed_query passes api_key and api_base when set."""
with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding(api_key="my-key", api_base="http://proxy:4000")

embedding.embed_query("test")

self.mock_litellm.embedding.assert_called_once_with(
model="text-embedding-ada-002",
input=["test"],
drop_params=True,
api_key="my-key",
api_base="http://proxy:4000",
)

def test_embed_query_omits_credentials_when_not_set(self):
"""Test embed_query omits api_key and api_base when not set."""
with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding()

embedding.embed_query("test")

call_kwargs = self.mock_litellm.embedding.call_args[1]
self.assertNotIn("api_key", call_kwargs)
self.assertNotIn("api_base", call_kwargs)

def test_dimension_property(self):
"""Test the dimension property."""
with patch.dict("os.environ", {}, clear=True):
embedding = LiteLLMEmbedding()
self.assertEqual(embedding.dimension, 1536)

embedding = LiteLLMEmbedding(dimension=768)
self.assertEqual(embedding.dimension, 768)


if __name__ == "__main__":
unittest.main()
Loading