@@ -31,7 +31,9 @@ def __init__(self, model_name: str, model_provider: str):
3131 self .model = setup_model (model_name , model_provider )
3232 self .chain = setup_chain (self .model , prompt = base_prompt )
3333
34- def stream_response (self , user_input : str , chain_args : dict = {}) -> Generator [str , None , None ]:
34+ def stream_response (
35+ self , user_input : str , chain_args : dict = {}, metadata : dict | None = None
36+ ) -> Generator [str , None , None ]:
3537 """Process user input and stream AI response."""
3638 # Add user message to history before streaming
3739 logger .debug (f'User input:\n "{ user_input } "' )
@@ -45,7 +47,11 @@ def stream_response(self, user_input: str, chain_args: dict = {}) -> Generator[s
4547 yield chunk
4648
4749 # After streaming is complete, add the full response to chat history
48- self .chat_history .append (AIMessage (content = full_response ))
50+ self .chat_history .append (
51+ AIMessage (
52+ content = full_response , response_metadata = {"context" : metadata } if metadata else {}
53+ )
54+ )
4955 logger .debug (f'AI response:\n "{ full_response } "' )
5056
5157
@@ -121,6 +127,13 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
121127 relevant_references = "\n " .join (
122128 [f"From { doc .metadata [RAG_DOC_ID ]} :\n { doc .page_content } " for doc in relevant_docs ]
123129 )
130+ relevant_metadata = [
131+ {
132+ "Document Title" : doc .metadata .get (RAG_DOC_ID , "N/A" ),
133+ "Content" : doc .page_content ,
134+ }
135+ for doc in relevant_docs
136+ ]
124137
125138 # Log the context documents
126139 logger .debug (f"Context: { len (relevant_docs )} documents returned." )
@@ -133,7 +146,9 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
133146 )
134147
135148 return super ().stream_response (
136- user_input , {"paper_text" : self .paper_text , "relevant_references" : relevant_references }
149+ user_input ,
150+ {"paper_text" : self .paper_text , "relevant_references" : relevant_references },
151+ relevant_metadata ,
137152 )
138153
139154
@@ -189,6 +204,13 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
189204 relevant_code = "\n " .join (
190205 [f"From { doc .metadata [RAG_DOC_ID ]} :\n { doc .page_content } " for doc in relevant_docs ]
191206 )
207+ relevant_metadata = [
208+ {
209+ "Document Title" : doc .metadata .get (RAG_DOC_ID , "N/A" ),
210+ "Content" : doc .page_content ,
211+ }
212+ for doc in relevant_docs
213+ ]
192214
193215 # Log the context documents
194216 logger .debug (f"Context: { len (relevant_docs )} documents returned." )
@@ -201,7 +223,9 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
201223 )
202224
203225 return super ().stream_response (
204- user_input , {"paper_text" : self .paper_text , "relevant_code" : relevant_code }
226+ user_input ,
227+ {"paper_text" : self .paper_text , "relevant_code" : relevant_code },
228+ relevant_metadata ,
205229 )
206230
207231
@@ -264,14 +288,22 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
264288 relevant_references = "\n " .join (
265289 [f"From { doc .metadata [RAG_DOC_ID ]} :\n { doc .page_content } " for doc in relevant_docs ]
266290 )
291+ relevant_metadata = [
292+ {
293+ "Document Title" : doc .metadata .get (RAG_DOC_ID , "N/A" ),
294+ "Content" : doc .page_content ,
295+ }
296+ for doc in relevant_docs
297+ ]
267298
268299 # Log the context documents
269300 logger .debug (f"Context: { len (relevant_docs )} documents returned." )
270- for i , doc in enumerate (relevant_docs , start = 1 ):
271- contents = doc . page_content [:200 ].replace ("\n " , " " )
301+ for i , doc in enumerate (relevant_metadata , start = 1 ):
302+ contents = doc [ "Content" ] [:200 ].replace ("\n " , " " )
272303 logger .debug (
273- f"""Context Document { i } :\n Document Title: { doc .metadata .get (RAG_DOC_ID , "N/A" )}
274- Page Content: { contents } ...
304+ f"""Context Document { i } :
305+ Document Title: { doc ["Document Title" ]}
306+ Page Content: { contents }
275307 """
276308 )
277309
@@ -281,4 +313,5 @@ def stream_response(self, user_input: str) -> Generator[str, None, None]:
281313 "paper_text" : self .paper_text ,
282314 "relevant_references" : relevant_references ,
283315 },
316+ relevant_metadata ,
284317 )
0 commit comments