Skip to content

Commit 4fc83de

Browse files
committed
fix(openai): support Langfuse arguments in stream helpers
1 parent 9593d37 commit 4fc83de

3 files changed

Lines changed: 700 additions & 39 deletions

File tree

langfuse/openai.py

Lines changed: 140 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -32,7 +32,7 @@
3232
from inspect import isawaitable, isclass
3333
from typing import Any, Optional, cast
3434

35-
from openai._types import NotGiven
35+
from openai._types import NotGiven, Omit
3636
from packaging.version import Version
3737
from pydantic import BaseModel
3838
from pydantic_core import to_jsonable_python
@@ -137,6 +137,24 @@ class OpenAiDefinition:
137137
min_version="1.50.0",
138138
max_version="1.92.0",
139139
),
140+
OpenAiDefinition(
141+
module="openai.resources.beta.chat.completions",
142+
object="Completions",
143+
method="stream",
144+
type="chat",
145+
sync=True,
146+
min_version="1.40.0",
147+
max_version="1.92.0",
148+
),
149+
OpenAiDefinition(
150+
module="openai.resources.beta.chat.completions",
151+
object="AsyncCompletions",
152+
method="stream",
153+
type="chat",
154+
sync=False,
155+
min_version="1.40.0",
156+
max_version="1.92.0",
157+
),
140158
OpenAiDefinition(
141159
module="openai.resources.chat.completions",
142160
object="Completions",
@@ -153,6 +171,22 @@ class OpenAiDefinition:
153171
sync=False,
154172
min_version="1.92.0",
155173
),
174+
OpenAiDefinition(
175+
module="openai.resources.chat.completions",
176+
object="Completions",
177+
method="stream",
178+
type="chat",
179+
sync=True,
180+
min_version="1.92.0",
181+
),
182+
OpenAiDefinition(
183+
module="openai.resources.chat.completions",
184+
object="AsyncCompletions",
185+
method="stream",
186+
type="chat",
187+
sync=False,
188+
min_version="1.92.0",
189+
),
156190
OpenAiDefinition(
157191
module="openai.resources.responses",
158192
object="Responses",
@@ -185,6 +219,22 @@ class OpenAiDefinition:
185219
sync=False,
186220
min_version="1.66.0",
187221
),
222+
OpenAiDefinition(
223+
module="openai.resources.responses",
224+
object="Responses",
225+
method="stream",
226+
type="chat",
227+
sync=True,
228+
min_version="1.66.0",
229+
),
230+
OpenAiDefinition(
231+
module="openai.resources.responses",
232+
object="AsyncResponses",
233+
method="stream",
234+
type="chat",
235+
sync=False,
236+
min_version="1.66.0",
237+
),
188238
OpenAiDefinition(
189239
module="openai.resources.embeddings",
190240
object="Embeddings",
@@ -204,10 +254,18 @@ class OpenAiDefinition:
204254

205255
_RESPONSES_PROMPT_FIELDS = ("tools", "tool_choice", "parallel_tool_calls")
206256
_STRUCTURED_OUTPUT_METADATA_FIELDS = ("response_format", "text_format")
257+
_LANGFUSE_STREAM_ARG_NAMES = (
258+
"name",
259+
"langfuse_prompt",
260+
"langfuse_public_key",
261+
"trace_id",
262+
"parent_observation_id",
263+
)
264+
_LANGFUSE_OPENAI_STREAM_ARGS = object()
207265

208266

209267
def _is_not_given(value: Any) -> bool:
210-
return isinstance(value, NotGiven)
268+
return isinstance(value, (NotGiven, Omit))
211269

212270

213271
def _get_attr_or_item(value: Any, key: str, default: Any = None) -> Any:
@@ -321,6 +379,20 @@ def __init__(
321379
self.args["trace_id"] = trace_id
322380
self.args["parent_observation_id"] = parent_observation_id
323381

382+
extra_body = kwargs.get("extra_body")
383+
if isinstance(extra_body, dict) and _LANGFUSE_OPENAI_STREAM_ARGS in extra_body:
384+
request_extra_body = extra_body.copy()
385+
stream_args = dict(request_extra_body.pop(_LANGFUSE_OPENAI_STREAM_ARGS))
386+
kwargs["extra_body"] = request_extra_body
387+
388+
if "metadata" in stream_args:
389+
self.metadata = stream_args.pop("metadata")
390+
self.args["metadata"] = _get_structured_output_metadata(
391+
self.metadata, kwargs
392+
)
393+
394+
self.args.update(stream_args)
395+
324396
self.kwargs = kwargs
325397

326398
def get_langfuse_args(self) -> Any:
@@ -356,13 +428,13 @@ def _extract_responses_prompt(kwargs: Any) -> Any:
356428
for key in _RESPONSES_PROMPT_FIELDS:
357429
value = kwargs.get(key, None)
358430

359-
if value is not None and not isinstance(value, NotGiven):
431+
if value is not None and not _is_not_given(value):
360432
prompt_fields[key] = _serialize_openai_value(value)
361433

362-
if isinstance(input_value, NotGiven):
434+
if _is_not_given(input_value):
363435
input_value = None
364436

365-
if isinstance(instructions, NotGiven):
437+
if _is_not_given(instructions):
366438
instructions = None
367439

368440
if instructions is None:
@@ -395,13 +467,15 @@ def _extract_chat_prompt(kwargs: Any) -> Any:
395467
"""Extracts the user input from prompts. Returns an array of messages or dict with messages and functions"""
396468
prompt = {}
397469

398-
if kwargs.get("functions") is not None:
470+
if kwargs.get("functions") is not None and not _is_not_given(kwargs["functions"]):
399471
prompt.update({"functions": kwargs["functions"]})
400472

401-
if kwargs.get("function_call") is not None:
473+
if kwargs.get("function_call") is not None and not _is_not_given(
474+
kwargs["function_call"]
475+
):
402476
prompt.update({"function_call": kwargs["function_call"]})
403477

404-
if kwargs.get("tools") is not None:
478+
if kwargs.get("tools") is not None and not _is_not_given(kwargs["tools"]):
405479
prompt.update({"tools": kwargs["tools"]})
406480

407481
if prompt:
@@ -531,11 +605,9 @@ def _get_langfuse_data_from_kwargs(resource: OpenAiDefinition, kwargs: Any) -> A
531605
raise ValueError("parent_observation_id requires trace_id to be set")
532606

533607
metadata = kwargs.get("metadata", {})
534-
if (
535-
metadata is not None
536-
and not isinstance(metadata, NotGiven)
537-
and not isinstance(metadata, dict)
538-
):
608+
if _is_not_given(metadata):
609+
metadata = {}
610+
elif metadata is not None and not isinstance(metadata, dict):
539611
if isinstance(metadata, BaseModel):
540612
metadata = _serialize_openai_value(metadata)
541613
else:
@@ -556,63 +628,61 @@ def _get_langfuse_data_from_kwargs(resource: OpenAiDefinition, kwargs: Any) -> A
556628

557629
parsed_temperature = (
558630
kwargs.get("temperature", 1)
559-
if not isinstance(kwargs.get("temperature", 1), NotGiven)
631+
if not _is_not_given(kwargs.get("temperature", 1))
560632
else 1
561633
)
562634

563635
parsed_max_tokens = (
564636
kwargs.get("max_tokens", float("inf"))
565-
if not isinstance(kwargs.get("max_tokens", float("inf")), NotGiven)
637+
if not _is_not_given(kwargs.get("max_tokens", float("inf")))
566638
else float("inf")
567639
)
568640

569641
parsed_max_completion_tokens = (
570642
kwargs.get("max_completion_tokens", None)
571-
if not isinstance(kwargs.get("max_completion_tokens", float("inf")), NotGiven)
643+
if not _is_not_given(kwargs.get("max_completion_tokens", float("inf")))
572644
else None
573645
)
574646

575647
parsed_top_p = (
576-
kwargs.get("top_p", 1)
577-
if not isinstance(kwargs.get("top_p", 1), NotGiven)
578-
else 1
648+
kwargs.get("top_p", 1) if not _is_not_given(kwargs.get("top_p", 1)) else 1
579649
)
580650

581651
parsed_frequency_penalty = (
582652
kwargs.get("frequency_penalty", 0)
583-
if not isinstance(kwargs.get("frequency_penalty", 0), NotGiven)
653+
if not _is_not_given(kwargs.get("frequency_penalty", 0))
584654
else 0
585655
)
586656

587657
parsed_presence_penalty = (
588658
kwargs.get("presence_penalty", 0)
589-
if not isinstance(kwargs.get("presence_penalty", 0), NotGiven)
659+
if not _is_not_given(kwargs.get("presence_penalty", 0))
590660
else 0
591661
)
592662

593663
parsed_seed = (
594664
kwargs.get("seed", None)
595-
if not isinstance(kwargs.get("seed", None), NotGiven)
665+
if not _is_not_given(kwargs.get("seed", None))
596666
else None
597667
)
598668

599-
parsed_n = kwargs.get("n", 1) if not isinstance(kwargs.get("n", 1), NotGiven) else 1
669+
parsed_n = kwargs.get("n", 1) if not _is_not_given(kwargs.get("n", 1)) else 1
600670

601671
parsed_service_tier = (
602672
kwargs.get("service_tier", None)
603-
if not isinstance(kwargs.get("service_tier", None), NotGiven)
673+
if not _is_not_given(kwargs.get("service_tier", None))
604674
else None
605675
)
606676

607677
if resource.type == "embedding":
608678
parsed_dimensions = (
609679
kwargs.get("dimensions", None)
610-
if not isinstance(kwargs.get("dimensions", None), NotGiven)
680+
if not _is_not_given(kwargs.get("dimensions", None))
611681
else None
612682
)
613683
parsed_encoding_format = (
614684
kwargs.get("encoding_format", "float")
615-
if not isinstance(kwargs.get("encoding_format", "float"), NotGiven)
685+
if not _is_not_given(kwargs.get("encoding_format", "float"))
616686
else "float"
617687
)
618688

@@ -1077,7 +1147,7 @@ def _instrument_openai_stream(
10771147
raw_iterator = response._iterator
10781148
completion_start_time: Optional[datetime] = None
10791149
is_finalized = False
1080-
close = response.close
1150+
close = response.response.close
10811151

10821152
def finalize_once() -> None:
10831153
nonlocal is_finalized
@@ -1115,7 +1185,7 @@ def traced_close() -> Any:
11151185
finalize_once()
11161186

11171187
response._iterator = traced_iterator()
1118-
response.close = traced_close
1188+
response.response.close = traced_close
11191189

11201190
return response
11211191

@@ -1139,7 +1209,7 @@ def _instrument_openai_async_stream(
11391209
raw_iterator = response._iterator
11401210
completion_start_time: Optional[datetime] = None
11411211
is_finalized = False
1142-
close = response.close
1212+
close = response.response.aclose
11431213

11441214
async def finalize_once() -> None:
11451215
nonlocal is_finalized
@@ -1180,7 +1250,7 @@ async def traced_aclose() -> Any:
11801250
return await traced_close()
11811251

11821252
response._iterator = traced_iterator()
1183-
response.close = traced_close
1253+
response.response.aclose = traced_close
11841254
response.aclose = traced_aclose
11851255

11861256
return response
@@ -1195,7 +1265,7 @@ def _get_raw_response_mode(kwargs: Any) -> Optional[str]:
11951265
"""
11961266
extra_headers = kwargs.get("extra_headers", None)
11971267

1198-
if extra_headers is None or isinstance(extra_headers, NotGiven):
1268+
if extra_headers is None or _is_not_given(extra_headers):
11991269
return None
12001270

12011271
try:
@@ -1244,6 +1314,36 @@ def _unwrap_raw_response(openai_response: Any) -> Any:
12441314
return openai_response
12451315

12461316

1317+
@_langfuse_wrapper
1318+
def _wrap_stream(
1319+
open_ai_resource: OpenAiDefinition, wrapped: Any, args: Any, kwargs: Any
1320+
) -> Any:
1321+
"""Forward Langfuse arguments to the create call owned by the stream manager."""
1322+
arg_names: tuple[str, ...] = _LANGFUSE_STREAM_ARG_NAMES
1323+
if open_ai_resource.module == "openai.resources.beta.chat.completions":
1324+
arg_names += ("metadata",)
1325+
1326+
langfuse_args = {name: kwargs[name] for name in arg_names if name in kwargs}
1327+
1328+
if not langfuse_args:
1329+
return wrapped(**kwargs)
1330+
1331+
if open_ai_resource.module == "openai.resources.responses" and any(
1332+
name in kwargs and not _is_not_given(kwargs[name])
1333+
for name in ("response_id", "starting_after")
1334+
):
1335+
return wrapped(**kwargs)
1336+
1337+
for name in langfuse_args:
1338+
kwargs.pop(name)
1339+
1340+
extra_body = dict(kwargs.get("extra_body") or {})
1341+
extra_body[_LANGFUSE_OPENAI_STREAM_ARGS] = langfuse_args
1342+
kwargs["extra_body"] = extra_body
1343+
1344+
return wrapped(**kwargs)
1345+
1346+
12471347
@_langfuse_wrapper
12481348
def _wrap(
12491349
open_ai_resource: OpenAiDefinition, wrapped: Any, args: Any, kwargs: Any
@@ -1436,10 +1536,18 @@ def register_tracing() -> None:
14361536
):
14371537
continue
14381538

1539+
wrapper = (
1540+
_wrap_stream(resource)
1541+
if resource.method == "stream"
1542+
else _wrap(resource)
1543+
if resource.sync
1544+
else _wrap_async(resource)
1545+
)
1546+
14391547
wrap_function_wrapper(
14401548
resource.module,
14411549
f"{resource.object}.{resource.method}",
1442-
_wrap(resource) if resource.sync else _wrap_async(resource),
1550+
wrapper,
14431551
)
14441552

14451553

0 commit comments

Comments
 (0)