Source code for klea_utils.nodes.answer_general

#!/usr/bin/env python3
"""
Answer general question node

File: klea_utils/nodes/answer_general.py

Copyright 2026 Ankur Sinha
Author: Ankur Sinha <sanjay DOT ankur AT gmail DOT com>
"""

import logging
from typing import Any, override

from langchain_core.messages import AIMessage
from pydantic import BaseModel

from ..llm import (
    content_to_str,
    extract_llm_output_content,
    format_alert,
    prompt_value_to_messages,
    split_output_by_section,
)
from ..nodes.abstract import NodeStreamData
from .base import BaseLLMNode


[docs] class FallbackConfig(BaseModel): enabled: bool = False warning: str = ""
[docs] class AnswerGeneral(BaseLLMNode): model_type = "chat" model_defaults = {"temperature": 0.3, "max_output_tokens": 2048} """Answer general (non-domain) questions using the LLM's training data. Provides a conversational, user-friendly response. Optionally appends conversation history for context and a fallback warning when configured. """ def __init__( self, logger: logging.Logger, label: str, llm_models: dict[str, Any], memory: bool = False, num_history_messages: int = 10, fallback_config: FallbackConfig | None = None, ): """Initialise the general answer node. :param logger: Logger instance :param label: Human-readable label for UI progress display :param llm_models: ``{role: LLMModel}`` dict (from ``BaseLangGraph.llm_models``) :param memory: Whether to include conversation history in the prompt :param num_history_messages: Number of recent messages to include when memory is enabled :param fallback_config: Optional config for fallback warning text """ super().__init__( logger=logger, label=label, llm_models=llm_models, output_schema=None, memory=memory, ) self.num_history_messages = num_history_messages self.fallback_config = fallback_config @override def _get_prompt_variables(self, state: BaseModel) -> dict: """Format prompt with the user's query.""" return {"query": state.query} # type: ignore @override def _update_state(self, result: Any, state: BaseModel) -> dict[str, Any]: """Extract answer, append fallback warning if configured, update messages.""" answer = "" # Add fallback warning if configured and query was domain-related. # The RAG stores classified domains in ``query_domains`` (a list); # a genuinely non-domain query is classified as ``["undefined"]``. # Warn only when a real domain matched, i.e. the query fell back to # training data after failed retrieval, not for plain general chat. # Default to ``["undefined"]`` so states without the attribute # (e.g. non-RAG graphs) never show the warning. fallback = self.fallback_config if fallback and fallback.enabled and fallback.warning: query_domains = getattr(state, "query_domains", ["undefined"]) if "undefined" not in query_domains: answer += f"\n\n{format_alert(fallback.warning)}\n\n" content = content_to_str(result.content) thought, answer_text = split_output_by_section(content, "<think>", "</think>") answer += answer_text messages = list(state.messages) # type: ignore result.content = answer messages.append(result) return {"messages": messages, "message_for_user": answer} @override def _get_default_error_result(self) -> AIMessage: """Return default result when processing fails.""" return AIMessage(content="") @override def _get_info(self) -> NodeStreamData | None: """Return answer summary.""" assert self._last_state_updates is not None result = content_to_str(self._last_state_updates.get("message_for_user", "")) char_count = len(result) return NodeStreamData( heading="General Answer", summary=f"Generated answer ({char_count} characters)" if char_count else "No answer generated", details={"character_count": char_count}, ) @override def _get_debug(self) -> NodeStreamData | None: """Return info + input/output triples.""" assert self._last_prompt is not None assert self._last_output is not None info = self._get_info() if not info: return None details = info.details.copy() details.update( { "input_prompt": prompt_value_to_messages(self._last_prompt), "unprocessed_output": extract_llm_output_content(self._last_output), } ) return NodeStreamData( heading=info.heading, summary=info.summary, details=details )