diff --git a/litellm/llms/nadir/chat/transformation.py b/litellm/llms/nadir/chat/transformation.py index d22b1a51bab..f3e8f70212c 100644 --- a/litellm/llms/nadir/chat/transformation.py +++ b/litellm/llms/nadir/chat/transformation.py @@ -1,7 +1,9 @@ import math +from collections.abc import Mapping from typing import Final import httpx +from pydantic import TypeAdapter, ValidationError from litellm.litellm_core_utils.core_helpers import set_response_cost_in_hidden_params from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig @@ -13,19 +15,39 @@ _SUPPORTED_OPENAI_PARAMS: Final = ( "frequency_penalty", "max_retries", "max_tokens", + "parallel_tool_calls", "presence_penalty", "response_format", + "service_tier", "stream", "temperature", + "tool_choice", + "tools", "top_p", + "user", ) +_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object]) + + +def _pricing_failed(reported: Mapping[str, object]) -> bool: + try: + return bool(_JSON_OBJECT.validate_python(reported.get("cost_breakdown")).get("pricing_failed")) + except ValidationError: + return False + + def _reported_cost_usd(raw_response: httpx.Response) -> float | None: try: - cost: Final = raw_response.json()["nadir_metadata"]["cost"]["total_cost_usd"] + reported: Final = _JSON_OBJECT.validate_python(raw_response.json()["nadir_metadata"]["cost"]) except (ValueError, KeyError, TypeError): return None + # An unpriced call still carries a total (0.0, or the sum of the attempts Nadir could price), + # flagged in its breakdown. Logging that total would record a paid call as free or cheaper than it was. + if _pricing_failed(reported): + return None + cost: Final = reported.get("total_cost_usd") if isinstance(cost, bool) or not isinstance(cost, (int, float)): return None if not math.isfinite(cost) or cost < 0: diff --git a/tests/unit/llms/nadir/test_nadir.py b/tests/unit/llms/nadir/test_nadir.py index 2b8b8387a42..2a8f9ef9b74 100644 --- a/tests/unit/llms/nadir/test_nadir.py +++ b/tests/unit/llms/nadir/test_nadir.py @@ -7,6 +7,7 @@ import pytest import litellm from litellm import get_llm_provider +from litellm.llms.custom_httpx.http_handler import HTTPHandler from litellm.types.utils import ModelResponse, Usage NADIR_BASE = "https://api.getnadir.com/v1" @@ -45,6 +46,86 @@ def _payload(**extra): } +_TOOLS = [ + { + "type": "function", + "function": { + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + }, + } +] +_TOOL_CALL_CHOICE = { + "index": 0, + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{"city": "Paris"}'}} + ], + }, + "finish_reason": "tool_calls", +} + + +class _ToolCallingNadir(HTTPHandler): + """Injected transport: records the body LiteLLM puts on the wire and answers with a tool call.""" + + def __init__(self) -> None: + super().__init__() + self.request_body: dict = {} + + def post(self, url: str, headers=None, data=None, **kwargs) -> httpx.Response: + self.request_body = json.loads(data) + return httpx.Response(200, json=_payload(choices=[_TOOL_CALL_CHOICE]), request=httpx.Request("POST", url)) + + +def _chunk(delta, finish_reason=None): + return { + "id": "req-1", + "object": "chat.completion.chunk", + "created": 0, + "model": "claude-haiku-4-5", + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], + } + + +# The frames Nadir streams for a tool-calling turn: the call's id and name first, then its arguments in +# fragments, then a finish frame that ends the turn with "tool_calls". +_TOOL_CALL_STREAM = ( + "".join( + f"data: {json.dumps(frame)}\n\n" + for frame in ( + _chunk({"role": "assistant", "content": ""}), + _chunk( + { + "tool_calls": [ + { + "index": 0, + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": ""}, + } + ] + } + ), + _chunk({"tool_calls": [{"index": 0, "function": {"arguments": '{"city": '}}]}), + _chunk({"tool_calls": [{"index": 0, "function": {"arguments": '"Paris"}'}}]}), + _chunk({}, "tool_calls"), + ) + ) + + "data: [DONE]\n\n" +).encode() + + +class _StreamingToolCallingNadir(_ToolCallingNadir): + """Injected transport that streams the tool call back as Nadir's SSE frames.""" + + def post(self, url: str, headers=None, data=None, **kwargs) -> httpx.Response: + self.request_body = json.loads(data) + return httpx.Response(200, content=_TOOL_CALL_STREAM, request=httpx.Request("POST", url)) + + def _cost(response, provider): return litellm.completion_cost(completion_response=response, custom_llm_provider=provider) @@ -131,15 +212,14 @@ class TestNadirParamMapping: assert params["temperature"] == 0.5 assert params["max_tokens"] == 64 - def test_streaming_is_advertised_and_tools_are_not(self): - params = litellm.get_supported_openai_params(model="auto", custom_llm_provider="nadir") - assert "stream" in params - assert "tools" not in params + @pytest.mark.parametrize("param", ["stream", "tools", "tool_choice", "parallel_tool_calls", "service_tier", "user"]) + def test_param_nadir_honours_is_advertised(self, param): + assert param in litellm.get_supported_openai_params(model="auto", custom_llm_provider="nadir") @pytest.mark.parametrize( "unsupported", [ - {"tools": [{"type": "function", "function": {"name": "f", "parameters": {}}}]}, + {"functions": [{"name": "f", "parameters": {}}]}, {"stop": ["\n"]}, {"seed": 7}, {"n": 2}, @@ -161,6 +241,50 @@ class TestNadirParamMapping: assert params["temperature"] == 0.2 +class TestNadirToolCalling: + @pytest.mark.parametrize("drop_params", [False, True]) + def test_tool_call_round_trip(self, drop_params): + client = _ToolCallingNadir() + response = litellm.completion( + model="nadir/auto", + messages=[{"role": "user", "content": "Weather in Paris?"}], + api_key="sk-test", + client=client, + drop_params=drop_params, + tools=_TOOLS, + tool_choice="auto", + parallel_tool_calls=False, + user="end-user-7", + service_tier="flex", + ) + body = client.request_body + assert body["tools"] == _TOOLS + assert (body["tool_choice"], body["parallel_tool_calls"]) == ("auto", False) + assert (body["user"], body["service_tier"]) == ("end-user-7", "flex") + assert response.choices[0].finish_reason == "tool_calls" + assert response.choices[0].message.tool_calls[0].function.name == "get_weather" + + def test_streamed_tool_call_round_trip(self): + client = _StreamingToolCallingNadir() + chunks = list( + litellm.completion( + model="nadir/auto", + messages=[{"role": "user", "content": "Weather in Paris?"}], + api_key="sk-test", + client=client, + stream=True, + tools=_TOOLS, + tool_choice="auto", + ) + ) + assert (client.request_body["stream"], client.request_body["tools"]) == (True, _TOOLS) + choices = [choice for chunk in chunks for choice in chunk.choices] + calls = [call for choice in choices for call in choice.delta.tool_calls or []] + assert calls[0].function.name == "get_weather" + assert "".join(call.function.arguments or "" for call in calls) == '{"city": "Paris"}' + assert [choice.finish_reason for choice in choices if choice.finish_reason] == ["tool_calls"] + + class TestNadirEnvValidation: def test_validate_environment_detects_key(self, monkeypatch): monkeypatch.setenv("NADIR_API_KEY", "sk-live-xyz") @@ -179,6 +303,22 @@ class TestNadirCostAttribution: assert _logged_cost(res) == pytest.approx(0.00123) assert _logged_cost(res) != _cost(res, "anthropic") + def test_priced_breakdown_keeps_the_reported_cost(self): + res = _transform( + _payload(nadir_metadata={"cost": {"total_cost_usd": 0.00123, "cost_breakdown": {"pricing_failed": False}}}) + ) + assert _logged_cost(res) == pytest.approx(0.00123) + + @pytest.mark.parametrize("reported_total", [0.0, 0.0004]) + def test_unpriced_call_prices_the_routed_model_instead(self, reported_total): + res = _transform( + _payload( + nadir_metadata={"cost": {"total_cost_usd": reported_total, "cost_breakdown": {"pricing_failed": True}}} + ) + ) + assert COST_HEADER not in res._hidden_params.get("additional_headers", {}) + assert _logged_cost(res) == _cost(res, "anthropic") > 0 + def test_routed_model_is_preserved(self): res = _transform(_payload(nadir_metadata={"cost": {"total_cost_usd": 0.001}})) assert res.model == "claude-haiku-4-5"