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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ docker = [
"any-llm-sdk[deepseek,openai,openrouter]>=1.19,<2",
]
dev = [
"any-llm-sdk[all]>=1.19,<2",
"any-llm-sdk[deepseek,openai,openrouter,otari]>=1.19,<2",
"anyio>=4.11,<5",
"pytest>=9,<10",
"pytest-cov>=7,<8",
Expand Down
1 change: 1 addition & 0 deletions tests/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
"""Test suite support package."""
15 changes: 15 additions & 0 deletions tests/any_llm_helpers.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
"""Shared AnyLLM test helpers."""

from any_llm import AnyLLM


def loadable_any_llm_providers() -> tuple[str, ...]:
"""Return providers whose SDK classes can be imported."""
result: list[str] = []
for provider in AnyLLM.get_supported_providers():
try:
AnyLLM.get_provider_class(provider)
result.append(provider)
except ImportError:
pass
return tuple(result)
22 changes: 11 additions & 11 deletions tests/test_any_llm_compatibility.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
from any_llm import AnyLLM
from any_llm.providers.openai.base import BaseOpenAIProvider

from tests.any_llm_helpers import loadable_any_llm_providers
from weather_briefing.data.any_llm_compatibility import (
UNSUPPORTED_DEFAULT_HEADER_PROVIDERS,
UNSUPPORTED_JSON_OBJECT_PROVIDERS,
Expand All @@ -22,20 +23,18 @@ def test_default_header_provider_compatibility_matches_the_pinned_sdk() -> None:
"watsonx",
"xai",
} == UNSUPPORTED_DEFAULT_HEADER_PROVIDERS
loadable = set(loadable_any_llm_providers())
checkable_blacklist = UNSUPPORTED_DEFAULT_HEADER_PROVIDERS & loadable
completion_providers = {
provider
for provider in AnyLLM.get_supported_providers()
if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION
provider for provider in loadable if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION
}

assert completion_providers > UNSUPPORTED_DEFAULT_HEADER_PROVIDERS
assert completion_providers > checkable_blacklist


def test_json_object_provider_compatibility_matches_the_pinned_sdk() -> None:
loadable = set(loadable_any_llm_providers())
completion_providers = {
provider
for provider in AnyLLM.get_supported_providers()
if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION
provider for provider in loadable if AnyLLM.get_provider_class(provider).SUPPORTS_COMPLETION
}
json_object_providers = {
provider
Expand All @@ -44,6 +43,7 @@ def test_json_object_provider_compatibility_matches_the_pinned_sdk() -> None:
}
unsupported_providers = completion_providers - json_object_providers

assert unsupported_providers == UNSUPPORTED_JSON_OBJECT_PROVIDERS
assert json_object_providers | UNSUPPORTED_JSON_OBJECT_PROVIDERS == completion_providers
assert json_object_providers.isdisjoint(UNSUPPORTED_JSON_OBJECT_PROVIDERS)
checkable_blacklist = UNSUPPORTED_JSON_OBJECT_PROVIDERS & loadable
assert unsupported_providers == checkable_blacklist
assert json_object_providers | checkable_blacklist == completion_providers
assert json_object_providers.isdisjoint(checkable_blacklist)
9 changes: 7 additions & 2 deletions tests/test_any_llm_provider.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from openai import AsyncOpenAI, BadRequestError
from pydantic import BaseModel, ValidationError

from tests.any_llm_helpers import loadable_any_llm_providers
from weather_briefing.api_client import LoggedAsyncClient
from weather_briefing.data.any_llm_compatibility import UNSUPPORTED_JSON_OBJECT_PROVIDERS
from weather_briefing.llm import (
Expand Down Expand Up @@ -411,11 +412,15 @@ async def test_provider_native_request_error_switches_to_fallback(monkeypatch) -
assert len(fallback_client.calls) == 1


def test_development_dependencies_load_runtime_providers() -> None:
assert {"deepseek", "openai", "openrouter"} <= set(loadable_any_llm_providers())


@pytest.mark.parametrize(
"provider",
tuple(AnyLLM.get_supported_providers()),
loadable_any_llm_providers(),
)
def test_factory_classifies_every_any_llm_provider(monkeypatch, provider: str) -> None:
def test_factory_classifies_every_loadable_provider(monkeypatch, provider: str) -> None:
created: list[tuple[str, dict[str, object]]] = []

def fake_create(provider: str, **options: object) -> _CompletionClientStub:
Expand Down
Loading