update model without structure output

This commit is contained in:
Xin Wang
2026-03-31 08:49:32 +08:00
parent 00ef163767
commit 2219d82eb6
8 changed files with 89 additions and 18 deletions

View File

@@ -21,7 +21,7 @@ def main() -> None:
)
parser.add_argument(
"--reasoning-model",
default="GLM-4.5-Air",
default="GLM-4.7",
help="Model for the final answer",
)
args = parser.parse_args()

View File

@@ -9,21 +9,21 @@ class Configuration(BaseModel):
"""The configuration for the agent."""
query_generator_model: str = Field(
default="GLM-4.5-Air",
default="GLM-4.7",
metadata={
"description": "The name of the language model to use for the agent's query generation."
},
)
reflection_model: str = Field(
default="GLM-4.5-Air",
default="GLM-4.7",
metadata={
"description": "The name of the language model to use for the agent's reflection."
},
)
answer_model: str = Field(
default="GLM-4.5-Air",
default="GLM-4.7",
metadata={
"description": "The name of the language model to use for the agent's answer."
},

View File

@@ -1,7 +1,9 @@
import os
from typing import TypeVar
from agent.tools_and_schemas import SearchQueryList, Reflection
from dotenv import load_dotenv
from pydantic import BaseModel, ValidationError
from langchain_core.messages import AIMessage
from langgraph.types import Send
from langgraph.graph import StateGraph
@@ -38,6 +40,7 @@ if os.getenv("TAVILY_API_KEY") is None:
raise ValueError("TAVILY_API_KEY is not set")
tavily_client = TavilyClient(api_key=os.getenv("TAVILY_API_KEY"))
SchemaT = TypeVar("SchemaT", bound=BaseModel)
def create_llm(model: str, temperature: float) -> ChatOpenAI:
@@ -51,6 +54,63 @@ def create_llm(model: str, temperature: float) -> ChatOpenAI:
)
def _extract_json_payload(content: str) -> str:
"""Extract a JSON object from a model response."""
text = content.strip()
if text.startswith("```"):
lines = text.splitlines()
if lines:
lines = lines[1:]
if lines and lines[-1].strip() == "```":
lines = lines[:-1]
text = "\n".join(lines).strip()
if text.startswith("json"):
text = text[4:].lstrip()
return text
def _invoke_json_model(
llm: ChatOpenAI,
prompt: str,
schema: type[SchemaT],
) -> SchemaT:
"""Invoke the model for JSON text and validate it with Pydantic."""
result = llm.invoke(prompt)
content = result.content if isinstance(result.content, str) else str(result.content)
json_payload = _extract_json_payload(content)
try:
return schema.model_validate_json(json_payload)
except ValidationError as exc:
raise ValueError(
f"Model returned invalid JSON for {schema.__name__}: {json_payload[:500]!r}"
) from exc
def _supports_native_structured_output(model: str) -> bool:
"""Return whether the selected model should use LangChain structured output."""
normalized = model.strip().lower()
return normalized.startswith("glm-4.7")
def _invoke_structured_model(
llm: ChatOpenAI,
model: str,
prompt: str,
schema: type[SchemaT],
) -> SchemaT:
"""Use native structured output when supported, otherwise fall back to JSON parsing."""
if _supports_native_structured_output(model):
try:
return llm.with_structured_output(schema).invoke(prompt)
except ValueError:
return _invoke_json_model(llm, prompt, schema)
return _invoke_json_model(llm, prompt, schema)
# Nodes
def generate_query(state: OverallState, config: RunnableConfig) -> QueryGenerationState:
"""LangGraph node that generates search queries based on the User's question.
@@ -71,7 +131,6 @@ def generate_query(state: OverallState, config: RunnableConfig) -> QueryGenerati
state["initial_search_query_count"] = configurable.number_of_initial_queries
llm = create_llm(configurable.query_generator_model, temperature=1.0)
structured_llm = llm.with_structured_output(SearchQueryList)
current_date = get_current_date()
formatted_prompt = query_writer_instructions.format(
@@ -79,7 +138,12 @@ def generate_query(state: OverallState, config: RunnableConfig) -> QueryGenerati
research_topic=get_research_topic(state["messages"]),
number_queries=state["initial_search_query_count"],
)
result = structured_llm.invoke(formatted_prompt)
result = _invoke_structured_model(
llm,
configurable.query_generator_model,
formatted_prompt,
SearchQueryList,
)
return {"search_query": result.query}
@@ -153,7 +217,12 @@ def reflection(state: OverallState, config: RunnableConfig) -> ReflectionState:
summaries="\n\n---\n\n".join(state["web_research_result"]),
)
llm = create_llm(reasoning_model, temperature=1.0)
result = llm.with_structured_output(Reflection).invoke(formatted_prompt)
result = _invoke_structured_model(
llm,
reasoning_model,
formatted_prompt,
Reflection,
)
return {
"is_sufficient": result.is_sufficient,

View File

@@ -15,6 +15,7 @@ Instructions:
- Queries should be diverse, if the topic is broad, generate more than 1 query.
- Don't generate multiple similar queries, 1 is enough.
- Query should ensure that the most current information is gathered. The current date is {current_date}.
- Return only valid JSON. Do not include markdown fences or any extra text.
Format:
- Format your response as a JSON object with ALL two of these exact keys:
@@ -27,7 +28,7 @@ Topic: What revenue grew more last year apple stock or the number of people buyi
```json
{{
"rationale": "To answer this comparative growth question accurately, we need specific data points on Apple's stock performance and iPhone sales metrics. These queries target the precise financial information needed: company revenue trends, product-specific unit sales figures, and stock price movement over the same fiscal period for direct comparison.",
"query": ["Apple total revenue growth fiscal year 2024", "iPhone unit sales growth fiscal year 2024", "Apple stock price growth fiscal year 2024"],
"query": ["Apple total revenue growth fiscal year 2024", "iPhone unit sales growth fiscal year 2024", "Apple stock price growth fiscal year 2024"]
}}
```
@@ -57,6 +58,7 @@ Instructions:
Requirements:
- Ensure the follow-up query is self-contained and includes necessary context for web search.
- Return only valid JSON. Do not include markdown fences or any extra text.
Output Format:
- Format your response as a JSON object with these exact keys:

View File

@@ -59,7 +59,7 @@ export function ActivityTimeline({
}, [isLoading, processedEvents]);
return (
<Card className="border-none rounded-lg bg-neutral-700 max-h-96">
<Card className="w-full min-w-0 border-none rounded-lg bg-neutral-700 max-h-96">
<CardHeader>
<CardDescription className="flex items-center justify-between">
<div
@@ -76,7 +76,7 @@ export function ActivityTimeline({
</CardDescription>
</CardHeader>
{!isTimelineCollapsed && (
<ScrollArea className="max-h-96 overflow-y-auto">
<ScrollArea className="w-full max-h-96 overflow-y-auto">
<CardContent>
{isLoading && processedEvents.length === 0 && (
<div className="relative pl-8 pb-4">

View File

@@ -187,9 +187,9 @@ const AiMessageBubble: React.FC<AiMessageBubbleProps> = ({
const isLiveActivityForThisBubble = isLastMessage && isOverallLoading;
return (
<div className={`relative break-words flex flex-col`}>
<div className="relative flex w-full min-w-0 max-w-[85%] break-words flex-col md:max-w-[80%]">
{activityForThisBubble && activityForThisBubble.length > 0 && (
<div className="mb-3 border-b border-neutral-700 pb-3 text-xs">
<div className="mb-3 w-full min-w-0 border-b border-neutral-700 pb-3 text-xs">
<ActivityTimeline
processedEvents={activityForThisBubble}
isLoading={isLiveActivityForThisBubble}
@@ -289,12 +289,12 @@ export function ChatMessagesView({
{isLoading &&
(messages.length === 0 ||
messages[messages.length - 1].type === "human") && (
<div className="flex items-start gap-3 mt-3">
<div className="mt-3 flex items-start gap-3">
{" "}
{/* AI message row structure */}
<div className="relative group max-w-[85%] md:max-w-[80%] rounded-xl p-3 shadow-sm break-words bg-neutral-800 text-neutral-100 rounded-bl-none w-full min-h-[56px]">
{liveActivityEvents.length > 0 ? (
<div className="text-xs">
<div className="w-full min-w-0 text-xs">
<ActivityTimeline
processedEvents={liveActivityEvents}
isLoading={true}

View File

@@ -26,7 +26,7 @@ export const InputForm: React.FC<InputFormProps> = ({
}) => {
const [internalInputValue, setInternalInputValue] = useState("");
const [effort, setEffort] = useState("medium");
const [model, setModel] = useState("GLM-4.5-Air");
const [model, setModel] = useState("GLM-4.7");
const handleInternalSubmit = (e?: React.FormEvent) => {
if (e) e.preventDefault();
@@ -136,11 +136,11 @@ export const InputForm: React.FC<InputFormProps> = ({
</SelectTrigger>
<SelectContent className="bg-neutral-700 border-neutral-600 text-neutral-300 cursor-pointer">
<SelectItem
value="GLM-4.5-Air"
value="GLM-4.7"
className="hover:bg-neutral-600 focus:bg-neutral-600 cursor-pointer"
>
<div className="flex items-center">
<Zap className="h-4 w-4 mr-2 text-cyan-400" /> GLM-4.5-Air
<Zap className="h-4 w-4 mr-2 text-cyan-400" /> GLM-4.7
</div>
</SelectItem>
</SelectContent>

View File

@@ -33,7 +33,7 @@ export const WelcomeScreen: React.FC<WelcomeScreenProps> = ({
/>
</div>
<p className="text-xs text-neutral-500">
Powered by GLM-4.5-Air, Tavily, and LangChain LangGraph.
Powered by GLM-4.7, Tavily, and LangChain LangGraph.
</p>
</div>
);