fix: prevent double prompt management in async path, preserve optional params

- aresponses() now pops prompt_id from kwargs after the async hook runs
  and passes merged_optional_params via _async_prompt_merged_params.
  responses() checks for this internal kwarg first and skips the sync
  hook entirely when present — eliminating double-merge of template
  messages.
- merged_optional_params from async_get_chat_completion_prompt is no
  longer discarded (_); it flows through to local_vars in responses().
- Async tests now assert get_chat_completion_prompt.assert_not_called()
  to directly detect any double-execution regression.

Made-with: Cursor
This commit is contained in:
Sameer Kankute
2026-03-18 17:09:28 +05:30
parent d333dc4077
commit 021540b2e2
2 changed files with 69 additions and 50 deletions
+57 -40
View File
@@ -466,6 +466,11 @@ async def aresponses(
#########################################################
# ASYNC PROMPT MANAGEMENT
# Run the async hook here so async-only prompt loggers are honoured.
# Then pop prompt_id from kwargs so the sync responses() path does NOT
# re-run the hook (which would double-prepend template messages).
# Pass merged_optional_params via an internal kwarg so responses()
# can apply them to local_vars without re-invoking the hook.
#########################################################
litellm_logging_obj = kwargs.get("litellm_logging_obj", None)
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
@@ -488,7 +493,7 @@ async def aresponses(
(
model,
merged_input,
_,
merged_optional_params,
) = await litellm_logging_obj.async_get_chat_completion_prompt(
model=model,
messages=client_input,
@@ -503,6 +508,8 @@ async def aresponses(
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model
)
kwargs.pop("prompt_id", None)
kwargs["_async_prompt_merged_params"] = merged_optional_params
func = partial(
responses,
@@ -666,47 +673,57 @@ def responses(
#########################################################
# PROMPT MANAGEMENT
# If aresponses() already ran the async hook, it pops prompt_id and
# passes the result via _async_prompt_merged_params — apply those
# directly and skip the sync hook to avoid double-merging.
#########################################################
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None))
original_model = model
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=kwargs
):
if isinstance(input, str):
client_input: List[AllMessageValues] = [
{"role": "user", "content": input}
]
else:
client_input = [
item # type: ignore[misc]
for item in input
if isinstance(item, dict) and "role" in item
]
(
model,
merged_input,
merged_optional_params,
) = litellm_logging_obj.get_chat_completion_prompt(
model=model,
messages=client_input,
non_default_params=kwargs,
prompt_id=prompt_id,
prompt_variables=prompt_variables,
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
input = cast(Union[str, ResponseInputParam], merged_input)
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model
)
local_vars["custom_llm_provider"] = custom_llm_provider
for k, v in merged_optional_params.items():
_async_merged = kwargs.pop("_async_prompt_merged_params", None)
if _async_merged is not None:
for k, v in _async_merged.items():
local_vars[k] = v
else:
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
prompt_variables = cast(
Optional[dict], kwargs.get("prompt_variables", None)
)
original_model = model
if isinstance(litellm_logging_obj, LiteLLMLoggingObj) and litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=kwargs
):
if isinstance(input, str):
client_input: List[AllMessageValues] = [
{"role": "user", "content": input}
]
else:
client_input = [
item # type: ignore[misc]
for item in input
if isinstance(item, dict) and "role" in item
]
(
model,
merged_input,
merged_optional_params,
) = litellm_logging_obj.get_chat_completion_prompt(
model=model,
messages=client_input,
non_default_params=kwargs,
prompt_id=prompt_id,
prompt_variables=prompt_variables,
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
input = cast(Union[str, ResponseInputParam], merged_input)
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model
)
local_vars["custom_llm_provider"] = custom_llm_provider
for k, v in merged_optional_params.items():
local_vars[k] = v
#########################################################
# Update input and tools with provider-specific file IDs if managed files are used
@@ -288,16 +288,16 @@ class TestResponsesAPIPromptManagement:
class TestAsyncResponsesAPIPromptManagement:
"""Tests for the async aresponses() prompt management path.
aresponses() calls async_get_chat_completion_prompt at the outer async level
(for async-only prompt loggers), then delegates to responses() via
run_in_executor where the sync hook also runs mirroring acompletion() in
main.py. Optional params are handled by the sync responses() path.
aresponses() calls async_get_chat_completion_prompt at the outer async
level, then pops prompt_id from kwargs and passes merged_optional_params
via an internal kwarg. The sync responses() path sees no prompt_id and
skips the sync hook entirely preventing double-merge of template messages.
"""
@pytest.mark.asyncio
async def test_async_calls_async_hook(self):
"""[H] aresponses() invokes async_get_chat_completion_prompt before
dispatching to the sync responses() path."""
async def test_async_calls_async_hook_not_sync(self):
"""[H] aresponses() invokes async_get_chat_completion_prompt and the
sync get_chat_completion_prompt is NOT called (no double-merge)."""
template_messages: List[AllMessageValues] = [
{"role": "system", "content": "You are helpful."}, # type: ignore[list-item]
]
@@ -318,14 +318,14 @@ class TestAsyncResponsesAPIPromptManagement:
)
logging_obj.async_get_chat_completion_prompt.assert_called_once()
logging_obj.get_chat_completion_prompt.assert_not_called()
call_kwargs = logging_obj.async_get_chat_completion_prompt.call_args.kwargs
assert call_kwargs["prompt_id"] == "async-test"
@pytest.mark.asyncio
async def test_async_optional_params_propagated(self):
"""[I] Template-defined optional params (e.g. temperature) reach the downstream
handler when called via aresponses(). The sync responses() path applies them
via local_vars."""
"""[I] Template-defined optional params (e.g. temperature) from the async
hook reach the downstream handler they are NOT silently discarded."""
template_messages: List[AllMessageValues] = [
{"role": "user", "content": "Hello"}, # type: ignore[list-item]
]
@@ -345,6 +345,7 @@ class TestAsyncResponsesAPIPromptManagement:
litellm_logging_obj=logging_obj,
)
logging_obj.get_chat_completion_prompt.assert_not_called()
handler_call_kwargs = mock_handler.call_args.kwargs
request_params = handler_call_kwargs.get("responses_api_request", {})
assert request_params.get("temperature") == 0.7
@@ -375,6 +376,7 @@ class TestAsyncResponsesAPIPromptManagement:
)
logging_obj.async_get_chat_completion_prompt.assert_called_once()
logging_obj.get_chat_completion_prompt.assert_not_called()
call_kwargs = logging_obj.async_get_chat_completion_prompt.call_args.kwargs
passed_messages = call_kwargs["messages"]
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)