Skip to content
Merged
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
10 changes: 8 additions & 2 deletions scripts/generate_api_docs.py
Original file line number Diff line number Diff line change
Expand Up @@ -434,9 +434,15 @@ def generate_any_llm_page() -> str:
| `model` | `str` | Combined identifier in `"provider:model"` format (e.g., `"openai:gpt-4.1-mini"`). The legacy `"provider/model"` format is also accepted but deprecated. |"""
)
parts.append("")
parts.append("**Returns:** A `(LLMProvider, model_name)` tuple.")
parts.append(
"**Returns:** A `(provider, model_name)` tuple. `provider` is an `LLMProvider` member "
"when the provider has one, and the bare name for a config-only registry gateway."
)
parts.append("")
parts.append("**Raises:** `ValueError` if the string does not contain a `:` or `/` delimiter.")
parts.append(
"**Raises:** `ValueError` if the string does not contain a `:` or `/` delimiter, or "
"`UnsupportedProviderError` if the provider is not resolvable."
)
parts.append("")
parts.append(
"""```python
Expand Down
74 changes: 65 additions & 9 deletions src/any_llm/any_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -250,7 +250,9 @@ def _create_provider(
if registry_class is not None:
return registry_class(api_key=api_key, api_base=api_base, **kwargs)

provider_key = LLMProvider.from_string(provider_key).value
# Resolve through the shared resolver so the unsupported-provider error lists
# registry gateways too, not just the enum.
provider_key = str(cls.resolve_provider_key(provider_key))

provider_class_name = f"{provider_key.capitalize()}Provider"
provider_module_name = f"{provider_key}"
Expand Down Expand Up @@ -296,7 +298,9 @@ def get_provider_class(cls, provider_key: str | LLMProvider) -> type[AnyLLM]:
if registry_class is not None:
return registry_class

provider_key = LLMProvider.from_string(provider_key).value
# Resolve through the shared resolver so the unsupported-provider error lists
# registry gateways too, not just the enum.
provider_key = str(cls.resolve_provider_key(provider_key))

provider_class_name = f"{provider_key.capitalize()}Provider"
provider_module_name = f"{provider_key}"
Expand All @@ -312,10 +316,54 @@ def get_provider_class(cls, provider_key: str | LLMProvider) -> type[AnyLLM]:
provider_class: type[AnyLLM] = getattr(module, provider_class_name)
return provider_class

@staticmethod
def get_registry_provider_names() -> list[str]:
"""Names of config-only gateways that exist only as registry rows.

Registry rows do not need an ``LLMProvider`` member, so these names are
absent from the enum and have to be added to any enumeration of
providers explicitly.
"""
from any_llm.providers.registry import PROVIDER_REGISTRY

enum_values = {provider.value for provider in LLMProvider}
return sorted(name for name in PROVIDER_REGISTRY if name not in enum_values)

@classmethod
def get_supported_providers(cls) -> list[str]:
"""Get a list of supported provider keys."""
return [provider.value for provider in LLMProvider]
"""Get a list of supported provider keys.

Includes registry-only gateways, which resolve by name without an
``LLMProvider`` member.
"""
return [provider.value for provider in LLMProvider] + cls.get_registry_provider_names()

@classmethod
def resolve_provider_key(cls, provider_key: str | LLMProvider) -> str | LLMProvider:
"""Resolve a provider key to an ``LLMProvider`` member where one exists.

Registry-only gateways have no enum member, so their name is returned
unchanged. Everything downstream (``create``, ``get_provider_class``)
accepts either form.

Raises:
UnsupportedProviderError: The key is neither an enum member nor a
registry row.

"""
if isinstance(provider_key, LLMProvider):
return provider_key
# Match LLMProvider.from_string's normalization so both resolution paths
# accept the same spellings.
normalized = provider_key.strip().lower()
try:
return LLMProvider(normalized)
except ValueError:
from any_llm.providers.registry import get_registry_config

if get_registry_config(normalized) is not None:
return normalized
raise UnsupportedProviderError(provider_key, cls.get_supported_providers()) from None

@classmethod
def get_all_provider_metadata(cls) -> list[ProviderMetadata]:
Expand All @@ -337,21 +385,29 @@ def get_all_provider_metadata(cls) -> list[ProviderMetadata]:

@classmethod
def get_provider_enum(cls, provider_key: str) -> LLMProvider:
"""Convert a string provider key to a ProviderName enum."""
"""Convert a string provider key to a ProviderName enum.

Registry-only gateways have no enum member, so this raises for them even
though they are resolvable. Use ``resolve_provider_key`` to accept both.
"""
try:
return LLMProvider(provider_key)
except ValueError as e:
supported = [provider.value for provider in LLMProvider]
raise UnsupportedProviderError(provider_key, supported) from e
# Report everything resolvable, not just the enum, so the message does
# not omit registry-only gateways.
raise UnsupportedProviderError(provider_key, cls.get_supported_providers()) from e

@classmethod
def split_model_provider(cls, model: str) -> tuple[LLMProvider, str]:
def split_model_provider(cls, model: str) -> tuple[str | LLMProvider, str]:
"""Extract the provider key from the model identifier.

Supports both new format 'provider:model' (e.g., 'mistral:mistral-small')
and legacy format 'provider/model' (e.g., 'mistral/mistral-small').

The legacy format will be deprecated in version 1.0.

Returns an ``LLMProvider`` member when the provider has one, and the bare
name for registry-only gateways. Both forms are accepted by ``create``.
"""
colon_index = model.find(":")
slash_index = model.find("/")
Expand All @@ -376,7 +432,7 @@ def split_model_provider(cls, model: str) -> tuple[LLMProvider, str]:
if not provider or not model_name:
msg = f"Invalid model format. Expected 'provider:model' or 'provider/model', got '{model}'"
raise ValueError(msg)
return cls.get_provider_enum(provider), model_name
return cls.resolve_provider_key(provider), model_name

@staticmethod
@abstractmethod
Expand Down
58 changes: 29 additions & 29 deletions src/any_llm/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,7 +95,7 @@ def completion(
if provider is None:
provider_key, model_id = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_id = model

llm = AnyLLM.create(
Expand Down Expand Up @@ -204,7 +204,7 @@ async def acompletion(
if provider is None:
provider_key, model_id = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_id = model

llm = AnyLLM.create(
Expand Down Expand Up @@ -345,7 +345,7 @@ def responses(
if provider is None:
provider_key, model_id = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_id = model

llm = AnyLLM.create(
Expand Down Expand Up @@ -494,7 +494,7 @@ async def aresponses(
if provider is None:
provider_key, model_id = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_id = model

llm = AnyLLM.create(
Expand Down Expand Up @@ -598,7 +598,7 @@ def messages(
if provider is None:
provider_key, model_id = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_id = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -682,7 +682,7 @@ async def amessages(
if provider is None:
provider_key, model_id = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_id = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -737,7 +737,7 @@ def embedding(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -775,7 +775,7 @@ async def aembedding(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -823,7 +823,7 @@ def image_generation(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -881,7 +881,7 @@ async def aimage_generation(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -942,7 +942,7 @@ def transcription(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -1000,7 +1000,7 @@ async def atranscription(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -1056,7 +1056,7 @@ def speech(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -1109,7 +1109,7 @@ async def aspeech(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand Down Expand Up @@ -1163,7 +1163,7 @@ def moderation(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand All @@ -1184,7 +1184,7 @@ async def amoderation(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

llm = AnyLLM.create(provider_key, api_key=api_key, api_base=api_base, **client_args or {})
Expand All @@ -1204,7 +1204,7 @@ def _resolve_rerank_target(
if provider is None:
provider_key, model_name = AnyLLM.split_model_provider(model)
else:
provider_key = LLMProvider.from_string(provider)
provider_key = AnyLLM.resolve_provider_key(provider)
model_name = model

if top_n is not None:
Expand Down Expand Up @@ -1287,7 +1287,7 @@ def list_models(
**kwargs: Any,
) -> Sequence[Model]:
"""List available models for a provider."""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return llm.list_models(**kwargs)


Expand All @@ -1299,7 +1299,7 @@ async def alist_models(
**kwargs: Any,
) -> Sequence[Model]:
"""List available models for a provider asynchronously."""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return await llm.alist_models(**kwargs)


Expand Down Expand Up @@ -1332,7 +1332,7 @@ def create_batch(
The created batch object

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return llm.create_batch(
input_file_path=input_file_path,
endpoint=endpoint,
Expand Down Expand Up @@ -1371,7 +1371,7 @@ async def acreate_batch(
The created batch object

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return await llm.acreate_batch(
input_file_path=input_file_path,
endpoint=endpoint,
Expand Down Expand Up @@ -1404,7 +1404,7 @@ def retrieve_batch(
The batch object

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return llm.retrieve_batch(batch_id, **kwargs)


Expand All @@ -1431,7 +1431,7 @@ async def aretrieve_batch(
The batch object

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return await llm.aretrieve_batch(batch_id, **kwargs)


Expand All @@ -1458,7 +1458,7 @@ def cancel_batch(
The cancelled batch object

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return llm.cancel_batch(batch_id, **kwargs)


Expand All @@ -1485,7 +1485,7 @@ async def acancel_batch(
The cancelled batch object

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return await llm.acancel_batch(batch_id, **kwargs)


Expand Down Expand Up @@ -1514,7 +1514,7 @@ def list_batches(
A list of batch objects

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return llm.list_batches(after=after, limit=limit, **kwargs)


Expand Down Expand Up @@ -1543,7 +1543,7 @@ async def alist_batches(
A list of batch objects

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return await llm.alist_batches(after=after, limit=limit, **kwargs)


Expand All @@ -1570,7 +1570,7 @@ def retrieve_batch_results(
The batch results containing per-request outcomes.

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return llm.retrieve_batch_results(batch_id, **kwargs)


Expand All @@ -1597,5 +1597,5 @@ async def aretrieve_batch_results(
The batch results containing per-request outcomes.

"""
llm = AnyLLM.create(LLMProvider.from_string(provider), api_key=api_key, api_base=api_base, **client_args or {})
llm = AnyLLM.create(AnyLLM.resolve_provider_key(provider), api_key=api_key, api_base=api_base, **client_args or {})
return await llm.aretrieve_batch_results(batch_id, **kwargs)
Loading