diff --git a/tests/pass_through_unit_tests/messages_api_structured_output/base_anthropic_messages_structured_output_test.py b/tests/pass_through_unit_tests/messages_api_structured_output/base_anthropic_messages_structured_output_test.py index e1df80c48e1..9346345eef6 100644 --- a/tests/pass_through_unit_tests/messages_api_structured_output/base_anthropic_messages_structured_output_test.py +++ b/tests/pass_through_unit_tests/messages_api_structured_output/base_anthropic_messages_structured_output_test.py @@ -9,7 +9,7 @@ import json import os import sys from abc import ABC, abstractmethod -from typing import Any, Dict, List +from typing import Any, Dict, List, Optional sys.path.insert(0, os.path.abspath("../../..")) @@ -23,6 +23,10 @@ class BaseAnthropicMessagesStructuredOutputTest(ABC): Subclasses must implement: - get_model(): Returns the model string to use for tests + + Subclasses may optionally implement: + - get_api_base(): Returns the API base URL (for Azure, etc.) + - get_api_key(): Returns the API key (for Azure, etc.) """ @abstractmethod @@ -32,6 +36,18 @@ class BaseAnthropicMessagesStructuredOutputTest(ABC): """ pass + def get_api_base(self) -> Optional[str]: + """ + Returns the API base URL. Override for providers like Azure. + """ + return None + + def get_api_key(self) -> Optional[str]: + """ + Returns the API key. Override for providers like Azure. + """ + return None + def get_output_format_schema(self) -> Dict[str, Any]: """ Returns a simple JSON schema for testing structured outputs. @@ -71,12 +87,23 @@ class BaseAnthropicMessagesStructuredOutputTest(ABC): messages = self.get_test_messages() output_format = self.get_output_format_schema() - response = await litellm.anthropic.messages.acreate( - model=self.get_model(), - messages=messages, - max_tokens=100, - output_format=output_format, - ) + # Build kwargs with optional api_base and api_key + kwargs: Dict[str, Any] = { + "model": self.get_model(), + "messages": messages, + "max_tokens": 100, + "output_format": output_format, + } + + api_base = self.get_api_base() + if api_base: + kwargs["api_base"] = api_base + + api_key = self.get_api_key() + if api_key: + kwargs["api_key"] = api_key + + response = await litellm.anthropic.messages.acreate(**kwargs) print(f"Response: {response}") diff --git a/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py b/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py index 815e9a2e198..96233d715ca 100644 --- a/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py +++ b/tests/pass_through_unit_tests/messages_api_structured_output/test_azure_anthropic_structured_output.py @@ -9,6 +9,7 @@ Requires Azure AI credentials and model deployment. import os import sys +from typing import Optional sys.path.insert(0, os.path.abspath("../../../..")) @@ -26,4 +27,10 @@ class TestAzureAnthropicStructuredOutput(BaseAnthropicMessagesStructuredOutputTe """ def get_model(self) -> str: - return "azure_ai/claude-3-5-sonnet-20241022" \ No newline at end of file + return "azure_ai/claude-haiku-4-5" + + def get_api_base(self) -> Optional[str]: + return "https://krish-mh44t553-eastus2.services.ai.azure.com/" + + def get_api_key(self) -> Optional[str]: + return os.environ.get("AZURE_ANTHROPIC_API_KEY") \ No newline at end of file