Files
audita/python/tests/test_framework_llm.py

256 lines
9.0 KiB
Python

from concurrent.futures import ThreadPoolExecutor
import sys
import types
import pytest
from audita.core.config import AuditaConfig
from audita.core.errors import AuditaLLMError
from audita.framework.llm import OpenAICompatibleStructuredLLMClient
class DummyResponseModel:
pass
def _config(**overrides):
base = AuditaConfig.from_sources(env={})
data = {
"api_key": "test-key",
"llm_concurrency": base.llm_concurrency,
"llm_timeout_seconds": base.llm_timeout_seconds,
"validation_llm_api_key": base.validation_llm_api_key,
"validation_llm_concurrency": base.validation_llm_concurrency,
"validation_llm_timeout_seconds": base.validation_llm_timeout_seconds,
"validation_model": base.validation_model,
"validation_base_url": base.validation_base_url,
"validation_max_retries": base.validation_max_retries,
"module_keys": base.module_keys,
"model": base.model,
"base_url": base.base_url,
"max_retries": base.max_retries,
"max_section_tokens": base.max_section_tokens,
"glossary_confidence_threshold": base.glossary_confidence_threshold,
"grammar_confidence_threshold": base.grammar_confidence_threshold,
"homophones_confidence_threshold": base.homophones_confidence_threshold,
"spoken_word_confidence_threshold": base.spoken_word_confidence_threshold,
"normalize_max_segment_gap": base.normalize_max_segment_gap,
"normalize_ellipsis_gap": base.normalize_ellipsis_gap,
"normalize_max_segment_duration": base.normalize_max_segment_duration,
"normalize_max_segment_tokens": base.normalize_max_segment_tokens,
"work_dir": base.work_dir,
"work_dir_retention": base.work_dir_retention,
}
data.update(overrides)
return AuditaConfig(**data)
def _install_fake_llm_modules(monkeypatch):
create_calls = []
openai_inits = []
class FakePatchedClient:
def __init__(self):
self.chat = types.SimpleNamespace(completions=types.SimpleNamespace(create=self._create))
def _create(self, **kwargs):
create_calls.append(kwargs)
return {"ok": True}
class FakeOpenAI:
def __init__(self, *, api_key, base_url, timeout):
openai_inits.append({"api_key": api_key, "base_url": base_url, "timeout": timeout})
fake_instructor = types.SimpleNamespace(
Mode=types.SimpleNamespace(TOOLS="TOOLS"),
patch=lambda client, mode: FakePatchedClient(),
)
fake_openai = types.SimpleNamespace(OpenAI=FakeOpenAI)
monkeypatch.setitem(sys.modules, "instructor", fake_instructor)
monkeypatch.setitem(sys.modules, "openai", fake_openai)
return create_calls, openai_inits
def test_openrouter_requests_strip_prefix_and_include_extra_body(monkeypatch):
create_calls, _ = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
client.run_structured(
stage_name="test-stage",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=_config(model="openrouter/google/gemma-4-31b-it"),
)
assert create_calls == [
{
"model": "google/gemma-4-31b-it",
"messages": [{"role": "user", "content": "Hello"}],
"response_model": DummyResponseModel,
"max_retries": 3,
"extra_body": {"provider": {"require_parameters": True}},
}
]
def test_openrouter_default_base_url_uses_openrouter_request_shape(monkeypatch):
create_calls, _ = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
client.run_structured(
stage_name="test-stage",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=_config(model="google/gemma-4-31b-it"),
)
assert create_calls == [
{
"model": "google/gemma-4-31b-it",
"messages": [{"role": "user", "content": "Hello"}],
"response_model": DummyResponseModel,
"max_retries": 3,
"extra_body": {"provider": {"require_parameters": True}},
}
]
def test_generic_endpoint_requests_keep_model_and_omit_extra_body(monkeypatch):
create_calls, _ = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
client.run_structured(
stage_name="test-stage",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=_config(
model="meta-llama/Llama-3.1-8B-Instruct",
base_url="http://localhost:8000/v1",
),
)
assert create_calls == [
{
"model": "meta-llama/Llama-3.1-8B-Instruct",
"messages": [{"role": "user", "content": "Hello"}],
"response_model": DummyResponseModel,
"max_retries": 3,
}
]
def test_client_cache_identity_uses_api_key_and_base_url(monkeypatch):
_, openai_inits = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
first = _config(api_key="key-1", base_url="http://localhost:8000/v1")
second = _config(api_key="key-1", base_url="http://localhost:8000/v1")
third = _config(api_key="key-1", base_url="https://api.openai.com/v1")
client.run_structured(
stage_name="one",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=first,
)
client.run_structured(
stage_name="two",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=second,
)
client.run_structured(
stage_name="three",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=third,
)
assert openai_inits == [
{"api_key": "key-1", "base_url": "http://localhost:8000/v1", "timeout": 600},
{"api_key": "key-1", "base_url": "https://api.openai.com/v1", "timeout": 600},
]
def test_client_cache_identity_uses_timeout(monkeypatch):
_, openai_inits = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
first = _config(api_key="key-1", base_url="http://localhost:8000/v1", llm_timeout_seconds=600)
second = _config(api_key="key-1", base_url="http://localhost:8000/v1", llm_timeout_seconds=1200)
client.run_structured(
stage_name="one",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=first,
)
client.run_structured(
stage_name="two",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=second,
)
assert openai_inits == [
{"api_key": "key-1", "base_url": "http://localhost:8000/v1", "timeout": 600},
{"api_key": "key-1", "base_url": "http://localhost:8000/v1", "timeout": 1200},
]
def test_missing_api_key_error_is_provider_neutral():
client = OpenAICompatibleStructuredLLMClient()
with pytest.raises(AuditaLLMError, match="OpenRouter endpoint"):
client.run_structured(
stage_name="test-stage",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=_config(api_key=None),
)
def test_missing_api_key_is_allowed_for_nondefault_endpoint(monkeypatch):
create_calls, openai_inits = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
client.run_structured(
stage_name="test-stage",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=_config(
api_key=None,
model="meta-llama/Llama-3.1-8B-Instruct",
base_url="http://localhost:8000/v1",
),
)
assert openai_inits == [{"api_key": "audita-no-key-required", "base_url": "http://localhost:8000/v1", "timeout": 600}]
assert create_calls == [
{
"model": "meta-llama/Llama-3.1-8B-Instruct",
"messages": [{"role": "user", "content": "Hello"}],
"response_model": DummyResponseModel,
"max_retries": 3,
}
]
def test_client_initialization_is_safe_under_concurrent_calls(monkeypatch):
_, openai_inits = _install_fake_llm_modules(monkeypatch)
client = OpenAICompatibleStructuredLLMClient()
config = _config(api_key="key-1", base_url="http://localhost:8000/v1")
with ThreadPoolExecutor(max_workers=4) as executor:
list(
executor.map(
lambda _: client.run_structured(
stage_name="test-stage",
messages=[{"role": "user", "content": "Hello"}],
response_model=DummyResponseModel,
config=config,
),
range(4),
)
)
assert openai_inits == [{"api_key": "key-1", "base_url": "http://localhost:8000/v1", "timeout": 600}]