diff --git a/codewiki/src/be/llm_services.py b/codewiki/src/be/llm_services.py index df8f1aa7..8a2ff5f6 100644 --- a/codewiki/src/be/llm_services.py +++ b/codewiki/src/be/llm_services.py @@ -6,6 +6,7 @@ Supports multiple providers: openai-compatible, anthropic, bedrock, azure-openai. """ +import inspect import logging from typing import Optional @@ -148,10 +149,20 @@ def __init__(self, model_name, *, prompt_caching=True, cache_registry_key="", ** def _prompt_caching_active(self) -> bool: return self._prompt_caching_enabled and self._cache_registry_key not in _CACHE_UNSUPPORTED - async def _map_messages(self, messages, model_request_parameters, *, model_settings=None): - openai_messages = await super()._map_messages( - messages, model_request_parameters, model_settings=model_settings - ) + async def _map_messages(self, messages, model_request_parameters=None, *, model_settings=None): + base = OpenAIChatModel._map_messages + base_params = set(inspect.signature(base).parameters) + if "model_settings" in base_params: + openai_messages = await super()._map_messages( + messages, model_request_parameters, model_settings=model_settings + ) + elif "model_request_parameters" in base_params: + # Intermediate pydantic-ai versions accept request parameters but + # predate the keyword-only model_settings argument. + openai_messages = await super()._map_messages(messages, model_request_parameters) + else: + # Early pydantic-ai versions accept only messages. + openai_messages = await super()._map_messages(messages) if self._prompt_caching_active: # Breakpoint 1: last system/developer message (covers tools + system). for message in reversed(openai_messages): diff --git a/tests/test_prompt_caching.py b/tests/test_prompt_caching.py index 12bf6c7c..97368009 100644 --- a/tests/test_prompt_caching.py +++ b/tests/test_prompt_caching.py @@ -57,6 +57,50 @@ def _map(model: CachingOpenAIModel) -> list: ) +def test_map_messages_adapts_to_all_base_signatures(): + """Forward only the arguments supported by each pydantic-ai API era.""" + model = _make_model(prompt_caching=False) + history = _history() + request_parameters = ModelRequestParameters() + model_settings = {"temperature": 0} + calls = [] + + async def one_argument(self, messages): + calls.append(("one", messages)) + return [] + + async def two_arguments(self, messages, model_request_parameters): + calls.append(("two", messages, model_request_parameters)) + return [] + + async def three_arguments( + self, messages, model_request_parameters, *, model_settings=None + ): + calls.append(("three", messages, model_request_parameters, model_settings)) + return [] + + original = OpenAIChatModel._map_messages + try: + for base_method in (one_argument, two_arguments, three_arguments): + OpenAIChatModel._map_messages = base_method + result = asyncio.run( + model._map_messages( + history, + request_parameters, + model_settings=model_settings, + ) + ) + assert result == [] + finally: + OpenAIChatModel._map_messages = original + + assert calls == [ + ("one", history), + ("two", history, request_parameters), + ("three", history, request_parameters, model_settings), + ] + + def test_injects_breakpoints_on_system_and_final_message(): _CACHE_UNSUPPORTED.clear() messages = _map(_make_model())