Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 15 additions & 0 deletions dspy/adapters/types/base_type.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
if TYPE_CHECKING:
from litellm import ModelResponseStream

from dspy.core.types import LMOutput
from dspy.signatures.signature import Signature

CUSTOM_TYPE_START_IDENTIFIER = "<<CUSTOM-TYPE-START-IDENTIFIER>>"
Expand Down Expand Up @@ -131,6 +132,20 @@ def parse_lm_response(cls, response: str | dict[str, Any]) -> Optional["Type"]:
"""
return None

@classmethod
def parse_lm_output(cls, output: "LMOutput") -> Optional["Type"]:
"""Parse a normalized LM output into the custom type.

Args:
output: A normalized LM output.

Returns:
A custom type object.
"""
data = output.to_output_dict()
response = data if any(key in data for key in ("reasoning_content", "citations", "tool_calls", "logprobs")) else output.text
return cls.parse_lm_response(response)


def split_message_content_for_custom_types(messages: list[dict[str, Any]]) -> list[dict[str, Any]]:
"""Split user message content into a list of content blocks.
Expand Down
25 changes: 23 additions & 2 deletions dspy/adapters/types/citation.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,13 @@
from typing import Any, Optional
from typing import TYPE_CHECKING, Any, Optional

import pydantic

from dspy.adapters.types.base_type import Type
from dspy.utils.annotation import experimental

if TYPE_CHECKING:
from dspy.core.types import LMOutput


@experimental(version="3.0.4")
class Citations(Type):
Expand Down Expand Up @@ -113,7 +116,7 @@ def from_dict_list(cls, citations_dicts: list[dict[str, Any]]) -> "Citations":
citations = Citations.from_dict_list(citations_dict)
```
"""
citations = [cls.Citation(**item) for item in citations_dicts]
citations = [cls.Citation(**_normalize_citation_dict(item)) for item in citations_dicts]
return cls(citations=citations)

@classmethod
Expand Down Expand Up @@ -219,3 +222,21 @@ def parse_lm_response(cls, response: str | dict[str, Any]) -> Optional["Citation
return cls.from_dict_list(citations_data)

return None

@classmethod
def parse_lm_output(cls, output: "LMOutput") -> Optional["Citations"]:
"""Parse a normalized LM output into Citations."""
if output.citations:
return cls.from_dict_list([citation.model_dump(exclude_none=True) for citation in output.citations])
return None


def _normalize_citation_dict(item: dict[str, Any]) -> dict[str, Any]:
data = {**item.get("metadata", {}), **item}
data.pop("metadata", None)
data.pop("type", None)
if "cited_text" not in data and "text" in data:
data["cited_text"] = data.pop("text")
if "document_title" not in data and "title" in data:
data["document_title"] = data.pop("title")
return data
8 changes: 8 additions & 0 deletions dspy/adapters/types/reasoning.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from dspy.clients.base_lm import BaseLM

if TYPE_CHECKING:
from dspy.core.types import LMOutput
from dspy.signatures.signature import Signature


Expand Down Expand Up @@ -83,6 +84,13 @@ def parse_lm_response(cls, response: str | dict[str, Any]) -> Optional["Reasonin
return Reasoning(content=response["reasoning_content"])
return None

@classmethod
def parse_lm_output(cls, output: "LMOutput") -> Optional["Reasoning"]:
"""Parse the LM output into a Reasoning object."""
if output.reasoning_content is not None:
return Reasoning(content=output.reasoning_content)
return None

@classmethod
def parse_stream_chunk(cls, chunk) -> str | None:
"""
Expand Down
21 changes: 16 additions & 5 deletions dspy/core/types.py
Original file line number Diff line number Diff line change
Expand Up @@ -1598,6 +1598,8 @@ def _history_message_parts_as_openai_content(parts: list[LMPart]) -> str | list[


def _history_part_as_openai_content(part: LMPart) -> dict[str, Any]:
if legacy_block := getattr(part, "metadata", {}).get("legacy_content_block"):
return dict(legacy_block)
if isinstance(part, LMTextPart):
return {"type": "text", "text": part.text}
if isinstance(part, LMImagePart):
Expand Down Expand Up @@ -1676,6 +1678,11 @@ def _history_media_format(media_type: str) -> str:
def _history_request_kwargs(request: LMRequest) -> dict[str, Any]:
data = request.config.model_dump(exclude_none=True)
extensions = data.pop("extensions", {}) or {}
if (reasoning := data.pop("reasoning", None)) is not None:
if effort := reasoning.get("effort"):
data["reasoning_effort"] = effort
else:
data["reasoning"] = reasoning
return {**extensions, **data}


Expand Down Expand Up @@ -1705,13 +1712,15 @@ def _messages_from_items(items: tuple[Any, ...], *, prompt: str | None = None) -
if len(items) == 1 and _is_message_sequence(items[0]):
items = tuple(items[0])

if all(isinstance(item, LMMessage) or isinstance(item, LMResponse) for item in items):
if all(_is_message_item(item) for item in items):
messages: list[LMMessage] = []
for item in items:
if isinstance(item, LMMessage):
messages.append(item)
else:
elif isinstance(item, LMResponse):
messages.extend(_messages_from_response(item))
else:
messages.append(_coerce_message(item))
return messages, []

parts = [_coerce_part(item) for item in items]
Expand All @@ -1723,9 +1732,11 @@ def _messages_from_response(response: LMResponse) -> list[LMMessage]:


def _is_message_sequence(value: Any) -> bool:
return isinstance(value, (list, tuple)) and all(
isinstance(item, LMMessage) or isinstance(item, LMResponse) for item in value
)
return isinstance(value, (list, tuple)) and all(_is_message_item(item) for item in value)


def _is_message_item(value: Any) -> bool:
return isinstance(value, (LMMessage, LMResponse)) or (isinstance(value, dict) and "role" in value)


def _coerce_part(value: Any) -> LMPart:
Expand Down