Skip to content

Commit aa1c5b2

Browse files
committed
clean up
refactor and addressing some @awilfox and copilot suggestions
1 parent 20fc9f9 commit aa1c5b2

1 file changed

Lines changed: 19 additions & 13 deletions

File tree

‎willa/chatbot/graph_manager.py‎

Lines changed: 19 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22
from typing import Optional, Annotated, NotRequired
33
from typing_extensions import TypedDict
44

5+
from langchain_core.documents import Document
56
from langchain_core.language_models import BaseChatModel
67
from langchain_core.messages import ChatMessage, HumanMessage, AIMessage
78
from langchain_core.vectorstores.base import VectorStore
@@ -72,7 +73,7 @@ def _filter_messages(self, state: WillaChatbotState) -> dict[str, list[AnyMessag
7273

7374
filtered = [
7475
msg for msg in messages
75-
if 'tind' not in msg.response_metadata and msg.type != "system"
76+
if "tind" not in getattr(msg, "response_metadata", {}) and msg.type != "system"
7677
]
7778
return {"filtered_messages": filtered}
7879

@@ -91,6 +92,21 @@ def _prepare_search_query(self, state: WillaChatbotState) -> dict[str, str]:
9192

9293
return {"search_query": search_query}
9394

95+
def _format_retrieved_documents(self, matching_docs: list[Document]) -> list[dict[str, str]]:
96+
"""Format documents from vector store into a list of dictionaries."""
97+
formatted_documents: list[dict[str, str]] = []
98+
for i, doc in enumerate(matching_docs, 1):
99+
tind_metadata = doc.metadata.get('tind_metadata', {})
100+
tind_id = tind_metadata.get('tind_id', [''])[0]
101+
formatted_documents.append({
102+
"id": f"{i}_{tind_id}",
103+
"page_content": doc.page_content,
104+
"title": tind_metadata.get('title', [''])[0],
105+
"project": tind_metadata.get('isPartOf', [''])[0],
106+
"tind_link": format_tind_context.get_tind_url(tind_id)
107+
})
108+
return formatted_documents
109+
94110
def _retrieve_context(self, state: WillaChatbotState) -> dict[str, str | list[dict[str, str]]]:
95111
"""Retrieve relevant context from vector store."""
96112
search_query = state.get("search_query", "")
@@ -102,17 +118,7 @@ def _retrieve_context(self, state: WillaChatbotState) -> dict[str, str | list[di
102118
# Search for relevant documents
103119
retriever = vector_store.as_retriever(search_kwargs={"k": int(CONFIG['K_VALUE'])})
104120
matching_docs = retriever.invoke(search_query)
105-
formatted_documents = [
106-
{
107-
"id": f"{i}_{doc.metadata.get('tind_metadata', {}).get('tind_id', [''])[0]}",
108-
"page_content": doc.page_content,
109-
"title": doc.metadata.get('tind_metadata', {}).get('title', [''])[0],
110-
"project": doc.metadata.get('tind_metadata', {}).get('isPartOf', [''])[0],
111-
"tind_link": format_tind_context.get_tind_url(
112-
doc.metadata.get('tind_metadata', {}).get('tind_id', [''])[0])
113-
}
114-
for i, doc in enumerate(matching_docs, 1)
115-
]
121+
formatted_documents = self._format_retrieved_documents(matching_docs)
116122

117123
# Format tind metadata
118124
tind_metadata = format_tind_context.get_tind_context(matching_docs)
@@ -142,7 +148,7 @@ def _generate_response(self, state: WillaChatbotState) -> dict[str, list[AnyMess
142148
tind_metadata = state.get("tind_metadata", "")
143149
model = self._model
144150
documents = state.get("documents", [])
145-
messages = state["messages_for_generation"]
151+
messages = state.get("messages_for_generation") or state.get("messages", [])
146152

147153
if not model:
148154
return {"messages": [AIMessage(content="Model not available.")]}

0 commit comments

Comments
 (0)