Add response_format arg

This commit is contained in:
Nuno Campos
2025-08-25 21:20:26 +01:00
parent e1aeb24a4e
commit 54272afe01
2 changed files with 21 additions and 2 deletions
+17 -2
View File
@@ -5,7 +5,13 @@ from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AIMessage, SystemMessage
from langchain_core.tools import BaseTool
from langgraph.agent.types import AgentInput, AgentMiddleware, AgentState, ModelRequest
from langgraph.agent.types import (
AgentInput,
AgentMiddleware,
AgentState,
ModelRequest,
ResponseFormat,
)
from langgraph.constants import END, START
from langgraph.graph.state import StateGraph
from langgraph.prebuilt.tool_node import ToolNode
@@ -17,6 +23,7 @@ def create_agent(
tools: Sequence[BaseTool | Callable],
system_prompt: str,
middleware: Sequence[AgentMiddleware] = (),
response_format: ResponseFormat | None = None,
) -> StateGraph[AgentState, None, AgentInput]:
# init chat model
if isinstance(model, str):
@@ -59,6 +66,7 @@ def create_agent(
tools=list(tool_node.tools_by_name.values()),
system_prompt=system_prompt,
middleware=middleware,
response_format=response_format,
),
)
graph.add_node("tools", tool_node)
@@ -123,6 +131,7 @@ def _make_model_request_node(
model: BaseChatModel,
tools: Sequence[BaseTool],
middleware: Sequence[AgentMiddleware],
response_format: ResponseFormat | None = None,
) -> Callable[[AgentState], AgentState]:
def model_request(state: AgentState) -> AgentState:
# create request
@@ -132,6 +141,7 @@ def _make_model_request_node(
messages=state.messages,
tool_choice=None,
tools=tools,
response_format=response_format,
)
# visit middleware in order
for mw in middleware:
@@ -141,8 +151,13 @@ def _make_model_request_node(
messages = [SystemMessage(request.system_prompt)] + request.messages
else:
messages = request.messages
# prepare model
if request.response_format:
model_ = request.model.with_structured_output(request.response_format)
else:
model_ = request.model
# call model
output = request.model.invoke(
output = model_.invoke(
messages, tools=request.tools, tool_choice=request.tool_choice
)
return {"messages": output}
+4
View File
@@ -7,10 +7,13 @@ from typing import Annotated, Any, Self
from langchain_core.language_models.chat_models import BaseChatModel
from langchain_core.messages import AnyMessage
from langchain_core.tools import BaseTool
from pydantic import BaseModel
from typing_extensions import TypedDict
from langgraph.graph.message import Messages, add_messages
ResponseFormat = dict | type[BaseModel]
@dataclass
class ModelRequest:
@@ -19,6 +22,7 @@ class ModelRequest:
messages: Sequence[AnyMessage] # excluding system prompt
tool_choice: Any
tools: Sequence[BaseTool]
response_format: ResponseFormat | None
class AgentMiddleware: