diff --git a/dspy/adapters/types/base_type.py b/dspy/adapters/types/base_type.py index 13a55727f2..0696f43427 100644 --- a/dspy/adapters/types/base_type.py +++ b/dspy/adapters/types/base_type.py @@ -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 = "<>" @@ -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. diff --git a/dspy/adapters/types/citation.py b/dspy/adapters/types/citation.py index a8c5461a58..0c4684ab05 100644 --- a/dspy/adapters/types/citation.py +++ b/dspy/adapters/types/citation.py @@ -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): @@ -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 @@ -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 diff --git a/dspy/adapters/types/reasoning.py b/dspy/adapters/types/reasoning.py index c4fa3de742..04ae55c0e9 100644 --- a/dspy/adapters/types/reasoning.py +++ b/dspy/adapters/types/reasoning.py @@ -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 @@ -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: """ diff --git a/dspy/core/types.py b/dspy/core/types.py index 6ceb9eb447..68baf075dd 100644 --- a/dspy/core/types.py +++ b/dspy/core/types.py @@ -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): @@ -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} @@ -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] @@ -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: