diff --git a/.github/workflows/ty.yml b/.github/workflows/ty.yml index 3c02d910c..ce438869c 100644 --- a/.github/workflows/ty.yml +++ b/.github/workflows/ty.yml @@ -27,6 +27,8 @@ jobs: run: uv run ty check --error-on-warning --output-format github instructor/ - name: Run type check with ty (tests) run: uv run ty check --config-file ty-tests.toml --error-on-warning --output-format github tests + - name: Run public surface inference check with Pyright + run: uv run --with 'pyright==1.1.408' pyright tests/typing/test_public_surface.py - name: Check installed-package public typing run: tests/typing/check_installed_package.sh - name: Check supported Python versions and platforms diff --git a/tests/llm/test_openai/test_multimodal.py b/tests/llm/test_openai/test_multimodal.py index 905ef6d8c..67b45ace9 100644 --- a/tests/llm/test_openai/test_multimodal.py +++ b/tests/llm/test_openai/test_multimodal.py @@ -8,6 +8,8 @@ import base64 import os +from instructor.core.exceptions import InstructorRetryException + audio_url = "https://raw.githubusercontent.com/instructor-ai/instructor/main/tests/assets/gettysburg.wav" image_url = "https://raw.githubusercontent.com/instructor-ai/instructor/main/tests/assets/image.jpg" @@ -60,21 +62,32 @@ def test_multimodal_audio_description(audio_file, mode, client): class AudioDescription(BaseModel): source: str - response = client.chat.completions.create( - model="gpt-audio-1.5", - response_model=AudioDescription, - modalities=["text"], - messages=[ - { - "role": "user", - "content": [ - "Where's this excerpt from?", - audio_file, - ], # type: ignore - }, - ], - audio={"voice": "alloy", "format": "wav"}, # type: ignore - ) + try: + response = client.chat.completions.create( + model="gpt-audio-1.5", + response_model=AudioDescription, + modalities=["text"], + messages=[ + { + "role": "user", + "content": [ + "Where's this excerpt from?", + audio_file, + ], # type: ignore + }, + ], + audio={"voice": "alloy", "format": "wav"}, # type: ignore + ) + except InstructorRetryException as exc: + message = str(exc).lower() + if ( + "model_not_found" in message + or "does not exist or you do not have access" in message + ): + pytest.skip(f"Audio model unavailable in this environment: {exc}") + raise + + assert isinstance(response, AudioDescription) class ImageDescription(BaseModel): diff --git a/tests/test_lazy_imports.py b/tests/test_lazy_imports.py index 485982d67..521d043fb 100644 --- a/tests/test_lazy_imports.py +++ b/tests/test_lazy_imports.py @@ -1,9 +1,13 @@ import importlib import json +import os import subprocess import sys +from pathlib import Path from typing import TypedDict +_PROJECT_ROOT = Path(__file__).resolve().parents[1] + class ColdImportState(TypedDict): modules: list[str] @@ -21,11 +25,17 @@ def _cold_import_state(code: str) -> ColdImportState: ) print(json.dumps(mods)) """ + env = dict(os.environ) + env["PYTHONPATH"] = os.pathsep.join( + path for path in (str(_PROJECT_ROOT), env.get("PYTHONPATH")) if path + ) result = subprocess.run( [sys.executable, "-c", probe, code], check=True, capture_output=True, text=True, + cwd=_PROJECT_ROOT, + env=env, ) modules = json.loads(result.stdout) return {"modules": modules, "count": len(modules)} diff --git a/tests/v2/provider_matrix.py b/tests/v2/provider_matrix.py index 586d66949..e642c2a15 100644 --- a/tests/v2/provider_matrix.py +++ b/tests/v2/provider_matrix.py @@ -4,7 +4,6 @@ import importlib.util from pathlib import Path -from typing import Any import pytest @@ -23,25 +22,31 @@ PROVIDER_HANDLER_MODES = { provider: spec.supported_modes for provider, spec in PROVIDER_SPECS.items() } - - -def legacy_config_dicts() -> dict[Provider, dict[str, Any]]: - """Expose the old dict shape while the baseline tests migrate.""" - return { - provider: { - "provider_string": spec.provider_string, - "supported_modes": list(spec.supported_modes), - "unsupported_modes": list(spec.unsupported_modes), - "legacy_modes": spec.legacy_modes, - "from_function": spec.from_function, - "sdk_module": spec.sdk_module, - "basic_modes": list(spec.basic_modes), - "async_modes": list(spec.async_modes), - "missing_sdk_message": spec.missing_sdk_message, - } - for provider, spec in TEST_PROVIDER_SPECS.items() - } - +PARTIAL_STREAM_CASES = tuple( + (provider, mode) + for provider, spec in TEST_PROVIDER_SPECS.items() + for mode in spec.capabilities.partial_stream_modes +) +ITERABLE_STREAM_CASES = tuple( + (provider, mode) + for provider, spec in TEST_PROVIDER_SPECS.items() + for mode in spec.capabilities.iterable_stream_modes +) +TYPED_MULTIMODAL_PROVIDERS = tuple( + provider + for provider, spec in TEST_PROVIDER_SPECS.items() + if spec.capabilities.multimodal_inputs +) +TYPED_MULTIMODAL_CASES = tuple( + (provider, media_type) + for provider, spec in TEST_PROVIDER_SPECS.items() + for media_type in spec.capabilities.multimodal_inputs +) +EXPLICIT_PARALLEL_PROVIDERS = tuple( + provider + for provider, spec in TEST_PROVIDER_SPECS.items() + if spec.capabilities.explicit_parallel_tools +) _PROJECT_ROOT = Path(__file__).resolve().parents[2] _HANDLERS_LOADED: set[Provider] = set() diff --git a/tests/v2/test_anthropic_parallel_runtime.py b/tests/v2/test_anthropic_parallel_runtime.py new file mode 100644 index 000000000..305b4dada --- /dev/null +++ b/tests/v2/test_anthropic_parallel_runtime.py @@ -0,0 +1,33 @@ +from __future__ import annotations + +from collections.abc import Iterable + +import pytest +from pydantic import BaseModel + +from instructor.v2.providers.anthropic import parallel + + +class Weather(BaseModel): + city: str + + +class Score(BaseModel): + value: int + + +def test_parallel_schema_generation_is_owned_by_anthropic( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[type[BaseModel]] = [] + + def fake_schema(model: type[BaseModel]) -> dict[str, str]: + calls.append(model) + return {"name": model.__name__} + + monkeypatch.setattr(parallel, "generate_anthropic_schema", fake_schema) + + schemas = parallel.handle_parallel_model(Iterable[Weather | Score]) + + assert schemas == [{"name": "Weather"}, {"name": "Score"}] + assert calls == [Weather, Score] diff --git a/tests/v2/test_bedrock_client.py b/tests/v2/test_bedrock_client.py deleted file mode 100644 index 40777d018..000000000 --- a/tests/v2/test_bedrock_client.py +++ /dev/null @@ -1,67 +0,0 @@ -"""Provider-specific tests for Bedrock v2 client factory.""" - -from __future__ import annotations - -import pytest - -from instructor import Mode - - -class TestBedrockClientWithSDK: - """Tests for Bedrock client factory that require botocore.""" - - @pytest.fixture - def bedrock_available(self): - """Check if botocore is available.""" - try: - from botocore.client import BaseClient # noqa: F401 - - return True - except ImportError: - return False - - def test_from_bedrock_raises_without_sdk(self, bedrock_available): - """from_bedrock should raise when botocore is missing.""" - if bedrock_available: - pytest.skip( - "botocore is installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.bedrock.client import from_bedrock - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="botocore is not installed"): - from_bedrock(None) # ty: ignore[no-matching-overload] - - def test_from_bedrock_with_invalid_client(self, bedrock_available): - """from_bedrock should reject non-BaseClient objects.""" - if not bedrock_available: - pytest.skip( - "botocore not installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.bedrock.client import from_bedrock - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="BaseClient"): - from_bedrock("not a client") # ty: ignore[no-matching-overload] - - def test_from_bedrock_with_invalid_mode(self, bedrock_available): - """from_bedrock should raise for unsupported modes.""" - if not bedrock_available: - pytest.skip( - "botocore not installed" # ty: ignore[too-many-positional-arguments] - ) - - from botocore.client import BaseClient - from instructor.v2.providers.bedrock.client import from_bedrock - from instructor.core.exceptions import ModeError - - def _converse(**_kwargs): - return {} - - client = BaseClient.__new__(BaseClient) - client.converse = _converse # type: ignore[assignment] - - with pytest.raises(ModeError): - from_bedrock(client, mode=Mode.JSON_SCHEMA) diff --git a/tests/v2/test_cerebras_client.py b/tests/v2/test_cerebras_client.py deleted file mode 100644 index e388d4bcb..000000000 --- a/tests/v2/test_cerebras_client.py +++ /dev/null @@ -1,72 +0,0 @@ -"""Unit tests for Cerebras v2 client factory. - -These tests verify client factory behavior without requiring API keys. -""" - -from __future__ import annotations - -import pytest - -from instructor import Mode - - -# ============================================================================ -# Integration Tests (require Cerebras SDK but not API key) -# ============================================================================ - - -class TestCerebrasClientWithSDK: - """Tests that require Cerebras SDK but not API key.""" - - @pytest.fixture - def cerebras_available(self): - """Check if cerebras SDK is available.""" - try: - from cerebras.cloud.sdk import Cerebras # noqa: F401 - - return True - except ImportError: - return False - - def test_from_cerebras_raises_without_sdk(self, cerebras_available): - """Test from_cerebras raises error when cerebras not installed.""" - if cerebras_available: - pytest.skip( - "cerebras is installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.cerebras.client import from_cerebras - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="cerebras is not installed"): - from_cerebras("not a client") # ty: ignore[no-matching-overload] - - def test_from_cerebras_with_invalid_client(self, cerebras_available): - """Test from_cerebras raises error with invalid client.""" - if not cerebras_available: - pytest.skip( - "cerebras not installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.cerebras.client import from_cerebras - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="must be an instance"): - from_cerebras("not a client") # ty: ignore[no-matching-overload] - - def test_from_cerebras_with_invalid_mode(self, cerebras_available): - """Test from_cerebras raises error with invalid mode.""" - if not cerebras_available: - pytest.skip( - "cerebras not installed" # ty: ignore[too-many-positional-arguments] - ) - - from cerebras.cloud.sdk import Cerebras - - from instructor.v2.providers.cerebras.client import from_cerebras - from instructor.core.exceptions import ModeError - - client = Cerebras(api_key="fake-key") - - with pytest.raises(ModeError): - from_cerebras(client, mode=Mode.RESPONSES_TOOLS) diff --git a/tests/v2/test_client_unified.py b/tests/v2/test_client_unified.py index 43da4bf6b..c6de6b8e1 100644 --- a/tests/v2/test_client_unified.py +++ b/tests/v2/test_client_unified.py @@ -1,528 +1,260 @@ -"""Unified parametrized tests for all provider client factories. - -These tests verify client factory behavior (mode normalization, registry, errors, imports) -across all providers without requiring API keys. -""" +"""Shared contracts for declarative provider client factories.""" from __future__ import annotations import importlib.util -from pathlib import Path -from typing import Any +import inspect +from types import SimpleNamespace +from typing import Any, cast import pytest from instructor import Mode, Provider -from instructor.v2.core.registry import mode_registry, normalize_mode -from tests.v2.provider_matrix import legacy_config_dicts - -_PROJECT_ROOT = Path(__file__).resolve().parents[2] -_HANDLER_MODULE_PATHS: dict[Provider, Path] = { - Provider.OPENAI: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.ANYSCALE: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.TOGETHER: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.DATABRICKS: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.DEEPSEEK: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.ANTHROPIC: _PROJECT_ROOT / "instructor/v2/providers/anthropic/handlers.py", - Provider.GENAI: _PROJECT_ROOT / "instructor/v2/providers/genai/handlers.py", - Provider.GEMINI: _PROJECT_ROOT / "instructor/v2/providers/gemini/handlers.py", - Provider.COHERE: _PROJECT_ROOT / "instructor/v2/providers/cohere/handlers.py", - Provider.OPENROUTER: _PROJECT_ROOT - / "instructor/v2/providers/openrouter/handlers.py", - Provider.PERPLEXITY: _PROJECT_ROOT - / "instructor/v2/providers/perplexity/handlers.py", - Provider.XAI: _PROJECT_ROOT / "instructor/v2/providers/xai/handlers.py", - Provider.GROQ: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.MISTRAL: _PROJECT_ROOT / "instructor/v2/providers/mistral/handlers.py", - Provider.FIREWORKS: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.CEREBRAS: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.WRITER: _PROJECT_ROOT / "instructor/v2/providers/writer/handlers.py", - Provider.BEDROCK: _PROJECT_ROOT / "instructor/v2/providers/bedrock/handlers.py", - Provider.VERTEXAI: _PROJECT_ROOT / "instructor/v2/providers/vertexai/handlers.py", -} -_HANDLERS_LOADED: set[Provider] = set() - - -def _clear_proxy_env(monkeypatch: pytest.MonkeyPatch) -> None: - for key in ( - "ALL_PROXY", - "all_proxy", - "HTTPS_PROXY", - "https_proxy", - "HTTP_PROXY", - "http_proxy", - ): - monkeypatch.delenv(key, raising=False) - - -def _ensure_handlers_loaded(provider: Provider) -> None: - if provider in _HANDLERS_LOADED: - return - provider_modes = PROVIDER_CLIENT_CONFIGS.get(provider, {}).get( - "supported_modes", [] - ) - if provider_modes and all( - mode_registry.is_registered(provider, mode) for mode in provider_modes - ): - _HANDLERS_LOADED.add(provider) - return - handler_path = _HANDLER_MODULE_PATHS.get(provider) - if handler_path is None or not handler_path.exists(): - return - spec = importlib.util.spec_from_file_location( - f"tests.v2.handlers_{provider.value}", - handler_path, - ) - if spec is None or spec.loader is None: - return - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - _HANDLERS_LOADED.add(provider) - - -PROVIDER_CLIENT_CONFIGS: dict[Provider, dict[str, Any]] = legacy_config_dicts() +from instructor.v2.core import client_factory +from instructor.v2.core.client import AsyncInstructor, Instructor +from instructor.v2.core.errors import ClientError, ModeError +from instructor.v2.core.provider_specs import PROVIDER_SPECS +from instructor.v2.core.registry import mode_registry +from instructor.v2.providers.anthropic import client as anthropic_client +from instructor.v2.providers.cohere import client as cohere_client +from instructor.v2.providers.gemini import client as gemini_client +from instructor.v2.providers.vertexai import client as vertexai_client +from tests.v2.provider_matrix import TEST_PROVIDER_SPECS, ensure_handlers_loaded def _dependency_missing(module: str) -> bool: - """Check if a dependency module is missing.""" try: - return importlib.util.find_spec(module.split(".")[0]) is None + return importlib.util.find_spec(module) is None except ModuleNotFoundError: return True -def _is_expected_missing_dependency(provider: Provider, exc: ImportError) -> bool: - """Return True when an import failed because the provider SDK is unavailable.""" - sdk_module = PROVIDER_CLIENT_CONFIGS[provider]["sdk_module"] - expected_root = str(sdk_module).split(".")[0] - missing_name = getattr(exc, "name", None) - if missing_name: - return missing_name.split(".")[0] == expected_root - - return f"No module named '{expected_root}'" in str(exc) - - -def _get_provider_params(): - """Generate provider parameters for parametrized tests.""" - return [ - pytest.param(provider, id=provider.value) - for provider in PROVIDER_CLIENT_CONFIGS.keys() - ] - - -def _get_provider_mode_params(): - """Generate (provider, mode) parameters for supported modes.""" - params = [] - for provider, config in PROVIDER_CLIENT_CONFIGS.items(): - for mode in config["supported_modes"]: - params.append( - pytest.param(provider, mode, id=f"{provider.value}-{mode.value}") - ) - return params - - -def _get_provider_unsupported_mode_params(): - """Generate (provider, mode) parameters for unsupported modes.""" - params = [] - for provider, config in PROVIDER_CLIENT_CONFIGS.items(): - for mode in config["unsupported_modes"]: - params.append( - pytest.param(provider, mode, id=f"{provider.value}-{mode.value}") - ) - return params - - -def _get_provider_legacy_mode_params(): - """Generate (provider, legacy_mode) parameters.""" - params = [] - for provider, config in PROVIDER_CLIENT_CONFIGS.items(): - for legacy_mode in config["legacy_modes"].keys(): - params.append( - pytest.param( - provider, - legacy_mode, - id=f"{provider.value}-{legacy_mode.value}", +def test_manifest_exposes_complete_declarative_client_contracts() -> None: + public = __import__("instructor.v2", fromlist=["__all__"]) + custom_factories = {Provider.LITELLM, Provider.XAI} + for provider, spec in TEST_PROVIDER_SPECS.items(): + ensure_handlers_loaded(provider) + assert callable(getattr(public, spec.from_function or "")), provider + for mode in spec.supported_modes: + handlers = mode_registry.get_handlers(provider, mode) + assert all( + ( + handlers.request_handler, + handlers.reask_handler, + handlers.response_parser, ) + ), (provider, mode) + if provider not in custom_factories: + assert spec.client and spec.client.sync_types and spec.client.create, ( + provider ) - return params - - -# ============================================================================ -# Mode Registry Tests -# ============================================================================ - - -@pytest.mark.parametrize("provider,mode", _get_provider_mode_params()) -def test_supported_mode_is_registered(provider: Provider, mode: Mode) -> None: - """Test that all supported modes are registered in the registry.""" - _ensure_handlers_loaded(provider) - assert mode_registry.is_registered(provider, mode), ( - f"Mode {mode.value} should be registered for {provider.value}" - ) - - -@pytest.mark.parametrize("provider,mode", _get_provider_unsupported_mode_params()) -def test_unsupported_mode_not_registered(provider: Provider, mode: Mode) -> None: - """Test that unsupported modes are NOT registered.""" - assert not mode_registry.is_registered(provider, mode), ( - f"Mode {mode.value} should NOT be registered for {provider.value}" - ) - -@pytest.mark.parametrize("provider", _get_provider_params()) -def test_get_modes_for_provider(provider: Provider) -> None: - """Test getting all modes for a provider.""" - _ensure_handlers_loaded(provider) - config = PROVIDER_CLIENT_CONFIGS[provider] - registered_modes = mode_registry.get_modes_for_provider(provider) - # All supported modes should be registered - for mode in config["supported_modes"]: - assert mode in registered_modes, ( - f"Mode {mode.value} should be in registered modes for {provider.value}" +def test_missing_optional_sdks_raise_shared_client_error() -> None: + for spec in TEST_PROVIDER_SPECS.values(): + if spec.sdk_module is None or not _dependency_missing(spec.sdk_module): + continue + module = __import__( + spec.client_module or "", fromlist=[spec.from_function or ""] ) + with pytest.raises(ClientError, match=spec.missing_sdk_message): + getattr(module, spec.from_function or "")("not a client") - # Unsupported modes should not be registered - for mode in config["unsupported_modes"]: - assert mode not in registered_modes, ( - f"Mode {mode.value} should NOT be in registered modes for {provider.value}" - ) +class _SyncClient: + def __init__(self) -> None: + self.calls: list[dict[str, Any]] = [] + create = lambda **kwargs: self.calls.append(kwargs) or object() + self.chat = SimpleNamespace(completions=SimpleNamespace(create=create)) + self.custom_create = create -@pytest.mark.parametrize("provider,mode", _get_provider_mode_params()) -def test_handlers_have_all_methods(provider: Provider, mode: Mode) -> None: - """Test that all handlers have required methods.""" - _ensure_handlers_loaded(provider) - handlers = mode_registry.get_handlers(provider, mode) - assert handlers.request_handler is not None - assert handlers.reask_handler is not None - assert handlers.response_parser is not None +class _AsyncClient(_SyncClient): + pass -# ============================================================================ -# Mode Normalization Tests -# ============================================================================ +def _fake_types(paths: tuple[str, ...], _message: str) -> tuple[type[Any], ...]: + return (_AsyncClient,) if any("Async" in path for path in paths) else (_SyncClient,) -@pytest.mark.parametrize("provider,mode", _get_provider_mode_params()) -def test_generic_mode_passes_through(provider: Provider, mode: Mode) -> None: - """Test that generic modes pass through unchanged.""" - result = normalize_mode(provider, mode) - assert result == mode, ( - f"Generic mode {mode.value} should pass through unchanged for {provider.value}" - ) - - -@pytest.mark.parametrize("provider,legacy_mode", _get_provider_legacy_mode_params()) -def test_legacy_mode_normalizes_to_registered_mode( - provider: Provider, legacy_mode: Mode +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("native", "wrapper_type"), + [(_SyncClient(), Instructor), (_AsyncClient(), AsyncInstructor)], +) +async def test_shared_factory_selects_wrapper_and_preserves_native_stream_flag( + monkeypatch: pytest.MonkeyPatch, + native: _SyncClient, + wrapper_type: type[Instructor] | type[AsyncInstructor], ) -> None: - """Legacy provider-specific modes normalize to registered v2 modes.""" - result = normalize_mode(provider, legacy_mode) - assert result != legacy_mode - assert mode_registry.is_registered(provider, legacy_mode), ( - f"Legacy mode {legacy_mode.value} should remain accepted for {provider.value}" + monkeypatch.setattr(client_factory, "_resolve_types", _fake_types) + monkeypatch.setattr(client_factory, "patch_v2", lambda *, func, **_kwargs: func) + wrapped = client_factory.create_instructor( + native, provider=Provider.GROQ, mode=Mode.TOOLS ) + result = wrapped.create_fn(stream=True) + if inspect.isawaitable(result): + await result + assert isinstance(wrapped, wrapper_type) + assert native.calls == [{"stream": True}] + + +def _set_path(root: Any, path: str, value: Any) -> None: + parts = path.split(".") + for part in parts[:-1]: + child = getattr(root, part, None) + if child is None: + child = SimpleNamespace() + setattr(root, part, child) + root = child + setattr(root, parts[-1], value) + + +_STREAM_CASES = tuple( + (provider, is_async) + for provider, spec in PROVIDER_SPECS.items() + if spec.client is not None + for is_async, path in ( + (False, spec.client.stream), + (True, spec.client.async_stream), + ) + if path is not None +) -# ============================================================================ -# Import Tests -# ============================================================================ - - -@pytest.mark.parametrize("provider", _get_provider_params()) -def test_from_function_importable(provider: Provider) -> None: - """Test that from_* function is importable from instructor.v2.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - from_function = config["from_function"] +@pytest.mark.asyncio +@pytest.mark.parametrize(("provider", "is_async"), _STREAM_CASES) +async def test_declared_stream_paths_resolve_and_switch( + monkeypatch: pytest.MonkeyPatch, provider: Provider, is_async: bool +) -> None: + contract = PROVIDER_SPECS[provider].client + assert contract is not None + calls: list[dict[str, Any]] = [] - # Import from instructor.v2 - module = __import__("instructor.v2", fromlist=[from_function]) - func = getattr(module, from_function, None) + async def async_stream(**kwargs: Any) -> object: + calls.append(kwargs) + return object() - # Should be None if SDK not installed, or a callable if installed - assert func is None or callable(func), ( - f"{from_function} should be None or callable, got {type(func)}" + native = SimpleNamespace() + create_path = ( + (contract.async_create or contract.create) if is_async else contract.create ) - - -@pytest.mark.parametrize("provider", _get_provider_params()) -def test_handlers_importable(provider: Provider) -> None: - """Test that handlers are importable.""" - handler_path = _HANDLER_MODULE_PATHS.get(provider) - assert handler_path is not None and handler_path.exists(), ( - f"Missing handler module path for {provider.value}" + stream_path = contract.async_stream if is_async else contract.stream + _set_path( + native, create_path, async_stream if is_async else lambda **_kwargs: object() ) - - _ensure_handlers_loaded(provider) - - assert any( - mode_registry.is_registered(provider, mode) - for mode in PROVIDER_CLIENT_CONFIGS[provider]["supported_modes"] - ), f"No registered handlers found for {provider.value}" - - -# ============================================================================ -# Error Handling Tests -# ============================================================================ - - -@pytest.mark.parametrize("provider,mode", _get_provider_unsupported_mode_params()) -def test_unsupported_mode_raises_error(provider: Provider, mode: Mode) -> None: - """Test that getting handlers for unsupported mode raises KeyError.""" - with pytest.raises(KeyError): - mode_registry.get_handlers(provider, mode) - - -@pytest.mark.parametrize("provider", _get_provider_params()) -def test_parallel_tools_not_supported_unless_registered(provider: Provider) -> None: - """Test that PARALLEL_TOOLS is not supported unless registered.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - is_supported = Mode.PARALLEL_TOOLS in config["supported_modes"] - is_registered = mode_registry.is_registered(provider, Mode.PARALLEL_TOOLS) - - assert is_supported == is_registered, ( - f"PARALLEL_TOOLS support mismatch for {provider.value}: " - f"supported={is_supported}, registered={is_registered}" + _set_path( + native, + stream_path or "", + async_stream if is_async else lambda **kwargs: calls.append(kwargs), ) - - -@pytest.mark.parametrize("provider", _get_provider_params()) -def test_responses_tools_not_supported_unless_registered(provider: Provider) -> None: - """Test that RESPONSES_TOOLS is not supported unless registered.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - is_supported = Mode.RESPONSES_TOOLS in config["supported_modes"] - is_registered = mode_registry.is_registered(provider, Mode.RESPONSES_TOOLS) - - assert is_supported == is_registered, ( - f"RESPONSES_TOOLS support mismatch for {provider.value}: " - f"supported={is_supported}, registered={is_registered}" + monkeypatch.setattr(client_factory, "patch_v2", lambda *, func, **_kwargs: func) + wrapped = client_factory.create_instructor( + native, + provider=provider, + mode=PROVIDER_SPECS[provider].supported_modes[0], + model="default-model", + use_async=is_async, + sync_types=(SimpleNamespace,), + async_types=(SimpleNamespace,), ) + result = wrapped.create_fn(stream=True, model=None, value=1) + if inspect.isawaitable(result): + await result + expected = { + "model": "default-model" if contract.falsey_model_fallback else None, + "value": 1, + } + assert calls == [expected] -# ============================================================================ -# SDK Availability Tests -# ============================================================================ - - -@pytest.mark.parametrize("provider", _get_provider_params()) -def test_from_function_raises_without_sdk(provider: Provider) -> None: - """Test that from_* function raises error when SDK not installed.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - sdk_module = config["sdk_module"] - from_function = config["from_function"] - - if not _dependency_missing(sdk_module): - pytest.skip( - f"{sdk_module} is installed" # ty: ignore[too-many-positional-arguments] - ) - - # Try to import the from_* function from the provider's client module - try: - client_module_path = f"instructor.v2.providers.{provider.value}.client" - client_module = __import__(client_module_path, fromlist=[from_function]) - from_function_obj = getattr(client_module, from_function, None) - - if from_function_obj is None: - pytest.skip( - f"{from_function} not found in client module" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.core.exceptions import ClientError - - expected_message = config.get( - "missing_sdk_message", - f"{sdk_module.split('.')[0]} is not installed", - ) - with pytest.raises(ClientError, match=expected_message): - from_function_obj("not a client") # type: ignore[call-arg] - except (ImportError, ModuleNotFoundError) as exc: - if _is_expected_missing_dependency(provider, exc): - pytest.skip( - f"{sdk_module} import path is unavailable in this environment" # ty: ignore[too-many-positional-arguments] - ) - raise - - -# ============================================================================ -# String-Based Initialization Tests -# ============================================================================ +def _missing_types(*_args: Any) -> tuple[type[Any], ...]: + raise ClientError("missing dependency") -# OpenAI-compatible providers that support string-based initialization -_OPENAI_COMPAT_PROVIDERS = [ - Provider.ANYSCALE, - Provider.TOGETHER, - Provider.DATABRICKS, - Provider.DEEPSEEK, -] +def _unexpected_types(*_args: Any) -> tuple[type[Any], ...]: + raise AssertionError("mode validation must run first") @pytest.mark.parametrize( - "provider", - [pytest.param(p, id=p.value) for p in _OPENAI_COMPAT_PROVIDERS], + ("provider", "resolver", "error"), + [ + (Provider.GROQ, _missing_types, ClientError), + (Provider.GROQ, _fake_types, ModeError), + (Provider.GEMINI, _unexpected_types, ModeError), + (Provider.GENAI, _fake_types, ClientError), + ], ) -def test_string_based_initialization_delegates_to_from_provider( - provider: Provider, +def test_shared_factory_preserves_validation_precedence( + monkeypatch: pytest.MonkeyPatch, provider: Provider, resolver: Any, error: Any ) -> None: - """Test that string-based initialization delegates to from_provider.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - from_function = config["from_function"] - - # Import the from_* function - module = __import__("instructor.v2", fromlist=[from_function]) - func = getattr(module, from_function, None) - - if func is None: - pytest.skip( - f"{from_function} not available (SDK may not be installed)" # ty: ignore[too-many-positional-arguments] + monkeypatch.setattr(client_factory, "_resolve_types", resolver) + with pytest.raises(error): + client_factory.create_instructor( + object(), provider=provider, mode=Mode.RESPONSES_TOOLS ) - assert callable(func) - - # Mock from_provider to verify it's called - from unittest.mock import patch - - with patch("instructor.from_provider") as mock_from_provider: - # Call with string (model name) - func("test-model", mode=Mode.TOOLS) - - # Verify from_provider was called with correct provider prefix - mock_from_provider.assert_called_once() - call_args = mock_from_provider.call_args - assert call_args[0][0] == f"{provider.value}/test-model" - assert call_args[1]["mode"] == Mode.TOOLS @pytest.mark.parametrize( - "provider", - [pytest.param(p, id=p.value) for p in _OPENAI_COMPAT_PROVIDERS], + ("module", "sdk_attribute", "factory"), + [ + (gemini_client, "genai", gemini_client.from_gemini), + (vertexai_client, "gm", vertexai_client.from_vertexai), + ], ) -def test_string_based_initialization_with_async_client(provider: Provider) -> None: - """Test that string-based initialization supports async_client parameter.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - from_function = config["from_function"] - - # Import the from_* function - module = __import__("instructor.v2", fromlist=[from_function]) - func = getattr(module, from_function, None) - - if func is None: - pytest.skip( - f"{from_function} not available (SDK may not be installed)" # ty: ignore[too-many-positional-arguments] - ) - assert callable(func) - - # Mock from_provider to verify it's called - from unittest.mock import patch - - with patch("instructor.from_provider") as mock_from_provider: - # Call with string and async_client=True - func("test-model", mode=Mode.TOOLS, async_client=True) - - # Verify from_provider was called with async_client=True - mock_from_provider.assert_called_once() - call_args = mock_from_provider.call_args - assert call_args[0][0] == f"{provider.value}/test-model" - assert call_args[1]["mode"] == Mode.TOOLS - assert call_args[1]["async_client"] is True - - -@pytest.mark.parametrize( - "provider", - [pytest.param(p, id=p.value) for p in _OPENAI_COMPAT_PROVIDERS], -) -def test_string_based_initialization_forwards_kwargs(provider: Provider) -> None: - """Test that string-based initialization forwards all kwargs to from_provider.""" - config = PROVIDER_CLIENT_CONFIGS[provider] - from_function = config["from_function"] - - # Import the from_* function - module = __import__("instructor.v2", fromlist=[from_function]) - func = getattr(module, from_function, None) - - if func is None: - pytest.skip( - f"{from_function} not available (SDK may not be installed)" # ty: ignore[too-many-positional-arguments] - ) - assert callable(func) - - # Mock from_provider to verify it's called - from unittest.mock import patch - - with patch("instructor.from_provider") as mock_from_provider: - # Call with string and additional kwargs - func( - "test-model", - mode=Mode.TOOLS, - api_key="test-key", - base_url="https://test.example.com", - timeout=30, - ) - - # Verify from_provider was called with all kwargs - mock_from_provider.assert_called_once() - call_args = mock_from_provider.call_args - assert call_args[0][0] == f"{provider.value}/test-model" - assert call_args[1]["mode"] == Mode.TOOLS - assert call_args[1]["api_key"] == "test-key" - assert call_args[1]["base_url"] == "https://test.example.com" - assert call_args[1]["timeout"] == 30 - - -@pytest.mark.parametrize( - "provider", - [pytest.param(p, id=p.value) for p in _OPENAI_COMPAT_PROVIDERS], -) -def test_client_based_initialization_still_works( - provider: Provider, monkeypatch: pytest.MonkeyPatch +def test_mode_first_adapters_validate_before_missing_sdk( + monkeypatch: pytest.MonkeyPatch, module: Any, sdk_attribute: str, factory: Any ) -> None: - """Test that client-based initialization still works (backward compatibility).""" - from unittest.mock import patch + monkeypatch.setattr(module, sdk_attribute, None) + with pytest.raises(ModeError): + factory(object(), mode=Mode.RESPONSES_TOOLS) - config = PROVIDER_CLIENT_CONFIGS[provider] - from_function = config["from_function"] - sdk_module = config["sdk_module"] - # Skip if SDK not installed - if _dependency_missing(sdk_module): - pytest.skip( - f"{sdk_module} not installed" # ty: ignore[too-many-positional-arguments] - ) - - # Import the from_* function - module = __import__("instructor.v2", fromlist=[from_function]) - func = getattr(module, from_function, None) - - if func is None: - pytest.skip( - f"{from_function} not available" # ty: ignore[too-many-positional-arguments] - ) - assert callable(func) - - # Import OpenAI client - try: - import openai - except ImportError: - pytest.skip( - "openai package not installed" # ty: ignore[too-many-positional-arguments] - ) +def _capture_factory(monkeypatch: pytest.MonkeyPatch, module: Any) -> dict[str, Any]: + captured: dict[str, Any] = {} + monkeypatch.setattr( + module, + "create_instructor", + lambda _client, **kwargs: captured.update(kwargs) or object(), + ) + return captured - _clear_proxy_env(monkeypatch) - # Create a mock OpenAI client - client = openai.OpenAI(api_key="test-key") +def test_anthropic_beta_declares_beta_method_override( + monkeypatch: pytest.MonkeyPatch, +) -> None: + sync_type, async_type = type("Sync", (), {}), type("Async", (), {}) + monkeypatch.setattr( + anthropic_client, + "anthropic", + SimpleNamespace( + Anthropic=sync_type, + AnthropicBedrock=sync_type, + AnthropicVertex=sync_type, + AsyncAnthropic=async_type, + AsyncAnthropicBedrock=async_type, + AsyncAnthropicVertex=async_type, + ), + ) + captured = _capture_factory(monkeypatch, anthropic_client) + cast(Any, anthropic_client.from_anthropic)(sync_type(), beta=True) + assert captured["create_path"] == "beta.messages.create" + assert captured["async_create_path"] == "beta.messages.create" - # Call with client (should use _from_openai_compat, not from_provider) - with patch( - "instructor.v2.providers.openai.client._from_openai_compat" - ) as mock_compat: - mock_compat.return_value = "mock_instructor" - result = func(client, mode=Mode.TOOLS) - # Verify _from_openai_compat was called (not from_provider) - mock_compat.assert_called_once() - call_args = mock_compat.call_args - assert call_args[0][0] == client - assert call_args[1]["provider"] == provider - assert call_args[1]["mode"] == Mode.TOOLS +def test_cohere_v2_declares_version_and_client_families( + monkeypatch: pytest.MonkeyPatch, +) -> None: + v1, v2, async_v1, async_v2 = (type(name, (), {}) for name in "V1 V2 A1 A2".split()) + monkeypatch.setattr( + cohere_client, + "cohere", + SimpleNamespace( + Client=v1, ClientV2=v2, AsyncClient=async_v1, AsyncClientV2=async_v2 + ), + ) + captured = _capture_factory(monkeypatch, cohere_client) + cast(Any, cohere_client.from_cohere)(v2()) + assert captured["_cohere_client_version"] == "v2" + assert captured["sync_types"] == (v1, v2) + assert captured["async_types"] == (async_v1, async_v2) diff --git a/tests/v2/test_core_multimodal_runtime.py b/tests/v2/test_core_multimodal_runtime.py index a7661fa73..291efaf58 100644 --- a/tests/v2/test_core_multimodal_runtime.py +++ b/tests/v2/test_core_multimodal_runtime.py @@ -139,6 +139,24 @@ def fake_convert_contents(contents: Any, mode: Mode) -> list[dict[str, str]]: ] +def test_convert_messages_accepts_provider_owned_media_encoder() -> None: + image = Image( + source="data:image/png;base64,AA==", + media_type="image/png", + data="AA==", + ) + + converted = convert_messages( + [{"role": "user", "content": [image]}], + Mode.TOOLS, + media_converter=lambda media: {"provider": type(media).__name__}, + ) + + assert converted == [ + {"role": "user", "content": [{"provider": "Image"}]}, + ] + + def test_convert_messages_rejects_unknown_typed_message() -> None: with pytest.raises(ValueError, match="Unsupported message type"): convert_messages( diff --git a/tests/v2/test_fireworks_client.py b/tests/v2/test_fireworks_client.py deleted file mode 100644 index 2fc891d11..000000000 --- a/tests/v2/test_fireworks_client.py +++ /dev/null @@ -1,74 +0,0 @@ -"""Provider-specific tests for Fireworks v2 client factory. - -Note: Common tests (mode normalization, registry, imports) are unified in -test_client_unified.py. This file only contains Fireworks-specific tests. -""" - -from __future__ import annotations - -import pytest - -from instructor import Mode - - -# ============================================================================ -# Provider-Specific Integration Tests -# ============================================================================ -# Note: Common SDK availability tests are in test_client_unified.py - - -class TestFireworksClientWithSDK: - """Tests that require Fireworks SDK but not API key.""" - - @pytest.fixture - def fireworks_available(self): - """Check if fireworks SDK is available.""" - try: - from fireworks.client import Fireworks # noqa: F401 - - return True - except ImportError: - return False - - def test_from_fireworks_raises_without_sdk(self, fireworks_available): - """Test from_fireworks raises error when fireworks not installed.""" - if fireworks_available: - pytest.skip( - "fireworks is installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.fireworks.client import from_fireworks - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="fireworks is not installed"): - from_fireworks("not a client") # ty: ignore[no-matching-overload] - - def test_from_fireworks_with_invalid_client(self, fireworks_available): - """Test from_fireworks raises error with invalid client.""" - if not fireworks_available: - pytest.skip( - "fireworks not installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.fireworks.client import from_fireworks - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="must be an instance"): - from_fireworks("not a client") # ty: ignore[no-matching-overload] - - def test_from_fireworks_with_invalid_mode(self, fireworks_available): - """Test from_fireworks raises error with invalid mode.""" - if not fireworks_available: - pytest.skip( - "fireworks not installed" # ty: ignore[too-many-positional-arguments] - ) - - from fireworks.client import Fireworks - - from instructor.v2.providers.fireworks.client import from_fireworks - from instructor.core.exceptions import ModeError - - client = Fireworks(api_key="fake-key") - - with pytest.raises(ModeError): - from_fireworks(client, mode=Mode.RESPONSES_TOOLS) diff --git a/tests/v2/test_gemini_utils_deterministic.py b/tests/v2/test_gemini_utils_deterministic.py index d290e3107..aff6d86b1 100644 --- a/tests/v2/test_gemini_utils_deterministic.py +++ b/tests/v2/test_gemini_utils_deterministic.py @@ -9,6 +9,7 @@ from pydantic import BaseModel from instructor.v2.core.errors import ConfigurationError +from instructor.v2.core.multimodal import PDFWithGenaiFile from instructor.v2.providers.gemini import utils @@ -208,11 +209,11 @@ def test_convert_to_genai_messages_supports_strings_existing_content_and_media( _install_fake_genai_types(monkeypatch) class FakeImage: - def to_genai(self) -> str: - return "image-part" + pass image = FakeImage() monkeypatch.setattr(utils, "Image", FakeImage) + monkeypatch.setattr(utils, "media_to_genai", lambda _image: "image-part") existing = FakeContent(role="user", parts=[FakePart.from_text("existing")]) uploaded = FakeFile() @@ -233,6 +234,22 @@ def to_genai(self) -> str: assert result[3].parts[1] == "image-part" +def test_convert_to_genai_messages_preserves_uploaded_pdf_conversion( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fake_genai_types(monkeypatch) + uploaded = PDFWithGenaiFile( + source="https://generativelanguage.googleapis.com/v1beta/files/abc", + media_type="application/pdf", + data=None, + ) + monkeypatch.setattr(utils, "media_to_genai", lambda media: ("media", media)) + + result = utils.convert_to_genai_messages([{"role": "user", "content": [uploaded]}]) + + assert result[0].parts == [("media", uploaded)] + + def test_handle_genai_message_conversion_extracts_system_and_contents( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -243,7 +260,8 @@ def test_handle_genai_message_conversion_extracts_system_and_contents( lambda messages: ["converted", *messages], ) monkeypatch.setattr( - "instructor.v2.core.multimodal.extract_genai_multimodal_content", + utils, + "extract_multimodal_content", lambda contents, autodetect_images: [*contents, autodetect_images], ) @@ -331,7 +349,10 @@ def test_handle_gemini_json_guards_empty_or_missing_messages( def test_handle_gemini_tools_sets_tool_config(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(utils, "_default_safety_thresholds", lambda: None) - monkeypatch.setattr(Answer, "gemini_schema", {"name": "Answer"}, raising=False) + monkeypatch.setattr( + "instructor.v2.providers.gemini.schema.generate_gemini_schema", + lambda model: {"name": model.__name__}, + ) model, kwargs = utils.handle_gemini_tools( Answer, @@ -432,7 +453,8 @@ def test_handle_genai_structured_outputs_builds_config( lambda _messages: ["converted"], ) monkeypatch.setattr( - "instructor.v2.core.multimodal.extract_genai_multimodal_content", + utils, + "extract_multimodal_content", lambda contents, autodetect_images: [*contents, autodetect_images], ) monkeypatch.setattr(utils, "map_to_gemini_function_schema", lambda schema: schema) @@ -469,7 +491,8 @@ def test_handle_genai_tools_builds_tool_declaration( lambda _messages: ["converted"], ) monkeypatch.setattr( - "instructor.v2.core.multimodal.extract_genai_multimodal_content", + utils, + "extract_multimodal_content", lambda contents, autodetect_images: [*contents, autodetect_images], ) monkeypatch.setattr(utils, "map_to_genai_schema", lambda _schema: {"schema": "ok"}) diff --git a/tests/v2/test_genai_handlers_deterministic.py b/tests/v2/test_genai_handlers_deterministic.py index 79827bb99..6da57e054 100644 --- a/tests/v2/test_genai_handlers_deterministic.py +++ b/tests/v2/test_genai_handlers_deterministic.py @@ -120,7 +120,7 @@ def test_tools_handler_prepare_request_without_response_model( lambda messages: ["converted", *messages], ) monkeypatch.setattr( - "instructor.v2.providers.genai.handlers.extract_genai_multimodal_content", + "instructor.v2.providers.genai.handlers.extract_multimodal_content", lambda contents, autodetect_images: [*contents, autodetect_images], ) diff --git a/tests/v2/test_genai_multimodal_runtime.py b/tests/v2/test_genai_multimodal_runtime.py index 92fe41764..e4162674a 100644 --- a/tests/v2/test_genai_multimodal_runtime.py +++ b/tests/v2/test_genai_multimodal_runtime.py @@ -191,6 +191,23 @@ def test_uploaded_pdf_to_genai_falls_back_to_pdf_encoder( assert multimodal.uploaded_pdf_to_genai(pdf) is sentinel +def test_media_to_genai_routes_uploaded_pdf_through_uri_encoder( + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fake_genai(monkeypatch) + multimodal = importlib.import_module("instructor.v2.providers.genai.multimodal") + sentinel = object() + monkeypatch.setattr(multimodal, "uploaded_pdf_to_genai", lambda _pdf: sentinel) + + pdf = PDFWithGenaiFile( + source="https://generativelanguage.googleapis.com/v1beta/files/abc", + media_type="application/pdf", + data=None, + ) + + assert multimodal.media_to_genai(pdf) is sentinel + + def test_upload_new_pdf_file_waits_until_active( monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -291,7 +308,7 @@ def test_extract_multimodal_content_converts_detected_media( ) -> None: _install_fake_genai(monkeypatch) multimodal = importlib.import_module("instructor.v2.providers.genai.multimodal") - monkeypatch.setattr(Image, "to_genai", lambda _self: "converted-image") + monkeypatch.setattr(multimodal, "media_to_genai", lambda _media: "converted-image") monkeypatch.setattr( multimodal, "autodetect_media", diff --git a/tests/v2/test_groq_client.py b/tests/v2/test_groq_client.py deleted file mode 100644 index 8932f1b6d..000000000 --- a/tests/v2/test_groq_client.py +++ /dev/null @@ -1,74 +0,0 @@ -"""Provider-specific tests for Groq v2 client factory. - -Note: Common tests (mode normalization, registry, imports) are unified in -test_client_unified.py. This file only contains Groq-specific tests. -""" - -from __future__ import annotations - -import pytest - -from instructor import Mode - - -# ============================================================================ -# Provider-Specific Integration Tests -# ============================================================================ -# Note: Common SDK availability tests are in test_client_unified.py - - -class TestGroqClientWithSDK: - """Tests that require Groq SDK but not API key.""" - - @pytest.fixture - def groq_available(self): - """Check if groq SDK is available.""" - try: - import groq # noqa: F401 - - return True - except ImportError: - return False - - def test_from_groq_raises_without_sdk(self, groq_available): - """Test from_groq raises error when groq not installed.""" - if groq_available: - pytest.skip( - "groq is installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.groq.client import from_groq - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="groq is not installed"): - from_groq("not a client") # ty: ignore[no-matching-overload] - - def test_from_groq_with_invalid_client(self, groq_available): - """Test from_groq raises error with invalid client.""" - if not groq_available: - pytest.skip( - "groq not installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.groq.client import from_groq - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="must be an instance"): - from_groq("not a client") # ty: ignore[no-matching-overload] - - def test_from_groq_with_invalid_mode(self, groq_available): - """Test from_groq raises error with invalid mode.""" - if not groq_available: - pytest.skip( - "groq not installed" # ty: ignore[too-many-positional-arguments] - ) - - import groq - - from instructor.v2.providers.groq.client import from_groq - from instructor.core.exceptions import ModeError - - client = groq.Groq(api_key="fake-key") - - with pytest.raises(ModeError): - from_groq(client, mode=Mode.RESPONSES_TOOLS) diff --git a/tests/v2/test_handler_registration_unified.py b/tests/v2/test_handler_registration_unified.py index 4dff0a8cf..30cf1d0ff 100644 --- a/tests/v2/test_handler_registration_unified.py +++ b/tests/v2/test_handler_registration_unified.py @@ -13,10 +13,7 @@ # Import handler loading utilities from existing test from tests.v2.conftest import get_registered_provider_mode_pairs -from tests.v2.test_handlers_parametrized import ( - PROVIDER_HANDLER_MODES, - _ensure_handlers_loaded, -) +from tests.v2.provider_matrix import PROVIDER_HANDLER_MODES, ensure_handlers_loaded def _get_provider_mode_params(): @@ -44,7 +41,7 @@ def _get_provider_params(): @pytest.mark.parametrize("provider,mode", _get_provider_mode_params()) def test_mode_is_registered(provider: Provider, mode: Mode) -> None: """Test that all expected modes are registered.""" - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) assert mode_registry.is_registered(provider, mode), ( f"Mode {mode.value} should be registered for {provider.value}" ) @@ -53,7 +50,7 @@ def test_mode_is_registered(provider: Provider, mode: Mode) -> None: @pytest.mark.parametrize("provider,mode", _get_provider_mode_params()) def test_handlers_have_all_methods(provider: Provider, mode: Mode) -> None: """Test that all handlers have required methods.""" - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) handlers = mode_registry.get_handlers(provider, mode) assert handlers.request_handler is not None, ( @@ -70,7 +67,7 @@ def test_handlers_have_all_methods(provider: Provider, mode: Mode) -> None: @pytest.mark.parametrize("provider", _get_provider_params()) def test_get_modes_for_provider(provider: Provider) -> None: """Test getting all modes for a provider.""" - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) expected_modes = set(PROVIDER_HANDLER_MODES.get(provider, [])) registered_modes = set(mode_registry.get_modes_for_provider(provider)) @@ -82,7 +79,7 @@ def test_get_modes_for_provider(provider: Provider) -> None: @pytest.mark.parametrize("provider", _get_provider_params()) def test_provider_in_mode_providers(provider: Provider) -> None: """Test that provider is listed for its supported modes.""" - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) expected_modes = PROVIDER_HANDLER_MODES.get(provider, []) for mode in expected_modes: @@ -175,7 +172,7 @@ def test_md_json_handler_inherits_from_openai(provider: Provider) -> None: @pytest.mark.parametrize("provider", _get_provider_params()) def test_parallel_tools_not_supported_unless_listed(provider: Provider) -> None: """Test that PARALLEL_TOOLS is not supported unless in PROVIDER_HANDLER_MODES.""" - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) expected_modes = PROVIDER_HANDLER_MODES.get(provider, []) is_expected = Mode.PARALLEL_TOOLS in expected_modes is_registered = mode_registry.is_registered(provider, Mode.PARALLEL_TOOLS) @@ -189,7 +186,7 @@ def test_parallel_tools_not_supported_unless_listed(provider: Provider) -> None: @pytest.mark.parametrize("provider", _get_provider_params()) def test_responses_tools_not_supported_unless_listed(provider: Provider) -> None: """Test that RESPONSES_TOOLS is not supported unless in PROVIDER_HANDLER_MODES.""" - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) expected_modes = PROVIDER_HANDLER_MODES.get(provider, []) is_expected = Mode.RESPONSES_TOOLS in expected_modes is_registered = mode_registry.is_registered(provider, Mode.RESPONSES_TOOLS) diff --git a/tests/v2/test_handlers_parametrized.py b/tests/v2/test_handlers_parametrized.py index 255309714..88434fdfd 100644 --- a/tests/v2/test_handlers_parametrized.py +++ b/tests/v2/test_handlers_parametrized.py @@ -9,7 +9,6 @@ import importlib.util import json from dataclasses import dataclass -from pathlib import Path from types import SimpleNamespace from typing import Any @@ -19,60 +18,11 @@ from instructor import Mode, Provider from instructor.processing.function_calls import ResponseSchema from instructor.v2.core.registry import mode_registry -from tests.v2.provider_matrix import PROVIDER_HANDLER_MODES - -_PROJECT_ROOT = Path(__file__).resolve().parents[2] -_HANDLER_MODULE_PATHS: dict[Provider, Path] = { - Provider.OPENAI: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.ANYSCALE: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.TOGETHER: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.DATABRICKS: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.DEEPSEEK: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.ANTHROPIC: _PROJECT_ROOT / "instructor/v2/providers/anthropic/handlers.py", - Provider.GENAI: _PROJECT_ROOT / "instructor/v2/providers/genai/handlers.py", - Provider.GEMINI: _PROJECT_ROOT / "instructor/v2/providers/gemini/handlers.py", - Provider.VERTEXAI: _PROJECT_ROOT / "instructor/v2/providers/vertexai/handlers.py", - Provider.COHERE: _PROJECT_ROOT / "instructor/v2/providers/cohere/handlers.py", - Provider.PERPLEXITY: _PROJECT_ROOT - / "instructor/v2/providers/perplexity/handlers.py", - Provider.XAI: _PROJECT_ROOT / "instructor/v2/providers/xai/handlers.py", - Provider.GROQ: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.MISTRAL: _PROJECT_ROOT / "instructor/v2/providers/mistral/handlers.py", - Provider.FIREWORKS: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.BEDROCK: _PROJECT_ROOT / "instructor/v2/providers/bedrock/handlers.py", - Provider.CEREBRAS: _PROJECT_ROOT / "instructor/v2/providers/openai/handlers.py", - Provider.WRITER: _PROJECT_ROOT / "instructor/v2/providers/writer/handlers.py", - Provider.OPENROUTER: _PROJECT_ROOT - / "instructor/v2/providers/openrouter/handlers.py", -} -_HANDLERS_LOADED: set[Provider] = set() - - -def _ensure_handlers_loaded(provider: Provider) -> None: - if provider in _HANDLERS_LOADED: - return - provider_modes = PROVIDER_HANDLER_MODES.get(provider, []) - if provider_modes and all( - mode_registry.is_registered(provider, mode) for mode in provider_modes - ): - _HANDLERS_LOADED.add(provider) - return - handler_path = _HANDLER_MODULE_PATHS.get(provider) - if handler_path is None: - return - spec = importlib.util.spec_from_file_location( - f"tests.v2.handlers_{provider.value}", - handler_path, - ) - if spec is None or spec.loader is None: - raise ImportError(f"Could not load handler module for {provider}") - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - _HANDLERS_LOADED.add(provider) +from tests.v2.provider_matrix import PROVIDER_HANDLER_MODES, ensure_handlers_loaded def _get_handlers(provider: Provider, mode: Mode): - _ensure_handlers_loaded(provider) + ensure_handlers_loaded(provider) return mode_registry.get_handlers(provider, mode) diff --git a/tests/v2/test_legacy_provider_compat.py b/tests/v2/test_legacy_provider_compat.py index 0b55b1a25..2be0d83bc 100644 --- a/tests/v2/test_legacy_provider_compat.py +++ b/tests/v2/test_legacy_provider_compat.py @@ -1,6 +1,32 @@ from __future__ import annotations import importlib +import os +from pathlib import Path +import subprocess +import sys + +import pytest + +from instructor.v2.core.provider_specs import PROVIDER_SPECS +from instructor.v2.core.providers import Provider + + +_ROOT = Path(__file__).parents[2] +_LEGACY_OPTIONAL_CLIENTS = ( + Provider.ANTHROPIC, + Provider.BEDROCK, + Provider.CEREBRAS, + Provider.COHERE, + Provider.FIREWORKS, + Provider.GEMINI, + Provider.GENAI, + Provider.GROQ, + Provider.MISTRAL, + Provider.VERTEXAI, + Provider.WRITER, + Provider.XAI, +) def test_legacy_provider_modules_remain_importable() -> None: @@ -58,3 +84,49 @@ def test_legacy_provider_utils_forward_to_v2_symbols() -> None: "instructor.v2.providers.gemini.utils" ).map_to_gemini_function_schema ) + + +@pytest.mark.parametrize("provider", _LEGACY_OPTIONAL_CLIENTS) +def test_legacy_client_facades_remain_lazy_without_optional_sdks( + provider: Provider, +) -> None: + spec = PROVIDER_SPECS[provider] + assert spec.sdk_module is not None + blocked_sdk = spec.sdk_module + legacy_module = f"instructor.providers.{provider.value}.client" + implementation_module = f"instructor.v2.providers.{provider.value}.client" + probe = """ +import builtins +import importlib +import sys + +blocked, facade, implementation = sys.argv[1:] +original_import = builtins.__import__ + +def block_sdk(name, globals=None, locals=None, fromlist=(), level=0): + if name == blocked or name.startswith(blocked + "."): + raise ImportError(f"blocked optional SDK: {name}") + return original_import(name, globals, locals, fromlist, level) + +builtins.__import__ = block_sdk +importlib.import_module(facade) +assert implementation not in sys.modules, implementation +""" + env = {**os.environ, "PYTHONPATH": str(_ROOT)} + + result = subprocess.run( + [ + sys.executable, + "-c", + probe, + blocked_sdk, + legacy_module, + implementation_module, + ], + check=False, + capture_output=True, + text=True, + env=env, + ) + + assert result.returncode == 0, result.stderr diff --git a/tests/v2/test_mistral_client.py b/tests/v2/test_mistral_client.py deleted file mode 100644 index f4a777f38..000000000 --- a/tests/v2/test_mistral_client.py +++ /dev/null @@ -1,58 +0,0 @@ -"""Provider-specific tests for Mistral v2 client factory. - -Note: Common tests (mode normalization, registry, imports, errors) are unified in -test_client_unified.py. This file only contains Mistral-specific tests. -""" - -from __future__ import annotations - -import pytest - - -# ============================================================================ -# Provider-Specific Integration Tests -# ============================================================================ -# Note: Common SDK availability tests are in test_client_unified.py - - -class TestMistralClientWithSDK: - """Tests for Mistral client factory that require the SDK.""" - - def test_from_mistral_raises_without_sdk(self): - """Test from_mistral raises helpful error when SDK not installed.""" - import importlib.util - - # This test checks behavior when mistralai is not installed - if importlib.util.find_spec("mistralai") is not None: - pytest.skip( - "mistralai is installed, skipping SDK-not-installed test" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.mistral.client import from_mistral - from instructor.core.exceptions import ClientError - - # Should raise ClientError about missing SDK - with pytest.raises(ClientError) as exc_info: - from_mistral(None) # type: ignore - - assert "mistralai is not installed" in str(exc_info.value) - - @pytest.mark.skipif(True, reason="Requires mistralai SDK") - def test_from_mistral_with_invalid_client(self): - """Test from_mistral raises error with invalid client type.""" - pass - - @pytest.mark.skipif(True, reason="Requires mistralai SDK") - def test_from_mistral_with_invalid_mode(self): - """Test from_mistral raises error with invalid mode.""" - pass - - @pytest.mark.skipif(True, reason="Requires mistralai SDK") - def test_from_mistral_sync_client(self): - """Test from_mistral creates sync Instructor.""" - pass - - @pytest.mark.skipif(True, reason="Requires mistralai SDK") - def test_from_mistral_async_client(self): - """Test from_mistral creates async Instructor with use_async=True.""" - pass diff --git a/tests/v2/test_model_mode_benchmark.py b/tests/v2/test_model_mode_benchmark.py new file mode 100644 index 000000000..f13bfd7f6 --- /dev/null +++ b/tests/v2/test_model_mode_benchmark.py @@ -0,0 +1,152 @@ +from __future__ import annotations + +from collections.abc import Iterable +import importlib.util +from pathlib import Path +import sys +from typing import Any, get_origin + +from instructor.v2.core.mode import Mode +from instructor.v2.core.provider_specs import PROVIDER_SPECS +from instructor.v2.core.providers import Provider + + +def _load_benchmark() -> Any: + path = Path(__file__).parents[2] / "examples" / "v2-model-mode-benchmark" / "run.py" + spec = importlib.util.spec_from_file_location("v2_model_mode_benchmark", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[spec.name] = module + spec.loader.exec_module(module) + return module + + +benchmark: Any = _load_benchmark() + + +def test_default_grid_includes_every_declared_provider_mode() -> None: + cases = benchmark.build_cases() + expected = { + (provider, mode) + for provider, spec in PROVIDER_SPECS.items() + if spec.handler_module is not None + for mode in spec.supported_modes + } + + assert {(case.provider, case.mode) for case in cases} == expected + + +def test_explicit_models_can_compare_multiple_models_in_selected_modes() -> None: + cases = benchmark.build_cases( + models=("openai/model-a", "openai/model-b"), + modes=(Mode.TOOLS, Mode.JSON_SCHEMA), + ) + + assert [(case.model, case.mode) for case in cases] == [ + ("openai/model-a", Mode.TOOLS), + ("openai/model-a", Mode.JSON_SCHEMA), + ("openai/model-b", Mode.TOOLS), + ("openai/model-b", Mode.JSON_SCHEMA), + ] + + +def test_missing_key_skips_cell_without_creating_client( + monkeypatch: Any, +) -> None: + monkeypatch.delenv("OPENAI_API_KEY", raising=False) + created = False + + def factory(*args: Any, **kwargs: Any) -> Any: # noqa: ARG001 + nonlocal created + created = True + raise AssertionError("client should not be created") + + result = benchmark.run_case( + benchmark.BenchmarkCase(Provider.OPENAI, "openai/gpt-4o-mini", Mode.TOOLS), + trials=1, + client_factory=factory, + ) + + assert result.status == "skipped" + assert result.detail == "missing credential: OPENAI_API_KEY" + assert created is False + + +def test_successful_cell_records_correctness_and_rendered_latency( + monkeypatch: Any, +) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "test-key") + monkeypatch.setattr(benchmark, "_module_available", lambda _module: True) + + class FakeClient: + def create(self, **kwargs: Any) -> Any: # noqa: ARG002 + return benchmark.Person(name="Jason", age=36) + + def factory(*_args: Any, **_kwargs: Any) -> FakeClient: + return FakeClient() + + result = benchmark.run_case( + benchmark.BenchmarkCase(Provider.OPENAI, "openai/gpt-4o-mini", Mode.TOOLS), + trials=2, + client_factory=factory, + ) + rendered = benchmark.render_markdown([result]) + + assert result.status == "passed" + assert result.successes == 2 + assert result.median_ms is not None + assert "## Ranked completed cells" in rendered + assert "| `openai/gpt-4o-mini` | `TOOLS` | passed | 2/2 |" in rendered + + +def test_parallel_cell_uses_iterable_response_model(monkeypatch: Any) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "test-key") + monkeypatch.setattr(benchmark, "_module_available", lambda _module: True) + seen_response_model: Any = None + + class FakeClient: + def create(self, **kwargs: Any) -> Any: + nonlocal seen_response_model + seen_response_model = kwargs["response_model"] + return [benchmark.Person(name="Jason", age=36)] + + result = benchmark.run_case( + benchmark.BenchmarkCase( + Provider.OPENAI, + "openai/gpt-4o-mini", + Mode.PARALLEL_TOOLS, + ), + trials=1, + client_factory=lambda *_args, **_kwargs: FakeClient(), + ) + + assert result.status == "passed" + assert get_origin(seen_response_model) is Iterable + + +def test_cloud_auth_cells_require_explicit_opt_in(monkeypatch: Any) -> None: + monkeypatch.setattr(benchmark, "_module_available", lambda _module: True) + case = benchmark.BenchmarkCase( + Provider.BEDROCK, + "bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0", + Mode.TOOLS, + ) + + skipped = benchmark.run_case(case, trials=1) + + assert skipped.status == "skipped" + assert skipped.detail == ( + "provider uses ambient/cloud credentials; pass --allow-cloud-auth" + ) + + class FakeClient: + def create(self, **kwargs: Any) -> Any: # noqa: ARG002 + return benchmark.Person(name="Jason", age=36) + + executed = benchmark.run_case( + case, + trials=1, + allow_cloud_auth=True, + client_factory=lambda *_args, **_kwargs: FakeClient(), + ) + assert executed.status == "passed" diff --git a/tests/v2/test_provider_capability_runtime.py b/tests/v2/test_provider_capability_runtime.py new file mode 100644 index 000000000..26c94bce2 --- /dev/null +++ b/tests/v2/test_provider_capability_runtime.py @@ -0,0 +1,151 @@ +from __future__ import annotations + +import inspect +from types import SimpleNamespace +from typing import Any + +import pytest + +from instructor import Mode, Provider +from instructor.v2.core.multimodal import Audio, Image, PDF +from instructor.v2.core.provider_specs import PROVIDER_SPECS +from instructor.v2.core.registry import mode_registry +from tests.v2.provider_matrix import ( + ITERABLE_STREAM_CASES, + PARTIAL_STREAM_CASES, + TYPED_MULTIMODAL_CASES, + ensure_handlers_loaded, +) + +_PAYLOAD = '{"answer": 4}' +_STREAM_CASES = tuple(dict.fromkeys((*PARTIAL_STREAM_CASES, *ITERABLE_STREAM_CASES))) + + +class _GeminiFunctionCall: + @staticmethod + def to_dict(_value: Any) -> dict[str, dict[str, int]]: + return {"args": {"answer": 4}} + + +def _stream_chunk(provider: Provider, mode: Mode) -> Any: + module = PROVIDER_SPECS[provider].handler_module + if module == "instructor.v2.providers.anthropic.handlers": + return SimpleNamespace( + delta=SimpleNamespace(partial_json=_PAYLOAD, text=_PAYLOAD) + ) + if module == "instructor.v2.providers.genai.handlers": + part = SimpleNamespace(function_call=SimpleNamespace(args={"answer": 4})) + return SimpleNamespace( + text=_PAYLOAD, + candidates=[SimpleNamespace(content=SimpleNamespace(parts=[part]))], + ) + if module == "instructor.v2.providers.gemini.handlers": + part = SimpleNamespace(function_call=_GeminiFunctionCall()) + return SimpleNamespace( + text=_PAYLOAD, + candidates=[SimpleNamespace(content=SimpleNamespace(parts=[part]))], + ) + if module == "instructor.v2.providers.vertexai.handlers": + part = SimpleNamespace( + text=_PAYLOAD, + function_call=SimpleNamespace(args={"answer": 4}), + ) + return SimpleNamespace( + candidates=[SimpleNamespace(content=SimpleNamespace(parts=[part]))] + ) + if module == "instructor.v2.providers.cohere.handlers": + return SimpleNamespace(text=_PAYLOAD) + if module == "instructor.v2.providers.mistral.handlers": + delta = SimpleNamespace( + content=_PAYLOAD, + tool_calls=[SimpleNamespace(function=SimpleNamespace(arguments=_PAYLOAD))], + ) + return SimpleNamespace( + data=SimpleNamespace(choices=[SimpleNamespace(delta=delta)]) + ) + if mode is Mode.RESPONSES_TOOLS: + from openai.types.responses import ResponseFunctionCallArgumentsDeltaEvent + + return ResponseFunctionCallArgumentsDeltaEvent.model_validate( + { + "delta": _PAYLOAD, + "item_id": "call_1", + "output_index": 0, + "sequence_number": 0, + "type": "response.function_call_arguments.delta", + } + ) + delta = SimpleNamespace( + content=_PAYLOAD, + tool_calls=[SimpleNamespace(function=SimpleNamespace(arguments=_PAYLOAD))], + ) + return SimpleNamespace(choices=[SimpleNamespace(delta=delta)]) + + +@pytest.mark.parametrize(("provider", "mode"), _STREAM_CASES) +def test_advertised_stream_contract_extracts_payload( + provider: Provider, mode: Mode +) -> None: + ensure_handlers_loaded(provider, skip_missing_dependency=True) + extractor = mode_registry.get_handlers(provider, mode).stream_extractor + + assert extractor is not None + assert "answer" in "".join(extractor([_stream_chunk(provider, mode)])) + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("provider", "mode"), _STREAM_CASES) +async def test_advertised_async_stream_contract_extracts_payload( + provider: Provider, mode: Mode +) -> None: + ensure_handlers_loaded(provider, skip_missing_dependency=True) + extractor = mode_registry.get_handlers(provider, mode).stream_extractor_async + + async def chunks(): + yield _stream_chunk(provider, mode) + + assert extractor is not None + extracted: Any = extractor(chunks()) + if inspect.isawaitable(extracted): + extracted = await extracted + assert "answer" in "".join([part async for part in extracted]) + + +def _media(provider: Provider, media_type: str) -> Image | Audio | PDF: + if media_type == "image": + return Image( + source="data:image/png;base64,AA==", media_type="image/png", data="AA==" + ) + if media_type == "audio": + return Audio( + source="data:audio/wav;base64,AA==", media_type="audio/wav", data="AA==" + ) + if provider is Provider.MISTRAL: + return PDF(source="https://example.com/document.pdf", data=None) + return PDF(source="data:application/pdf;base64,AA==", data="AA==") + + +@pytest.mark.parametrize(("provider", "media_type"), TYPED_MULTIMODAL_CASES) +def test_advertised_multimodal_contract_dispatches_typed_media( + provider: Provider, media_type: str, monkeypatch: pytest.MonkeyPatch +) -> None: + media = _media(provider, media_type) + if provider is Provider.GENAI: + from instructor.v2.providers.genai import multimodal + + part = SimpleNamespace( + from_bytes=lambda **kwargs: kwargs, + from_uri=lambda **kwargs: kwargs, + ) + monkeypatch.setattr(multimodal, "_types", lambda: SimpleNamespace(Part=part)) + assert multimodal.media_to_genai(media) + return + + ensure_handlers_loaded(provider, skip_missing_dependency=True) + modes = PROVIDER_SPECS[provider].supported_modes + mode = Mode.TOOLS if Mode.TOOLS in modes else modes[0] + converter = mode_registry.get_handlers(provider, mode).message_converter + + assert converter is not None + converted = converter([{"role": "user", "content": [media]}]) + assert converted[0]["content"] diff --git a/tests/v2/test_provider_multimodal_ownership.py b/tests/v2/test_provider_multimodal_ownership.py new file mode 100644 index 000000000..fcc28e16f --- /dev/null +++ b/tests/v2/test_provider_multimodal_ownership.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +import pytest + +from instructor.v2.core.multimodal import Image, PDF, PDFWithCacheControl +from instructor.v2.providers.anthropic import handlers as anthropic_handlers +from instructor.v2.providers.mistral import handlers as mistral_handlers +from instructor.v2.providers.openai import handlers as openai_handlers + + +def _image() -> Image: + return Image( + source="data:image/png;base64,AA==", + media_type="image/png", + data="AA==", + ) + + +def test_openai_handler_uses_openai_media_encoder( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + openai_handlers, + "media_to_openai", + lambda _media, mode: {"provider": "openai", "mode": mode.value}, + ) + + converted = openai_handlers.OpenAIToolsHandler().convert_messages( + [{"role": "user", "content": [_image()]}] + ) + + assert converted[0]["content"] == [{"provider": "openai", "mode": "tool_call"}] + + +def test_anthropic_handler_uses_anthropic_media_encoder( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + anthropic_handlers, + "media_to_anthropic", + lambda _media: {"provider": "anthropic"}, + ) + + converted = anthropic_handlers.AnthropicToolsHandler().convert_messages( + [{"role": "user", "content": [_image()]}] + ) + + assert converted[0]["content"] == [{"provider": "anthropic"}] + + +def test_mistral_handler_uses_mistral_media_encoder( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + mistral_handlers, + "media_to_mistral", + lambda _media, mode: {"provider": "mistral", "mode": mode.value}, + ) + + converted = mistral_handlers.MistralToolsHandler().convert_messages( + [{"role": "user", "content": [_image()]}] + ) + + assert converted[0]["content"] == [{"provider": "mistral", "mode": "mistral_tools"}] + + +@pytest.mark.parametrize( + ("handlers", "handler_cls", "factory_name"), + [ + (openai_handlers, openai_handlers.OpenAIToolsHandler, "image_from_params"), + ( + anthropic_handlers, + anthropic_handlers.AnthropicToolsHandler, + "image_from_params", + ), + (mistral_handlers, mistral_handlers.MistralToolsHandler, "image_from_params"), + ], +) +def test_active_handlers_own_image_shorthand_conversion( + monkeypatch: pytest.MonkeyPatch, + handlers: object, + handler_cls: type[object], + factory_name: str, +) -> None: + image = _image() + monkeypatch.setattr(handlers, factory_name, lambda _params: image) + + converted = handler_cls().convert_messages( # type: ignore[attr-defined] + [{"role": "user", "content": [{"type": "image", "source": "ignored"}]}], + autodetect_images=True, + ) + + assert converted[0]["content"] + + +def test_anthropic_cache_compatibility_is_data_driven() -> None: + ordinary_pdf = PDF( + source="data:application/pdf;base64,AA==", + media_type="application/pdf", + data="AA==", + ) + cacheable_pdf = PDFWithCacheControl( + source=ordinary_pdf.source, + media_type=ordinary_pdf.media_type, + data=ordinary_pdf.data, + ) + + assert "cache_control" not in anthropic_handlers.media_to_anthropic(ordinary_pdf) + assert anthropic_handlers.media_to_anthropic(cacheable_pdf)["cache_control"] == { + "type": "ephemeral" + } diff --git a/tests/v2/test_provider_specs.py b/tests/v2/test_provider_specs.py index 4c59de961..295fc8276 100644 --- a/tests/v2/test_provider_specs.py +++ b/tests/v2/test_provider_specs.py @@ -2,9 +2,22 @@ from __future__ import annotations +import warnings + +import pytest + +from instructor import Mode from instructor import Provider -from instructor.v2.auto_client import _PROVIDER_BUILDERS, supported_providers +from instructor.v2.auto_client import supported_providers from instructor.v2.core.provider_specs import ALIAS_TO_PROVIDER, PROVIDER_SPECS +from instructor.v2.core.providers import ( + normalize_mode_for_provider, + provider_from_mode, +) +from tests.v2.provider_matrix import ( + EXPLICIT_PARALLEL_PROVIDERS, + TYPED_MULTIMODAL_PROVIDERS, +) def test_supported_provider_aliases_come_from_manifest() -> None: @@ -12,7 +25,12 @@ def test_supported_provider_aliases_come_from_manifest() -> None: def test_supported_provider_aliases_have_auto_client_builders() -> None: - assert set(supported_providers) <= set(_PROVIDER_BUILDERS) + routed_aliases = { + alias + for alias, provider in ALIAS_TO_PROVIDER.items() + if PROVIDER_SPECS[provider].model_builder_module is not None + } + assert set(supported_providers) <= routed_aliases def test_compatibility_aliases_point_to_canonical_providers() -> None: @@ -30,3 +48,75 @@ def test_first_class_specs_are_self_canonical() -> None: }: continue assert spec.canonical_provider is spec.provider + + +@pytest.mark.parametrize( + "spec", + [spec for spec in PROVIDER_SPECS.values() if spec.handler_module is not None], + ids=lambda spec: spec.provider.value, +) +def test_advertised_streaming_modes_are_supported_modes(spec) -> None: + advertised_modes = { + *spec.capabilities.partial_stream_modes, + *spec.capabilities.iterable_stream_modes, + } + assert advertised_modes <= set(spec.supported_modes) + + +@pytest.mark.parametrize( + "spec", + [spec for spec in PROVIDER_SPECS.values() if spec.handler_module is not None], + ids=lambda spec: spec.provider.value, +) +def test_explicit_parallel_contract_matches_supported_mode(spec) -> None: + assert spec.capabilities.explicit_parallel_tools is ( + Mode.PARALLEL_TOOLS in spec.supported_modes + ) + + +def test_known_public_streaming_gaps_are_not_advertised() -> None: + assert ( + Mode.MD_JSON + not in PROVIDER_SPECS[Provider.XAI].capabilities.partial_stream_modes + ) + assert not PROVIDER_SPECS[Provider.GEMINI].capabilities.iterable_stream_modes + assert not PROVIDER_SPECS[Provider.VERTEXAI].capabilities.iterable_stream_modes + + +def test_multimodal_contract_is_explicitly_typed_media_only() -> None: + assert Provider.GENAI in TYPED_MULTIMODAL_PROVIDERS + assert Provider.BEDROCK not in TYPED_MULTIMODAL_PROVIDERS + assert PROVIDER_SPECS[Provider.ANTHROPIC].capabilities.multimodal_inputs == ( + "image", + "pdf", + ) + + +def test_explicit_parallel_contract_is_driven_by_manifest() -> None: + assert set(EXPLICIT_PARALLEL_PROVIDERS) == { + provider + for provider, spec in PROVIDER_SPECS.items() + if spec.handler_module is not None + and Mode.PARALLEL_TOOLS in spec.supported_modes + } + + +def test_legacy_mode_ownership_and_normalization_come_from_manifest() -> None: + owners: dict[Mode, list[Provider]] = {} + with warnings.catch_warnings(): + warnings.simplefilter("ignore", DeprecationWarning) + for provider, spec in PROVIDER_SPECS.items(): + for legacy_mode, canonical_mode in spec.legacy_modes.items(): + owners.setdefault(legacy_mode, []).append(provider) + assert ( + normalize_mode_for_provider(legacy_mode, provider) is canonical_mode + ) + + for legacy_mode, providers in owners.items(): + canonical_owners = { + PROVIDER_SPECS[provider].canonical_provider for provider in providers + } + if len(canonical_owners) == 1: + assert provider_from_mode(legacy_mode) is canonical_owners.pop() + else: + assert provider_from_mode(legacy_mode, Provider.COHERE) is Provider.COHERE diff --git a/tests/v2/test_retry_runtime.py b/tests/v2/test_retry_runtime.py index de956b304..f84231eff 100644 --- a/tests/v2/test_retry_runtime.py +++ b/tests/v2/test_retry_runtime.py @@ -57,6 +57,13 @@ def test_initialize_usage_returns_openai_usage_shape() -> None: def test_retry_sync_v2_returns_raw_result_when_no_response_model() -> None: + arguments: list[tuple[tuple[Any, ...], dict[str, Any]]] = [] + hooks = Hooks() + hooks.on( + "completion:kwargs", + lambda *args, **kwargs: arguments.append((args, kwargs)), + ) + def fake_func(*args: Any, **kwargs: Any) -> str: return f"{args[0]}:{kwargs['suffix']}" @@ -70,10 +77,11 @@ def fake_func(*args: Any, **kwargs: Any) -> str: args=("hello",), kwargs={"suffix": "world"}, strict=True, - hooks=None, + hooks=hooks, ) assert result == "hello:world" + assert arguments == [(("hello",), {"suffix": "world"})] def test_retry_sync_v2_reasks_after_validation_error( @@ -804,3 +812,246 @@ def fake_parser(**_kwargs: Any) -> Answer: assert result == Answer(value=42) assert parser_calls == 2 + + +def test_retry_sync_v2_emits_raw_call_failure_hooks() -> None: + events: list[tuple[str, Exception, dict[str, Any]]] = [] + hooks = Hooks() + hooks.on( + "completion:error", + lambda error, **metadata: events.append(("error", error, metadata)), + ) + hooks.on( + "completion:last_attempt", + lambda error, **metadata: events.append(("last", error, metadata)), + ) + + def fail(**_kwargs: Any) -> None: + raise RuntimeError("provider failed") + + with pytest.raises(RuntimeError, match="provider failed"): + retry_sync_v2( + func=fail, + response_model=None, + provider=Provider.OPENAI, + mode=Mode.TOOLS, + context=None, + max_retries=1, + args=(), + kwargs={}, + strict=True, + hooks=hooks, + ) + + assert [event[0] for event in events] == ["error", "last"] + assert events[0][2] == { + "attempt_number": 1, + "max_attempts": 1, + "is_last_attempt": True, + } + + +def test_retry_sync_v2_defers_last_attempt_for_custom_retry_policy() -> None: + events: list[tuple[str, bool]] = [] + hooks = Hooks() + hooks.on( + "completion:error", + lambda _error, **metadata: events.append( + ("error", metadata["is_last_attempt"]) + ), + ) + hooks.on( + "completion:last_attempt", + lambda _error, **metadata: events.append(("last", metadata["is_last_attempt"])), + ) + attempts = 0 + + def fail(**_kwargs: Any) -> None: + nonlocal attempts + attempts += 1 + raise RuntimeError("provider failed") + + with pytest.raises(InstructorRetryException, match="provider failed"): + retry_sync_v2( + func=fail, + response_model=Answer, + provider=Provider.OPENAI, + mode=Mode.TOOLS, + context=None, + max_retries=Retrying( + stop=stop_after_attempt(2), + retry=retry_if_exception_type(RuntimeError), + reraise=True, + ), + args=(), + kwargs={}, + strict=True, + hooks=hooks, + ) + + assert attempts == 2 + assert events == [("error", False), ("error", True), ("last", True)] + + +def test_retry_sync_v2_preserves_provider_error_when_tenacity_wraps() -> None: + failures: list[Exception] = [] + hooks = Hooks() + hooks.on("completion:last_attempt", lambda error: failures.append(error)) + + def fail(**_kwargs: Any) -> None: + raise RuntimeError("provider failed") + + with pytest.raises(InstructorRetryException) as exc_info: + retry_sync_v2( + func=fail, + response_model=Answer, + provider=Provider.OPENAI, + mode=Mode.TOOLS, + context=None, + max_retries=Retrying( + stop=stop_after_attempt(1), + retry=retry_if_exception_type(RuntimeError), + ), + args=(), + kwargs={}, + strict=True, + hooks=hooks, + ) + + assert str(exc_info.value) == "provider failed" + assert len(failures) == 1 + assert isinstance(failures[0], RuntimeError) + + +@pytest.mark.asyncio +async def test_retry_async_v2_emits_raw_call_failure_hooks() -> None: + events: list[str] = [] + hooks = Hooks() + hooks.on("completion:error", lambda _error, **_metadata: events.append("error")) + hooks.on( + "completion:last_attempt", + lambda _error, **_metadata: events.append("last"), + ) + + async def fail(**_kwargs: Any) -> None: + raise RuntimeError("provider failed") + + with pytest.raises(RuntimeError, match="provider failed"): + await retry_async_v2( + func=fail, + response_model=None, + provider=Provider.OPENAI, + mode=Mode.TOOLS, + context=None, + max_retries=1, + args=(), + kwargs={}, + strict=True, + hooks=hooks, + ) + + assert events == ["error", "last"] + + +def test_retry_sync_v2_emits_post_response_failure_hooks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: list[tuple[str, Exception]] = [] + hooks = Hooks() + hooks.on( + "completion:error", lambda error, **_metadata: events.append(("error", error)) + ) + hooks.on( + "completion:last_attempt", + lambda error, **_metadata: events.append(("last", error)), + ) + + monkeypatch.setattr( + "instructor.v2.core.retry.RegistryValidationMixin.validate_mode_registration", + lambda _provider, _mode: None, + ) + monkeypatch.setattr( + "instructor.v2.core.retry.mode_registry.get_handlers", + lambda _provider, _mode: SimpleNamespace( + response_parser=lambda **_kwargs: (_ for _ in ()).throw( + ValueError("parser failed") + ), + reask_handler=lambda **kwargs: kwargs["kwargs"], + ), + ) + monkeypatch.setattr( + "instructor.v2.core.retry.update_total_usage", + lambda **_kwargs: None, + ) + + with pytest.raises(InstructorRetryException) as exc_info: + retry_sync_v2( + func=lambda **_kwargs: object(), + response_model=Answer, + provider=Provider.OPENAI, + mode=Mode.TOOLS, + context=None, + max_retries=2, + args=(), + kwargs={}, + strict=True, + hooks=hooks, + ) + + assert isinstance(exc_info.value.__cause__, ValueError) + assert [name for name, _error in events] == ["error", "last"] + assert all(str(error) == "parser failed" for _name, error in events) + + +@pytest.mark.asyncio +async def test_retry_async_v2_emits_post_response_failure_hooks( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: list[tuple[str, Exception]] = [] + hooks = Hooks() + hooks.on( + "completion:error", lambda error, **_metadata: events.append(("error", error)) + ) + hooks.on( + "completion:last_attempt", + lambda error, **_metadata: events.append(("last", error)), + ) + + monkeypatch.setattr( + "instructor.v2.core.retry.RegistryValidationMixin.validate_mode_registration", + lambda _provider, _mode: None, + ) + monkeypatch.setattr( + "instructor.v2.core.retry.mode_registry.get_handlers", + lambda _provider, _mode: SimpleNamespace( + response_parser=lambda **_kwargs: (_ for _ in ()).throw( + ValueError("parser failed") + ), + reask_handler=lambda **kwargs: kwargs["kwargs"], + ), + ) + monkeypatch.setattr( + "instructor.v2.core.retry.update_total_usage", + lambda **_kwargs: None, + ) + + async def response(**_kwargs: Any) -> object: + return object() + + with pytest.raises(InstructorRetryException) as exc_info: + await retry_async_v2( + func=response, + response_model=Answer, + provider=Provider.OPENAI, + mode=Mode.TOOLS, + context=None, + max_retries=2, + args=(), + kwargs={}, + strict=True, + hooks=hooks, + ) + + assert isinstance(exc_info.value.__cause__, ValueError) + assert [name for name, _error in events] == ["error", "last"] + assert all(str(error) == "parser failed" for _name, error in events) diff --git a/tests/v2/test_writer_client.py b/tests/v2/test_writer_client.py deleted file mode 100644 index dc9c912d9..000000000 --- a/tests/v2/test_writer_client.py +++ /dev/null @@ -1,93 +0,0 @@ -"""Unit tests for Writer v2 client factory. - -These tests verify client factory behavior without requiring API keys. -""" - -from __future__ import annotations - -import pytest - -from instructor import Mode - - -# ============================================================================ -# Import Tests -# ============================================================================ - - -class TestWriterImports: - """Tests for Writer module imports.""" - - def test_handlers_importable(self): - """Test Writer handlers are importable.""" - from instructor.v2.providers.writer.handlers import ( - WriterJSONSchemaHandler, - WriterMDJSONHandler, - WriterToolsHandler, - ) - - assert WriterToolsHandler is not None - assert WriterJSONSchemaHandler is not None - assert WriterMDJSONHandler is not None - - -# ============================================================================ -# Integration Tests (require Writer SDK but not API key) -# ============================================================================ - - -class TestWriterClientWithSDK: - """Tests that require Writer SDK but not API key.""" - - @pytest.fixture - def writer_available(self): - """Check if writerai SDK is available.""" - try: - from writerai import Writer # noqa: F401 - - return True - except ImportError: - return False - - def test_from_writer_raises_without_sdk(self, writer_available): - """Test from_writer raises error when writerai not installed.""" - if writer_available: - pytest.skip( - "writerai is installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.writer.client import from_writer - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="writerai is not installed"): - from_writer("not a client") # ty: ignore[no-matching-overload] - - def test_from_writer_with_invalid_client(self, writer_available): - """Test from_writer raises error with invalid client.""" - if not writer_available: - pytest.skip( - "writerai not installed" # ty: ignore[too-many-positional-arguments] - ) - - from instructor.v2.providers.writer.client import from_writer - from instructor.core.exceptions import ClientError - - with pytest.raises(ClientError, match="must be an instance"): - from_writer("not a client") # ty: ignore[no-matching-overload] - - def test_from_writer_with_invalid_mode(self, writer_available): - """Test from_writer raises error with invalid mode.""" - if not writer_available: - pytest.skip( - "writerai not installed" # ty: ignore[too-many-positional-arguments] - ) - - from writerai import Writer - - from instructor.v2.providers.writer.client import from_writer - from instructor.core.exceptions import ModeError - - client = Writer(api_key="fake-key") - - with pytest.raises(ModeError): - from_writer(client, mode=Mode.RESPONSES_TOOLS) diff --git a/ty-tests.toml b/ty-tests.toml index 53f75c554..cadfd8088 100644 --- a/ty-tests.toml +++ b/ty-tests.toml @@ -4,4 +4,5 @@ exclude = [ ".venv/", "tests/llm/", "tests/docs/", + "tests/typing/test_public_surface.py", ]