mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor(responses): give the background row store an injectable seam
The store lived inline in responses_api, so covering it meant patching five proxy_server globals per test, which the test-quality gate counts as pinning the test to the wiring rather than the behaviour. It is now a module-level function that takes the managed-files hook as an argument, and the tests hand it a fake directly. The poller tests reach the same fetch through the router fixture that is already injected, instead of patching litellm.aget_responses. A response with no model_id now returns early rather than raising into the caller's except block. Nothing is stored either way and the warning is unchanged, so only the redundant second log line goes away. Claude-Session: https://claude.ai/code/session_01RHAjRxNhXTpKHeGMZ1nDKi
This commit is contained in:
parent
49a8ce5e42
commit
492336a50b
3 changed files with 147 additions and 151 deletions
|
|
@ -39,10 +39,47 @@ from litellm.types.responses.main import DeleteResponseResult
|
|||
from litellm.types.utils import TokenCountResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles
|
||||
|
||||
from litellm.router import Router
|
||||
|
||||
router: Final = APIRouter()
|
||||
|
||||
|
||||
async def store_background_response_object(
|
||||
response: ResponsesAPIResponse,
|
||||
managed_files_obj: "_PROXY_LiteLLMManagedFiles",
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> None:
|
||||
"""Record a queued background response so the cost poller can find and bill it.
|
||||
|
||||
``model_object_id`` carries the provider's own id because the advertised ``response.id``
|
||||
is re-encrypted with a fresh nonce on every call, leaving the row no stable handle on
|
||||
the generation it describes.
|
||||
"""
|
||||
from litellm.proxy.hooks.responses_id_security import ResponsesIDSecurity
|
||||
|
||||
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
|
||||
if not hidden_params.get("model_id"):
|
||||
verbose_proxy_logger.warning(
|
||||
"No model_id found in response hidden params for response %s, skipping managed object storage",
|
||||
response.id,
|
||||
)
|
||||
return
|
||||
|
||||
provider_response_id, _, _ = ResponsesIDSecurity()._decrypt_response_id(response.id)
|
||||
await managed_files_obj.store_unified_object_id(
|
||||
unified_object_id=response.id,
|
||||
file_object=response,
|
||||
litellm_parent_otel_span=None,
|
||||
model_object_id=provider_response_id,
|
||||
file_purpose="response",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
persist_attribution=True,
|
||||
)
|
||||
verbose_proxy_logger.info("Stored background response %s in managed objects table", response.id)
|
||||
|
||||
|
||||
_user_api_key_auth_dep: Final = Depends(user_api_key_auth)
|
||||
_RESPONSES_TAGS: Final[list[str | Enum]] = ["responses"] # mutable-ok: fastapi's route signature requires list tags
|
||||
|
||||
|
|
@ -366,54 +403,29 @@ async def responses_api(
|
|||
)
|
||||
|
||||
# Store in managed objects table if background mode is enabled
|
||||
if data.get("background") and isinstance(response, ResponsesAPIResponse):
|
||||
if response.status in ["queued", "in_progress"]:
|
||||
from litellm_enterprise.proxy.hooks.managed_files import (
|
||||
_PROXY_LiteLLMManagedFiles,
|
||||
)
|
||||
if (
|
||||
data.get("background")
|
||||
and isinstance(response, ResponsesAPIResponse)
|
||||
and response.status in ("queued", "in_progress")
|
||||
):
|
||||
from litellm_enterprise.proxy.hooks.managed_files import (
|
||||
_PROXY_LiteLLMManagedFiles,
|
||||
)
|
||||
|
||||
managed_files_obj: Final = cast(
|
||||
_PROXY_LiteLLMManagedFiles | None,
|
||||
proxy_logging_obj.get_proxy_hook("managed_files"),
|
||||
)
|
||||
managed_files_obj: Final = cast(
|
||||
_PROXY_LiteLLMManagedFiles | None,
|
||||
proxy_logging_obj.get_proxy_hook("managed_files"),
|
||||
)
|
||||
|
||||
if managed_files_obj and llm_router:
|
||||
try:
|
||||
from litellm.proxy.hooks.responses_id_security import (
|
||||
ResponsesIDSecurity,
|
||||
)
|
||||
|
||||
# Get the actual deployment model_id from hidden params
|
||||
hidden_params: Final = getattr(response, "_hidden_params", {}) or {}
|
||||
model_id: Final = hidden_params.get("model_id", None)
|
||||
|
||||
if not model_id:
|
||||
verbose_proxy_logger.warning(
|
||||
"No model_id found in response hidden params for response %s, skipping managed object storage",
|
||||
response.id,
|
||||
)
|
||||
raise Exception("No model_id found in response hidden params")
|
||||
provider_response_id, _, _ = ResponsesIDSecurity()._decrypt_response_id(response.id)
|
||||
# Store in managed objects table
|
||||
await managed_files_obj.store_unified_object_id(
|
||||
unified_object_id=response.id,
|
||||
file_object=response,
|
||||
litellm_parent_otel_span=None,
|
||||
model_object_id=provider_response_id,
|
||||
file_purpose="response",
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
persist_attribution=True,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.info(
|
||||
"Stored background response %s in managed objects table with unified_id=%s",
|
||||
response.id,
|
||||
response.id,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(
|
||||
"Failed to store background response in managed objects table: %s", e
|
||||
)
|
||||
if managed_files_obj and llm_router:
|
||||
try:
|
||||
await store_background_response_object(
|
||||
response=response,
|
||||
managed_files_obj=managed_files_obj,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error("Failed to store background response in managed objects table: %s", e)
|
||||
|
||||
return response
|
||||
except ModifyResponseException as e:
|
||||
|
|
|
|||
|
|
@ -1191,18 +1191,23 @@ class TestCheckResponsesCost:
|
|||
"""The row's provider id drives the fetch, not the nonce-encrypted advertised id.
|
||||
|
||||
A background create advertises a freshly encrypted id per call, so unified_object_id
|
||||
is not a stable handle on the generation.
|
||||
is no handle on the generation.
|
||||
"""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids")
|
||||
|
||||
provider_response_id = "resp_provider_stable_1"
|
||||
provider_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai",
|
||||
model_id="deployment-xyz",
|
||||
response_id="resp_upstream_stable",
|
||||
)
|
||||
stale_advertised_id = "resp_" + str(
|
||||
encrypt_value_helper(
|
||||
value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
|
||||
"resp_some_other_encoding", "test-user", "test-team"
|
||||
"resp_a_previous_encoding", "test-user", "test-team"
|
||||
)
|
||||
)
|
||||
)
|
||||
|
|
@ -1216,21 +1221,20 @@ class TestCheckResponsesCost:
|
|||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
|
||||
|
||||
mock_response = ResponsesAPIResponse(
|
||||
id=provider_response_id,
|
||||
object="response",
|
||||
status="completed",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
|
||||
mock_llm_router.aget_responses = AsyncMock(
|
||||
return_value=ResponsesAPIResponse(
|
||||
id=provider_response_id,
|
||||
object="response",
|
||||
status="completed",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
|
||||
)
|
||||
)
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget:
|
||||
mock_aget.return_value = mock_response
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
assert mock_aget.call_args[1]["response_id"] == provider_response_id
|
||||
assert mock_llm_router.aget_responses.call_args[1]["response_id"] == provider_response_id
|
||||
assert _completed_job_ids(mock_prisma_client) == ["job-provider-id"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -1239,11 +1243,16 @@ class TestCheckResponsesCost:
|
|||
):
|
||||
"""Rows created earlier carry the encrypted advertised id in both columns."""
|
||||
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids")
|
||||
|
||||
provider_response_id = "resp_legacy_upstream_9"
|
||||
provider_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai",
|
||||
model_id="deployment-xyz",
|
||||
response_id="resp_legacy_upstream",
|
||||
)
|
||||
legacy_id = "resp_" + str(
|
||||
encrypt_value_helper(
|
||||
value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
|
||||
|
|
@ -1261,19 +1270,18 @@ class TestCheckResponsesCost:
|
|||
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(return_value=[mock_job])
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1)
|
||||
|
||||
mock_response = ResponsesAPIResponse(
|
||||
id=provider_response_id,
|
||||
object="response",
|
||||
status="completed",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
|
||||
mock_llm_router.aget_responses = AsyncMock(
|
||||
return_value=ResponsesAPIResponse(
|
||||
id=provider_response_id,
|
||||
object="response",
|
||||
status="completed",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=10, output_tokens=5, total_tokens=15),
|
||||
)
|
||||
)
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget:
|
||||
mock_aget.return_value = mock_response
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
assert mock_aget.call_args[1]["response_id"] == provider_response_id
|
||||
assert mock_llm_router.aget_responses.call_args[1]["response_id"] == provider_response_id
|
||||
assert _completed_job_ids(mock_prisma_client) == ["job-legacy"]
|
||||
|
|
|
|||
|
|
@ -1962,11 +1962,11 @@ class TestResponsesInputTokens:
|
|||
|
||||
|
||||
class TestBackgroundResponseManagedObjectId:
|
||||
"""The managed row for a background response must be keyed by the provider's own id.
|
||||
"""The managed row for a background response is keyed by the provider's own id.
|
||||
|
||||
The advertised ``response.id`` is encrypted with a fresh nonce per call, so storing it
|
||||
in ``model_object_id`` leaves the row with no stable lookup key and every later read
|
||||
of the same generation looks like a new object.
|
||||
in ``model_object_id`` leaves the row with no stable handle on the generation and every
|
||||
later read of the same generation looks like a new object.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
|
|
@ -1979,16 +1979,10 @@ class TestBackgroundResponseManagedObjectId:
|
|||
)
|
||||
return f"resp_{encrypt_value_helper(value=managed_id)}"
|
||||
|
||||
async def _store_call_for(self, provider_response_id: str) -> dict:
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.response_api_endpoints.endpoints import responses_api
|
||||
@staticmethod
|
||||
def _queued_response(advertised_id: str, model_id: str | None = "deployment-1"):
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
advertised_id = self._encrypted_id(provider_response_id)
|
||||
assert advertised_id != self._encrypted_id(provider_response_id), (
|
||||
"advertised ids must be nonce-encrypted, otherwise this regression cannot occur"
|
||||
)
|
||||
|
||||
response = ResponsesAPIResponse(
|
||||
id=advertised_id,
|
||||
created_at=0,
|
||||
|
|
@ -2000,50 +1994,60 @@ class TestBackgroundResponseManagedObjectId:
|
|||
tools=[],
|
||||
status="queued",
|
||||
)
|
||||
response._hidden_params = {"model_id": "deployment-1"}
|
||||
response._hidden_params = {"model_id": model_id} if model_id else {}
|
||||
return response
|
||||
|
||||
async def _stored_kwargs(self, advertised_id: str, model_id: str | None = "deployment-1"):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.response_api_endpoints.endpoints import (
|
||||
store_background_response_object,
|
||||
)
|
||||
|
||||
managed_files_obj = MagicMock()
|
||||
managed_files_obj.store_unified_object_id = AsyncMock()
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.get_proxy_hook = MagicMock(return_value=managed_files_obj)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server._read_request_body",
|
||||
AsyncMock(return_value={"model": "gpt-4o", "input": "hi", "background": True}),
|
||||
), patch("litellm.proxy.proxy_server.polling_via_cache_enabled", False), patch(
|
||||
"litellm.proxy.proxy_server.llm_router", MagicMock()
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
|
||||
), patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_process_llm_request",
|
||||
AsyncMock(return_value=response),
|
||||
):
|
||||
await responses_api(
|
||||
request=MagicMock(),
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"),
|
||||
)
|
||||
|
||||
managed_files_obj.store_unified_object_id.assert_awaited_once()
|
||||
return managed_files_obj.store_unified_object_id.await_args.kwargs
|
||||
await store_background_response_object(
|
||||
response=self._queued_response(advertised_id, model_id),
|
||||
managed_files_obj=managed_files_obj,
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"),
|
||||
)
|
||||
return managed_files_obj.store_unified_object_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_object_id_is_the_provider_response_id(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt")
|
||||
provider_response_id = "resp_provider68abc123"
|
||||
advertised_id = self._encrypted_id(provider_response_id)
|
||||
assert advertised_id != self._encrypted_id(provider_response_id), (
|
||||
"advertised ids must be nonce-encrypted, otherwise this regression cannot occur"
|
||||
)
|
||||
|
||||
kwargs = await self._store_call_for(provider_response_id)
|
||||
store = await self._stored_kwargs(advertised_id)
|
||||
|
||||
store.assert_awaited_once()
|
||||
kwargs = store.await_args.kwargs
|
||||
assert kwargs["model_object_id"] == provider_response_id
|
||||
assert kwargs["unified_object_id"] != provider_response_id
|
||||
assert kwargs["unified_object_id"] == kwargs["file_object"].id
|
||||
assert kwargs["unified_object_id"] == advertised_id
|
||||
assert kwargs["file_object"].id == advertised_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_two_background_creates_are_distinguishable_by_provider_id(self, monkeypatch):
|
||||
async def test_two_creates_of_one_generation_share_a_provider_id(self, monkeypatch):
|
||||
"""Re-encrypting the same generation must not look like a second object."""
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt")
|
||||
provider_response_id = "resp_provider_same_gen"
|
||||
|
||||
first = (await self._stored_kwargs(self._encrypted_id(provider_response_id))).await_args.kwargs
|
||||
second = (await self._stored_kwargs(self._encrypted_id(provider_response_id))).await_args.kwargs
|
||||
|
||||
assert first["unified_object_id"] != second["unified_object_id"]
|
||||
assert first["model_object_id"] == second["model_object_id"] == provider_response_id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_distinct_generations_keep_distinct_provider_ids(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt")
|
||||
|
||||
first = await self._store_call_for("resp_providerAAA")
|
||||
second = await self._store_call_for("resp_providerBBB")
|
||||
first = (await self._stored_kwargs(self._encrypted_id("resp_providerAAA"))).await_args.kwargs
|
||||
second = (await self._stored_kwargs(self._encrypted_id("resp_providerBBB"))).await_args.kwargs
|
||||
|
||||
assert first["model_object_id"] == "resp_providerAAA"
|
||||
assert second["model_object_id"] == "resp_providerBBB"
|
||||
|
|
@ -2051,45 +2055,17 @@ class TestBackgroundResponseManagedObjectId:
|
|||
@pytest.mark.asyncio
|
||||
async def test_unencrypted_advertised_id_is_stored_as_is(self, monkeypatch):
|
||||
"""With response-id security disabled the advertised id is already the provider's."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.response_api_endpoints.endpoints import responses_api
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt")
|
||||
response = ResponsesAPIResponse(
|
||||
id="resp_rawprovider999",
|
||||
created_at=0,
|
||||
model="gpt-4o",
|
||||
object="response",
|
||||
output=[],
|
||||
parallel_tool_calls=False,
|
||||
tool_choice="auto",
|
||||
tools=[],
|
||||
status="queued",
|
||||
)
|
||||
response._hidden_params = {"model_id": "deployment-1"}
|
||||
|
||||
managed_files_obj = MagicMock()
|
||||
managed_files_obj.store_unified_object_id = AsyncMock()
|
||||
proxy_logging_obj = MagicMock()
|
||||
proxy_logging_obj.get_proxy_hook = MagicMock(return_value=managed_files_obj)
|
||||
store = await self._stored_kwargs("resp_rawprovider999")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server._read_request_body",
|
||||
AsyncMock(return_value={"model": "gpt-4o", "input": "hi", "background": True}),
|
||||
), patch("litellm.proxy.proxy_server.polling_via_cache_enabled", False), patch(
|
||||
"litellm.proxy.proxy_server.llm_router", MagicMock()
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.proxy_logging_obj", proxy_logging_obj
|
||||
), patch(
|
||||
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing.base_process_llm_request",
|
||||
AsyncMock(return_value=response),
|
||||
):
|
||||
await responses_api(
|
||||
request=MagicMock(),
|
||||
fastapi_response=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-1234", user_id="u-1", team_id="t-1"),
|
||||
)
|
||||
assert store.await_args.kwargs["model_object_id"] == "resp_rawprovider999"
|
||||
|
||||
kwargs = managed_files_obj.store_unified_object_id.await_args.kwargs
|
||||
assert kwargs["model_object_id"] == "resp_rawprovider999"
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_without_a_deployment_is_not_stored(self, monkeypatch):
|
||||
"""No model_id means the poller could never route the read, so no row is written."""
|
||||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-regression-salt")
|
||||
|
||||
store = await self._stored_kwargs(self._encrypted_id("resp_no_deployment"), model_id=None)
|
||||
|
||||
store.assert_not_awaited()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue