diff --git a/langfuse/openai.py b/langfuse/openai.py index 16e7f1c1f..eec423e75 100644 --- a/langfuse/openai.py +++ b/langfuse/openai.py @@ -475,6 +475,9 @@ def _extract_chat_response(kwargs: Any) -> Any: if kwargs.get("tool_calls") is not None: response.update({"tool_calls": kwargs["tool_calls"]}) + if kwargs.get("refusal") is not None: + response.update({"refusal": kwargs["refusal"]}) + if kwargs.get("audio") is not None: audio = kwargs["audio"].__dict__ @@ -827,12 +830,9 @@ def _extract_streamed_openai_response(resource: Any, chunks: Any) -> Any: if delta.get("role", None) is not None: completion["role"] = delta["role"] - if delta.get("content", None) is not None: - completion["content"] = ( - delta.get("content", None) - if completion["content"] is None - else completion["content"] + delta.get("content", None) - ) + for field in ("content", "refusal"): + if delta.get(field) is not None: + completion[field] = (completion[field] or "") + delta[field] if delta.get("function_call", None) is not None: curr = completion["function_call"] @@ -907,6 +907,9 @@ def _extract_streamed_openai_response(resource: Any, chunks: Any) -> Any: def get_response_for_chat() -> Any: content = completion["content"] + if completion["refusal"] is not None: + return _extract_chat_response({**completion, "role": "assistant"}) + if completion["tool_calls"]: response = { "role": "assistant", diff --git a/tests/unit/test_openai.py b/tests/unit/test_openai.py index cacdae5c2..a27809166 100644 --- a/tests/unit/test_openai.py +++ b/tests/unit/test_openai.py @@ -1,7 +1,9 @@ import asyncio +import json from types import SimpleNamespace from unittest.mock import patch +import httpx import pytest from openai.types.responses import ParsedResponseOutputMessage, ParsedResponseOutputText from pydantic import BaseModel @@ -1367,6 +1369,95 @@ def handler(request: httpx.Request) -> httpx.Response: ) +@pytest.mark.asyncio +@pytest.mark.parametrize("async_client", [False, True]) +@pytest.mark.parametrize("stream", [False, True]) +async def test_chat_completion_captures_refusal( + langfuse_memory_client, get_span, json_attr, async_client, stream +): + refusal = "I cannot help with that request." + payload = _chat_completion_payload() + payload["choices"][0]["message"] = { + "role": "assistant", + "content": None, + "refusal": refusal, + } + + def handler(request): + if not stream: + return httpx.Response(200, json=payload) + + events = [] + for part in [None, "I cannot help ", "with that request."]: + events.append( + { + "id": payload["id"], + "object": "chat.completion.chunk", + "created": payload["created"], + "model": payload["model"], + "choices": [ + { + "index": 0, + "delta": {"role": "assistant", "refusal": part}, + "finish_reason": None, + } + ], + } + ) + events.append( + { + **events[-1], + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + "usage": payload["usage"], + } + ) + body = "".join(f"data: {json.dumps(event)}\n\n" for event in events) + return httpx.Response( + 200, + content=body + "data: [DONE]\n\n", + headers={"content-type": "text/event-stream"}, + ) + + kwargs = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "test request"}], + "stream": stream, + } + if async_client: + async with lf_openai.AsyncOpenAI( + api_key="test", + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) as client: + result = await client.chat.completions.create(**kwargs) + if stream: + result = [chunk async for chunk in result] + else: + with lf_openai.OpenAI( + api_key="test", + http_client=httpx.Client(transport=httpx.MockTransport(handler)), + ) as client: + result = client.chat.completions.create(**kwargs) + if stream: + result = list(result) + + if stream: + assert ( + "".join(chunk.choices[0].delta.refusal or "" for chunk in result) == refusal + ) + else: + assert result.choices[0].message.refusal == refusal + + langfuse_memory_client.flush() + span = get_span("OpenAI-generation") + assert json_attr(span, LangfuseOtelSpanAttributes.OBSERVATION_OUTPUT) == { + "role": "assistant", + "content": None, + "refusal": refusal, + } + usage = json_attr(span, LangfuseOtelSpanAttributes.OBSERVATION_USAGE_DETAILS) + assert usage["total_tokens"] == payload["usage"]["total_tokens"] + + def test_with_raw_response_chat_completion_captures_output_and_usage( langfuse_memory_client, get_span, json_attr ):