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
3 changes: 2 additions & 1 deletion spacy_llm/models/rest/__init__.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
from . import anthropic, azure, base, cohere, noop, openai
from . import anthropic, azure, base, cohere, litellm, noop, openai

__all__ = [
"anthropic",
"azure",
"base",
"cohere",
"litellm",
"openai",
"noop",
]
3 changes: 3 additions & 0 deletions spacy_llm/models/rest/litellm/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
from .model import LiteLLM

__all__ = ["LiteLLM"]
67 changes: 67 additions & 0 deletions spacy_llm/models/rest/litellm/model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
import os
import warnings
from typing import Any, Dict, Iterable, List, Optional

import srsly # type: ignore[import]

from ..base import REST


class LiteLLM(REST):
"""Queries LLMs via the LiteLLM Python SDK, supporting 100+ providers."""

@property
def credentials(self) -> Dict[str, str]:
api_key = os.getenv("LITELLM_API_KEY", "")
if not api_key:
warnings.warn(
"No LITELLM_API_KEY found. LiteLLM will attempt to read provider-specific "
"API keys from environment variables (e.g. OPENAI_API_KEY, ANTHROPIC_API_KEY)."
)
return {"api_key": api_key}

def _verify_auth(self) -> None:
pass

def __call__(self, prompts: Iterable[Iterable[str]]) -> Iterable[Iterable[str]]:
import litellm

all_api_responses: List[List[str]] = []

for prompts_for_doc in prompts:
api_responses: List[str] = []
prompts_for_doc = list(prompts_for_doc)

for prompt in prompts_for_doc:
kwargs: Dict[str, Any] = {
"model": self._name,
"messages": [{"role": "user", "content": prompt}],
"drop_params": True,
**self._config,
}
api_key = self._credentials.get("api_key")
if api_key:
kwargs["api_key"] = api_key
if self._endpoint:
kwargs["api_base"] = self._endpoint

try:
response = litellm.completion(**kwargs)
api_responses.append(
response.choices[0].message.content or ""
)
except Exception as e:
if self._strict:
raise ValueError(
f"Request to LiteLLM API failed: {e}"
) from e
else:
api_responses.append(srsly.json_dumps({"error": str(e)}))

all_api_responses.append(api_responses)

return all_api_responses

@staticmethod
def _get_context_lengths() -> Dict[str, int]:
return {}
42 changes: 42 additions & 0 deletions spacy_llm/models/rest/litellm/registry.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,42 @@
from typing import Any, Dict, Optional

from confection import SimpleFrozenDict

from ....registry import registry
from .model import LiteLLM

_DEFAULT_TEMPERATURE = 0.0


@registry.llm_models("spacy.LiteLLM.v1")
def litellm_v1(
config: Dict[Any, Any] = SimpleFrozenDict(temperature=_DEFAULT_TEMPERATURE),
name: str = "openai/gpt-4o",
strict: bool = LiteLLM.DEFAULT_STRICT,
max_tries: int = LiteLLM.DEFAULT_MAX_TRIES,
interval: float = LiteLLM.DEFAULT_INTERVAL,
max_request_time: float = LiteLLM.DEFAULT_MAX_REQUEST_TIME,
endpoint: Optional[str] = None,
context_length: Optional[int] = None,
) -> LiteLLM:
"""Returns LiteLLM instance for any model supported by LiteLLM (100+ providers).

Uses the LiteLLM Python SDK to route requests to any LLM provider.
Model names follow the litellm format: provider/model-name
(e.g. "anthropic/claude-haiku-4-5", "openai/gpt-4o", "bedrock/anthropic.claude-3-haiku").

config (Dict[Any, Any]): LLM config passed on to the model's initialization.
name (str): Model name in litellm format (provider/model-name).
context_length (Optional[int]): Context length for this model.
RETURNS (LiteLLM): LiteLLM instance.
"""
return LiteLLM(
name=name,
endpoint=endpoint or "",
config=config,
strict=strict,
max_tries=max_tries,
interval=interval,
max_request_time=max_request_time,
context_length=context_length,
)