Chat-Service / app /agent /graph.py
ArabicNewsAnalyzer's picture
Update app/agent/graph.py
cf34999 verified
Raw
History Blame Contribute Delete
3.65 kB
from __future__ import annotations
from app.memory.history_trimmer import maybe_trim_history
import json
from typing import Any, Optional
from langchain.agents import create_agent
from langchain_nvidia_ai_endpoints import ChatNVIDIA
from langgraph.checkpoint.redis.aio import AsyncRedisSaver
from app.agent.prompts import SYSTEM_PROMPT
from app.agent.state import AgentState
from app.config import get_settings
from app.tools.graph_tool import graph_query_tool
from app.tools.sql_tool import sql_query_tool
from app.tools.article_tool import get_article_detail
from app.tools.hybrid_tool import hybrid_search_tool
DEFAULT_MODEL = "z-ai/glm-5.2"
_agent = None
def build_agent(checkpointer: AsyncRedisSaver):
"""Called once at startup after the checkpointer is ready."""
global _agent
settings = get_settings()
llm = ChatNVIDIA(
api_key=settings.nvidia_api_key,
model=DEFAULT_MODEL,
temperature=0,
)
_agent = create_agent(
model=llm,
tools=[sql_query_tool, graph_query_tool, hybrid_search_tool],
state_schema=AgentState,
system_prompt=SYSTEM_PROMPT,
checkpointer=checkpointer,
)
return _agent
def get_agent():
if _agent is None:
raise RuntimeError(
"Agent not initialized — call build_agent() at startup first")
return _agent
def _extract_sources(messages: list) -> list[dict]:
sources: list[dict] = []
for m in messages:
if getattr(m, "type", None) != "tool":
continue
try:
payload = json.loads(m.content) if isinstance(
m.content, str) else m.content
except (json.JSONDecodeError, TypeError):
continue
refs = payload.get("source_refs") if isinstance(
payload, dict) else None
if refs:
sources.extend(refs)
seen: set[tuple] = set()
deduped: list[dict] = []
for s in sources:
key = (s.get("type"), s.get("id"))
if key not in seen:
seen.add(key)
deduped.append(s)
return deduped
async def _enrich_sources(raw_refs: list[dict]) -> list[dict]:
enriched: list[dict] = []
for ref in raw_refs:
if ref.get("type") == "article" and ref.get("id") is not None:
detail = await get_article_detail(int(ref["id"]))
rows = detail.get("rows") or []
if rows:
article = rows[0]
enriched.append({
"type": "article",
"id": article["id"],
"title": article.get("title"),
"url": article.get("url"),
})
continue
enriched.append(ref)
return enriched
async def run_agent(message: str, session_id: str, user_id: Optional[str]) -> dict[str, Any]:
agent = get_agent()
config = {"configurable": {"thread_id": session_id}}
result = await agent.ainvoke(
{
"messages": [{"role": "user", "content": message}],
"sources": [],
"session_id": session_id,
"user_id": user_id,
},
config=config,
)
final_message = result["messages"][-1]
answer = final_message.content if isinstance(
final_message.content, str) else str(final_message.content)
raw_sources = _extract_sources(result["messages"])
sources = await _enrich_sources(raw_sources)
await maybe_trim_history(agent, config)
return {"answer": answer, "sources": sources}