项目文件夹

文件
2025-10-14 13:02:15 -07:00

566 行
21 KiB
Python

from unittest.mock import Mock, patch
import io
import pytest
from aisuite import Client
from aisuite.framework.message import TranscriptionResult
from aisuite.provider import ASRError
@pytest.fixture(scope="module")
def provider_configs():
return {
"openai": {"api_key": "test_openai_api_key"},
"aws": {
"aws_access_key": "test_aws_access_key",
"aws_secret_key": "test_aws_secret_key",
"aws_session_token": "test_aws_session_token",
"aws_region": "us-west-2",
},
"azure": {
"api_key": "azure-api-key",
"base_url": "https://model.ai.azure.com",
},
"groq": {
"api_key": "groq-api-key",
},
"mistral": {
"api_key": "mistral-api-key",
},
"google": {
"project_id": "test_google_project_id",
"region": "us-west4",
"application_credentials": "test_google_application_credentials",
},
"fireworks": {
"api_key": "fireworks-api-key",
},
"nebius": {
"api_key": "nebius-api-key",
},
"inception": {
"api_key": "inception-api-key",
},
"deepgram": {
"api_key": "deepgram-api-key",
},
}
@pytest.mark.parametrize(
argnames=("patch_target", "provider", "model"),
argvalues=[
(
"aisuite.providers.openai_provider.OpenaiProvider.chat_completions_create",
"openai",
"gpt-4o",
),
(
"aisuite.providers.mistral_provider.MistralProvider.chat_completions_create",
"mistral",
"mistral-model",
),
(
"aisuite.providers.groq_provider.GroqProvider.chat_completions_create",
"groq",
"groq-model",
),
(
"aisuite.providers.aws_provider.AwsProvider.chat_completions_create",
"aws",
"claude-v3",
),
(
"aisuite.providers.azure_provider.AzureProvider.chat_completions_create",
"azure",
"azure-model",
),
(
"aisuite.providers.anthropic_provider.AnthropicProvider.chat_completions_create",
"anthropic",
"anthropic-model",
),
(
"aisuite.providers.google_provider.GoogleProvider.chat_completions_create",
"google",
"google-model",
),
(
"aisuite.providers.fireworks_provider.FireworksProvider.chat_completions_create",
"fireworks",
"fireworks-model",
),
(
"aisuite.providers.nebius_provider.NebiusProvider.chat_completions_create",
"nebius",
"nebius-model",
),
(
"aisuite.providers.inception_provider.InceptionProvider.chat_completions_create",
"inception",
"mercury",
),
],
)
def test_client_chat_completions(
provider_configs: dict, patch_target: str, provider: str, model: str
):
expected_response = f"{patch_target}_{provider}_{model}"
with patch(patch_target) as mock_provider:
mock_provider.return_value = expected_response
client = Client()
client.configure(provider_configs)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Who won the world series in 2020?"},
]
model_str = f"{provider}:{model}"
model_response = client.chat.completions.create(model_str, messages=messages)
assert model_response == expected_response
def test_invalid_provider_in_client_config():
# Testing an invalid provider name in the configuration
invalid_provider_configs = {
"invalid_provider": {"api_key": "invalid_api_key"},
}
# With lazy loading, Client initialization should succeed
client = Client()
client.configure(invalid_provider_configs)
messages = [
{"role": "user", "content": "Hello"},
]
# Expect ValueError when actually trying to use the invalid provider
with pytest.raises(
ValueError,
match=r"Invalid provider key 'invalid_provider'. Supported providers: ",
):
client.chat.completions.create("invalid_provider:some-model", messages=messages)
def test_invalid_model_format_in_create(monkeypatch):
from aisuite.providers.openai_provider import OpenaiProvider
monkeypatch.setattr(
target=OpenaiProvider,
name="chat_completions_create",
value=Mock(),
)
# Valid provider configurations
provider_configs = {
"openai": {"api_key": "test_openai_api_key"},
}
# Initialize the client with valid provider
client = Client()
client.configure(provider_configs)
messages = [
{"role": "system", "content": "You are a helpful assistant."},
{"role": "user", "content": "Tell me a joke."},
]
# Invalid model format
invalid_model = "invalidmodel"
# Expect ValueError when calling create with invalid model format and verify message
with pytest.raises(
ValueError, match=r"Invalid model format. Expected 'provider:model'"
):
client.chat.completions.create(invalid_model, messages=messages)
class TestClientASR:
"""Test suite for Client ASR functionality - essential tests only."""
def test_audio_interface_initialization(self):
"""Test that Audio interface is properly initialized."""
client = Client()
assert hasattr(client, "audio")
assert hasattr(client.audio, "transcriptions")
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_transcriptions_create_success(
self, mock_create_provider, provider_configs
):
"""Test successful audio transcription with OpenAI."""
mock_result = TranscriptionResult(
text="Hello, this is a test transcription.",
language="en",
confidence=0.95,
task="transcribe",
)
# Create a mock provider with audio support
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
client = Client()
client.configure(provider_configs)
audio_data = io.BytesIO(b"fake audio data")
result = client.audio.transcriptions.create(
model="openai:whisper-1", file=audio_data, language="en"
)
assert isinstance(result, TranscriptionResult)
assert result.text == "Hello, this is a test transcription."
mock_provider.audio.transcriptions.create.assert_called_once()
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_transcriptions_create_deepgram(
self, mock_create_provider, provider_configs
):
"""Test audio transcription with Deepgram provider."""
mock_result = TranscriptionResult(
text="Deepgram transcription result.",
language="en",
confidence=0.92,
task="transcribe",
)
# Create a mock provider with audio support
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
client = Client()
client.configure(provider_configs)
result = client.audio.transcriptions.create(
model="deepgram:nova-2", file="test_audio.wav", language="en"
)
assert isinstance(result, TranscriptionResult)
assert result.text == "Deepgram transcription result."
mock_provider.audio.transcriptions.create.assert_called_once()
def test_transcriptions_invalid_model_format(self, provider_configs):
"""Test that invalid model format raises ValueError."""
client = Client()
client.configure(provider_configs)
with pytest.raises(ValueError, match="Invalid model format"):
client.audio.transcriptions.create(
model="invalid-format", file="test.wav", language="en"
)
def test_transcriptions_unsupported_provider(self, provider_configs):
"""Test error handling for unsupported ASR provider."""
client = Client()
client.configure(provider_configs)
with pytest.raises(ValueError, match="Invalid provider key"):
client.audio.transcriptions.create(
model="unsupported:model", file="test.wav", language="en"
)
class TestClientASRParameterValidation:
"""Test suite for Client-level ASR parameter validation."""
def test_client_initialization_strict_mode(self):
"""Test Client initialization with strict extra_param_mode."""
client = Client(extra_param_mode="strict")
assert client.extra_param_mode == "strict"
assert client.param_validator.extra_param_mode == "strict"
def test_client_initialization_warn_mode(self):
"""Test Client initialization with warn extra_param_mode (default)."""
client = Client()
assert client.extra_param_mode == "warn"
assert client.param_validator.extra_param_mode == "warn"
def test_client_initialization_permissive_mode(self):
"""Test Client initialization with permissive extra_param_mode."""
client = Client(extra_param_mode="permissive")
assert client.extra_param_mode == "permissive"
assert client.param_validator.extra_param_mode == "permissive"
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_strict_mode_rejects_unknown_param(self, mock_create_provider):
"""Test that strict mode raises ValueError for unknown parameters."""
client = Client(
provider_configs={"openai": {"api_key": "test"}}, extra_param_mode="strict"
)
# Mock provider shouldn't be called due to validation error
mock_provider = Mock()
mock_create_provider.return_value = mock_provider
with pytest.raises(ValueError, match="Unknown parameters for openai"):
client.audio.transcriptions.create(
model="openai:whisper-1",
file=io.BytesIO(b"audio"),
language="en",
invalid_param=True, # Unknown param
)
# Provider should not have been called (validation failed first)
mock_provider.audio.transcriptions.create.assert_not_called()
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_strict_mode_typo_detection(self, mock_create_provider):
"""Test that strict mode catches typos in parameter names."""
client = Client(
provider_configs={"openai": {"api_key": "test"}}, extra_param_mode="strict"
)
mock_provider = Mock()
mock_create_provider.return_value = mock_provider
with pytest.raises(
ValueError, match="Unknown parameters for openai: \\['langauge'\\]"
):
client.audio.transcriptions.create(
model="openai:whisper-1",
file=io.BytesIO(b"audio"),
langauge="en", # TYPO: should be "language"
)
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_warn_mode_continues_execution(self, mock_create_provider):
"""Test that warn mode continues execution after warning."""
import warnings
client = Client(
provider_configs={"openai": {"api_key": "test"}}, extra_param_mode="warn"
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
# Should warn but continue
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
result = client.audio.transcriptions.create(
model="openai:whisper-1",
file=io.BytesIO(b"audio"),
language="en",
invalid_param=True, # Unknown param
)
# Should have issued a warning
assert len(w) == 1
assert "Unknown parameters" in str(w[0].message)
# But execution should continue
assert result.text == "Test"
mock_provider.audio.transcriptions.create.assert_called_once()
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_permissive_mode_allows_unknown_params(self, mock_create_provider):
"""Test that permissive mode allows unknown parameters."""
import warnings
client = Client(
provider_configs={"openai": {"api_key": "test"}},
extra_param_mode="permissive",
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
# Should not warn or raise
with warnings.catch_warnings(record=True) as w:
warnings.simplefilter("always")
result = client.audio.transcriptions.create(
model="openai:whisper-1",
file=io.BytesIO(b"audio"),
experimental_feature=True, # Unknown param
)
# Should not have issued any warnings
assert len(w) == 0
# Execution should succeed
assert result.text == "Test"
mock_provider.audio.transcriptions.create.assert_called_once()
# Unknown param should be passed through
call_kwargs = mock_provider.audio.transcriptions.create.call_args.kwargs
assert call_kwargs.get("experimental_feature") is True
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_common_param_mapping_at_client_level(self, mock_create_provider):
"""Test that common parameters are mapped correctly at Client level."""
client = Client(
provider_configs={"google": {"project_id": "test", "region": "us"}},
extra_param_mode="strict",
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
# Use common param "language" which should map to "language_code" for Google
result = client.audio.transcriptions.create(
model="google:latest_long",
file=io.BytesIO(b"audio"),
language="en", # Common param
)
assert result.text == "Test"
mock_provider.audio.transcriptions.create.assert_called_once()
# Verify parameter was mapped to language_code
call_kwargs = mock_provider.audio.transcriptions.create.call_args.kwargs
assert "language_code" in call_kwargs
assert call_kwargs["language_code"] == "en-US" # Expanded
assert "language" not in call_kwargs # Original key should be mapped
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_provider_specific_params_passthrough(self, mock_create_provider):
"""Test that provider-specific parameters pass through correctly."""
client = Client(
provider_configs={"deepgram": {"api_key": "test"}},
extra_param_mode="strict",
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
result = client.audio.transcriptions.create(
model="deepgram:nova-2",
file=io.BytesIO(b"audio"),
punctuate=True,
diarize=True,
)
assert result.text == "Test"
# Verify provider-specific params passed through
call_kwargs = mock_provider.audio.transcriptions.create.call_args.kwargs
assert call_kwargs["punctuate"] is True
assert call_kwargs["diarize"] is True
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_mixed_common_and_provider_params(self, mock_create_provider):
"""Test mixing common and provider-specific parameters."""
client = Client(
provider_configs={"deepgram": {"api_key": "test"}},
extra_param_mode="strict",
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
result = client.audio.transcriptions.create(
model="deepgram:nova-2",
file=io.BytesIO(b"audio"),
language="en", # Common param
prompt="meeting", # Common param that maps to keywords
punctuate=True, # Deepgram-specific
diarize=True, # Deepgram-specific
)
assert result.text == "Test"
# Verify both common and provider params processed correctly
call_kwargs = mock_provider.audio.transcriptions.create.call_args.kwargs
assert call_kwargs["language"] == "en"
assert call_kwargs["keywords"] == ["meeting"] # prompt mapped to keywords
assert call_kwargs["punctuate"] is True
assert call_kwargs["diarize"] is True
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_validation_happens_before_provider_call(self, mock_create_provider):
"""Test that validation occurs before provider SDK is called."""
client = Client(
provider_configs={"openai": {"api_key": "test"}}, extra_param_mode="strict"
)
mock_provider = Mock()
mock_create_provider.return_value = mock_provider
# Validation should fail before provider is even initialized
with pytest.raises(ValueError, match="Unknown parameters"):
client.audio.transcriptions.create(
model="openai:whisper-1",
file=io.BytesIO(b"audio"),
completely_invalid_param=True,
)
# Provider create method should still have been called to initialize
# but the transcription method should never be called
mock_provider.audio.transcriptions.create.assert_not_called()
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_unsupported_common_param_ignored(self, mock_create_provider):
"""Test that unsupported common params are gracefully ignored."""
client = Client(
provider_configs={"deepgram": {"api_key": "test"}},
extra_param_mode="strict",
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
# temperature is not supported by Deepgram (should be ignored)
result = client.audio.transcriptions.create(
model="deepgram:nova-2",
file=io.BytesIO(b"audio"),
language="en",
temperature=0.5, # Not supported by Deepgram
)
assert result.text == "Test"
# Verify temperature was not passed to provider
call_kwargs = mock_provider.audio.transcriptions.create.call_args.kwargs
assert "temperature" not in call_kwargs
assert call_kwargs["language"] == "en"
@patch("aisuite.provider.ProviderFactory.create_provider")
def test_multiple_providers_with_same_client(self, mock_create_provider):
"""Test that the same client can handle multiple providers with different validation."""
client = Client(
provider_configs={
"openai": {"api_key": "test1"},
"deepgram": {"api_key": "test2"},
},
extra_param_mode="strict",
)
mock_result = TranscriptionResult(text="Test", language="en")
mock_provider = Mock()
mock_provider.audio.transcriptions.create.return_value = mock_result
mock_create_provider.return_value = mock_provider
# Test OpenAI with temperature (supported)
result1 = client.audio.transcriptions.create(
model="openai:whisper-1", file=io.BytesIO(b"audio"), temperature=0.5
)
assert result1.text == "Test"
call_kwargs1 = mock_provider.audio.transcriptions.create.call_args.kwargs
assert call_kwargs1.get("temperature") == 0.5
# Reset mock
mock_provider.reset_mock()
# Test Deepgram with temperature (not supported, should be ignored)
result2 = client.audio.transcriptions.create(
model="deepgram:nova-2", file=io.BytesIO(b"audio"), temperature=0.5
)
assert result2.text == "Test"
call_kwargs2 = mock_provider.audio.transcriptions.create.call_args.kwargs
assert "temperature" not in call_kwargs2