chATLAS_Chains 0.1.3__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
benchmark/basic.py ADDED
@@ -0,0 +1,34 @@
1
+ import os
2
+
3
+ from tqdm import tqdm
4
+
5
+ from chATLAS_Benchmark import BenchmarkTest
6
+ from chATLAS_Chains.chains.basic import basic_retrieval_chain
7
+ from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
8
+ from chATLAS_Chains.vectorstore import vectorstore
9
+
10
+ chain = basic_retrieval_chain(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vectorstore, model_name="gpt-4o-mini")
11
+
12
+ # Initialize the test set
13
+ questions_path = os.getenv("CHATLAS_BENCHMARK_QUESTIONS")
14
+ test = BenchmarkTest(questions_path)
15
+ # test = BenchmarkTest(questions_path, keys={"test_questions": "question", "test_documents":"documents", "test_answer": "answer"})
16
+
17
+ # --- Run the RAG on the questions ---
18
+ # Assuming RAG.run() returns an answer and list of docs for each question
19
+ gen_answers = []
20
+ gen_docs = []
21
+ for q in tqdm(test.questions):
22
+ result = chain.invoke(q)
23
+
24
+ gen_answers.append(result["answer"].content)
25
+ gen_docs.append(result["docs"])
26
+
27
+ # Set generated answers and documents on the test instance
28
+ test.set_generated_data(gen_answers, gen_docs)
29
+
30
+ # Run the scoring with any metrics you want
31
+ scores = test.score_test_set("LexicalMetrics", "SemanticSimilarity", "DocumentMatch")
32
+
33
+ # Save the results to the db
34
+ test.store_results(scores, db_name="results.db", name="basic_retrieval_chain")
File without changes
File without changes
File without changes
@@ -0,0 +1,42 @@
1
+ from operator import itemgetter
2
+
3
+ from langchain_core.prompts import ChatPromptTemplate
4
+ from langchain_core.runnables import RunnableParallel
5
+
6
+ from chATLAS_Chains.llm.model_selection import get_chat_model
7
+ from chATLAS_Chains.search.basic import search_runnable
8
+ from chATLAS_Chains.utils.doc_utils import combine_documents
9
+
10
+
11
+ def basic_retrieval_chain(prompt: str, vectorstore, model_name: str) -> RunnableParallel:
12
+ """
13
+ Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
14
+
15
+ :param prompt: The prompt template to use with the model.
16
+ :type prompt: str
17
+ :param vectorstore: The vectorstore or list of vectorstores to search over.
18
+ :type vectorstore: Any
19
+ :param model_name: The name of the chat model to use for generating responses.
20
+ :type model_name: str
21
+
22
+ :return: A LangChain RunnableParallel chain that performs retrieval and response generation.
23
+ :rtype: RunnableParallel
24
+
25
+ """
26
+ prompt = ChatPromptTemplate.from_template(prompt)
27
+ model = get_chat_model(model_name)
28
+
29
+ search = search_runnable(vectorstore)
30
+
31
+ final_inputs = {
32
+ "context": lambda x: combine_documents(x["docs"]),
33
+ "question": itemgetter("question"),
34
+ }
35
+
36
+ answer = {
37
+ "answer": final_inputs | prompt | model,
38
+ "docs": lambda x: x["docs"],
39
+ }
40
+
41
+ chain = search | answer
42
+ return chain
@@ -0,0 +1,163 @@
1
+ """
2
+ Example graph for running langgraph with this general setup and the postgres vector stores.
3
+ """
4
+
5
+ from typing import TypedDict
6
+
7
+ import langgraph.graph as lg
8
+ from langchain_core.documents import Document
9
+ from langchain_core.prompts import ChatPromptTemplate
10
+ from langgraph.graph import END, StateGraph
11
+
12
+ from chATLAS_Chains.llm.model_selection import get_chat_model
13
+ from chATLAS_Chains.utils.doc_utils import combine_documents
14
+ from chATLAS_Embed import LangChainVectorStore
15
+
16
+
17
+ # Define TypedDict for the state
18
+ class GraphState(TypedDict, total=False):
19
+ question: str
20
+ search_kwargs: dict
21
+ retrieved_docs: dict[str, list[Document]]
22
+ merged_docs: list[Document]
23
+ context: str
24
+ answer: str
25
+
26
+
27
+ def basic_retrieval_graph(prompt: str, vectorstore, model_name: str) -> lg.Graph:
28
+ """
29
+ Baseline RAG retrieval graph using LangGraph. Searches one or several vectorstores,
30
+ passes retrieved documents to the model, and returns the final answer.
31
+
32
+ Args:
33
+ prompt: str, the prompt to use for the model
34
+ vectorstore: the vectorstore(s) to search
35
+ model_name: str, the name of the model to use for the response
36
+
37
+ Returns:
38
+ A LangGraph graph that can be executed for RAG
39
+ """
40
+ # Initialize the model and prompt template
41
+ model = get_chat_model(model_name)
42
+ prompt_template = ChatPromptTemplate.from_template(prompt)
43
+
44
+ # Create a list of retrievers from the vectorstore(s)
45
+ if isinstance(vectorstore, list):
46
+ retrievers = [LangChainVectorStore(vector_store=vs) for vs in vectorstore]
47
+ else:
48
+ retrievers = [LangChainVectorStore(vector_store=vectorstore)]
49
+
50
+ # Define the retrieval function
51
+ def retrieve_documents(state: GraphState) -> GraphState:
52
+ """Retrieve documents from all vectorstores."""
53
+ question = state["question"]
54
+ search_kwargs = state.get("search_kwargs", {}) # Get search params from state
55
+ docs_dict = {}
56
+
57
+ for i, retriever in enumerate(retrievers):
58
+ docs_dict[f"docs_{i}"] = retriever.invoke(
59
+ question,
60
+ config={"metadata": {"search_kwargs": search_kwargs}}, # Pass search params to retriever
61
+ )
62
+
63
+ return {
64
+ "question": question,
65
+ "search_kwargs": search_kwargs,
66
+ "retrieved_docs": docs_dict,
67
+ }
68
+
69
+ # Define the document merging function
70
+ def merge_docs(state: GraphState) -> GraphState:
71
+ """Merge all retrieved documents into a single list."""
72
+ docs_dict = state["retrieved_docs"]
73
+ all_docs = []
74
+
75
+ for i in range(len(retrievers)):
76
+ all_docs.extend(docs_dict[f"docs_{i}"])
77
+
78
+ return {
79
+ "question": state["question"],
80
+ "retrieved_docs": state["retrieved_docs"],
81
+ "merged_docs": all_docs,
82
+ }
83
+
84
+ # Define the document processing function
85
+ def process_docs(state: GraphState) -> GraphState:
86
+ """Process the merged documents into a context string."""
87
+ docs = state["merged_docs"]
88
+ context = combine_documents(docs)
89
+ return {
90
+ "question": state["question"],
91
+ "retrieved_docs": state["retrieved_docs"],
92
+ "merged_docs": state["merged_docs"],
93
+ "context": context,
94
+ }
95
+
96
+ # Define the answer generation function
97
+ def generate_answer(state: GraphState) -> GraphState:
98
+ """Generate an answer using the LLM."""
99
+ question = state["question"]
100
+ context = state["context"]
101
+
102
+ prompt_input = {"context": context, "question": question}
103
+
104
+ chain = prompt_template | model
105
+ response = chain.invoke(prompt_input)
106
+ answer = response.content
107
+
108
+ return {
109
+ "question": state["question"],
110
+ "retrieved_docs": state["retrieved_docs"],
111
+ "merged_docs": state["merged_docs"],
112
+ "context": state["context"],
113
+ "answer": answer,
114
+ }
115
+
116
+ # Build the graph with the defined state schema
117
+ graph = StateGraph(GraphState)
118
+
119
+ # Add nodes
120
+ graph.add_node("retrieve", retrieve_documents)
121
+ graph.add_node("merge", merge_docs)
122
+ graph.add_node("process", process_docs)
123
+ graph.add_node("generate", generate_answer)
124
+
125
+ # Define the edges
126
+ graph.add_edge("retrieve", "merge")
127
+ graph.add_edge("merge", "process")
128
+ graph.add_edge("process", "generate")
129
+ graph.add_edge("generate", END)
130
+
131
+ # Set the entry point
132
+ graph.set_entry_point("retrieve")
133
+
134
+ # Compile the graph
135
+ return graph.compile()
136
+
137
+
138
+ if __name__ == "__main__":
139
+ # Example of how to run the graph correctly
140
+ import os
141
+
142
+ os.environ["CHATLAS_EMBEDDING_MODEL_PATH"] = "<PATH TO YOUR EMBEDDING MODEL>"
143
+ os.environ["CHATLAS_OPENAI_KEY"] = "YOUR OPENAI API KEY"
144
+ os.environ["CHATLAS_DB_PASSWORD"] = "<>"
145
+
146
+ from ..prompt.starters import CHAT_PROMPT_TEMPLATE
147
+ from ..vectorstore import vectorstore
148
+
149
+ graph = basic_retrieval_graph(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=vectorstore, model_name="gpt-4o-mini")
150
+
151
+ ans = graph.invoke(
152
+ {
153
+ "question": "How many onions are in ATLAS",
154
+ "search_kwargs": {
155
+ "k_text": 3,
156
+ "k": 10,
157
+ "date_filter": "01-01-2010",
158
+ "type": ["CDS", "twiki", "Indico"],
159
+ },
160
+ }
161
+ )
162
+
163
+ print(ans)
File without changes
@@ -0,0 +1,447 @@
1
+ """
2
+ Example graph for running langgraph with RAG retrieval plus web search capability.
3
+ Combines database retrieval with web search results for enhanced information gathering.
4
+ """
5
+
6
+ import json
7
+ import logging
8
+ import sys
9
+ from collections.abc import Callable
10
+ from operator import itemgetter
11
+ from pathlib import Path
12
+ from typing import Annotated, Any, Optional, TypedDict, Union
13
+
14
+ import langgraph.graph as lg
15
+ from langchain_core.documents import Document
16
+ from langchain_core.prompts import ChatPromptTemplate
17
+ from langgraph.graph import END, StateGraph
18
+
19
+ from chATLAS_Embed import LangChainVectorStore
20
+
21
+ # Handle both relative and absolute imports for different execution contexts
22
+ try:
23
+ from ..llm.model_selection import get_chat_model
24
+ from ..utils.doc_utils import combine_documents
25
+ except ImportError:
26
+ # For direct execution, use absolute imports
27
+ package_root = Path(__file__).parent.parent.parent
28
+ sys.path.insert(0, str(package_root))
29
+ from chATLAS_Chains.llm.model_selection import get_chat_model
30
+ from chATLAS_Chains.utils.doc_utils import combine_documents
31
+
32
+ logger = logging.getLogger(__name__)
33
+
34
+
35
+ # Define TypedDict for the state with web search support
36
+ class GraphState(TypedDict, total=False):
37
+ question: str
38
+ search_kwargs: dict
39
+ retrieved_docs: dict[str, list[Document]]
40
+ websearch_source: list[dict[str, Any]]
41
+ merged_docs: list[Document]
42
+ context: str
43
+ answer: str
44
+
45
+
46
+ def web_search(query: str) -> str:
47
+ """
48
+ Performs a web search using OpenAI's gpt-4o-search-preview model
49
+ and formats the result as a JSON object containing title, url, and content.
50
+
51
+ Parameters
52
+ ----------
53
+ query : str
54
+ The search query string.
55
+
56
+ Returns
57
+ -------
58
+ str
59
+ A JSON string representing the search result, or an error message in JSON format.
60
+ """
61
+ try:
62
+ import os
63
+
64
+ from openai import OpenAI
65
+
66
+ api_key = os.getenv("CHATLAS_OPENAI_KEY")
67
+ if not api_key:
68
+ raise ValueError("CHATLAS_OPENAI_KEY not set in environment")
69
+ api_key = api_key.strip()
70
+
71
+ client = OpenAI(api_key=api_key)
72
+
73
+ logger.info(f"Performing web search for query: {query}")
74
+ completion = client.chat.completions.create(
75
+ model="gpt-4o-search-preview",
76
+ web_search_options={"search_context_size": "low"},
77
+ messages=[
78
+ {
79
+ "role": "user",
80
+ "content": query,
81
+ }
82
+ ],
83
+ )
84
+
85
+ message = completion.choices[0].message
86
+ content = message.content
87
+ title = None
88
+ url = None
89
+ citations = []
90
+
91
+ # Extract title and url from url_citation annotations
92
+ if hasattr(message, "annotations") and message.annotations:
93
+ for annotation in message.annotations:
94
+ if annotation.type == "url_citation" and hasattr(annotation, "url_citation"):
95
+ citation = {
96
+ "title": getattr(annotation.url_citation, "title", None),
97
+ "url": getattr(annotation.url_citation, "url", None),
98
+ }
99
+ citations.append(citation)
100
+ # Use the first citation for main title/url if not already set
101
+ if title is None and url is None:
102
+ title = citation["title"]
103
+ url = citation["url"]
104
+
105
+ result = {
106
+ "title": title or "Web Search Result",
107
+ "url": url or "#",
108
+ "content": content,
109
+ "type": "WebSearch",
110
+ "citations": citations,
111
+ }
112
+
113
+ return json.dumps(result, ensure_ascii=False, indent=4)
114
+
115
+ except Exception as e:
116
+ logger.error(f"Web search error: {e!s}")
117
+ error_result = {"error": str(e), "type": "WebSearchError"}
118
+ return json.dumps(error_result, ensure_ascii=False, indent=4)
119
+
120
+
121
+ def is_low_quality_web_result(content: str) -> bool:
122
+ """
123
+ Detect if web search result is low quality or contains "no information" responses.
124
+
125
+ Parameters
126
+ ----------
127
+ content : str
128
+ Content of the web search result
129
+
130
+ Returns
131
+ -------
132
+ bool
133
+ True if the result is low quality and should be filtered out
134
+ """
135
+ content_lower = content.lower()
136
+
137
+ # Patterns that indicate "no information" responses
138
+ no_info_patterns = [
139
+ "unable to locate",
140
+ "unable to find",
141
+ "i don't have access",
142
+ "i cannot find",
143
+ "no specific information",
144
+ "not publicly available",
145
+ "may not be publicly available",
146
+ "i recommend the following steps",
147
+ "contact the developers",
148
+ "check internal documentation",
149
+ "i may be able to assist you further",
150
+ "if you can provide more context",
151
+ ]
152
+
153
+ # Check if content contains any "no information" patterns
154
+ for pattern in no_info_patterns:
155
+ if pattern in content_lower:
156
+ logger.info(f"Filtering out low-quality web result containing: '{pattern}'")
157
+ return True
158
+
159
+ # Check if content is too short (likely not useful)
160
+ if len(content.strip()) < 100:
161
+ logger.info("Filtering out web result: content too short")
162
+ return True
163
+
164
+ return False
165
+
166
+
167
+ def process_web_search_results(query: str, results_json: str) -> list[dict[str, Any]]:
168
+ """
169
+ Process web search results, calculate similarity for each result and filter out low-quality results.
170
+
171
+ Parameters
172
+ ----------
173
+ query : str
174
+ Search query
175
+ results_json : str
176
+ JSON string returned by web_search function
177
+
178
+ Returns
179
+ -------
180
+ list[dict[str, Any]]
181
+ Processed results list, each result contains similarity field, filtered for quality
182
+ """
183
+ try:
184
+ results = json.loads(results_json)
185
+
186
+ # If it's an error result, return empty list
187
+ if "error" in results:
188
+ return []
189
+
190
+ # Convert single result to list
191
+ if not isinstance(results, list):
192
+ results = [results]
193
+
194
+ filtered_results = []
195
+
196
+ # Process and filter each result
197
+ for result in results:
198
+ content = result.get("content", "")
199
+
200
+ # Skip low-quality results
201
+ if is_low_quality_web_result(content):
202
+ continue
203
+
204
+ # Add similarity score (placeholder, could be improved with semantic similarity)
205
+ similarity_score = calculate_basic_similarity(query, content)
206
+ result["similarity"] = similarity_score
207
+ result["source_priority"] = 2 # Web search sources have lower priority
208
+
209
+ filtered_results.append(result)
210
+
211
+ # Sort by similarity score (descending)
212
+ filtered_results.sort(key=lambda x: x.get("similarity", 0), reverse=True)
213
+
214
+ logger.info(f"Processed {len(filtered_results)} web search results (filtered from {len(results)} total)")
215
+
216
+ return filtered_results
217
+
218
+ except Exception as e:
219
+ logger.error(f"Error processing web search results: {e!s}")
220
+ return []
221
+
222
+
223
+ def calculate_basic_similarity(query: str, content: str) -> float:
224
+ """
225
+ Calculate basic similarity between query and content based on word overlap.
226
+ This is a simple implementation - could be improved with semantic similarity.
227
+
228
+ Parameters
229
+ ----------
230
+ query : str
231
+ Search query
232
+ content : str
233
+ Content to compare against
234
+
235
+ Returns
236
+ -------
237
+ float
238
+ Similarity score between 0 and 1
239
+ """
240
+ query_words = set(query.lower().split())
241
+ content_words = set(content.lower().split())
242
+
243
+ if not query_words:
244
+ return 0.0
245
+
246
+ overlap = len(query_words.intersection(content_words))
247
+ return overlap / len(query_words)
248
+
249
+
250
+ def websearch_retrieval_graph(vectorstore, model_name: str, enable_web_search: bool = True) -> lg.Graph:
251
+ """
252
+ Enhanced RAG retrieval graph using LangGraph with web search capability.
253
+ Searches vectorstore(s), optionally performs web search, and generates answers.
254
+
255
+ Args:
256
+ vectorstore: the vectorstore(s) to search
257
+ model_name: str, the name of the model to use for the response
258
+ enable_web_search: bool, whether to enable web search functionality
259
+
260
+ Returns:
261
+ A LangGraph graph that can be executed for RAG with web search
262
+ """
263
+ from chATLAS_Chains.prompt.starters import WEB_SEARCH_CHAT_PROMPT_TEMPLATE
264
+
265
+ # Initialize the model and prompt template
266
+ model = get_chat_model(model_name)
267
+ prompt_template = ChatPromptTemplate.from_template(WEB_SEARCH_CHAT_PROMPT_TEMPLATE)
268
+
269
+ # Create a list of retrievers from the vectorstore(s)
270
+ if isinstance(vectorstore, list):
271
+ retrievers = [LangChainVectorStore(vector_store=vs) for vs in vectorstore]
272
+ else:
273
+ retrievers = [LangChainVectorStore(vector_store=vectorstore)]
274
+
275
+ # Define the retrieval function
276
+ def retrieve_documents(state: GraphState) -> GraphState:
277
+ """Retrieve documents from all vectorstores."""
278
+ question = state["question"]
279
+ search_kwargs = state.get("search_kwargs", {}) # Get search params from state
280
+ docs_dict = {}
281
+
282
+ for i, retriever in enumerate(retrievers):
283
+ try:
284
+ docs_dict[f"docs_{i}"] = retriever.invoke(
285
+ question,
286
+ config={"metadata": {"search_kwargs": search_kwargs}}, # Pass search params to retriever
287
+ )
288
+ logger.info(f"Retrieved {len(docs_dict[f'docs_{i}'])} documents from vectorstore {i}")
289
+ except Exception as e:
290
+ logger.error(f"Error retrieving from vectorstore {i}: {e}")
291
+ docs_dict[f"docs_{i}"] = []
292
+
293
+ # Web search functionality
294
+ websearch_results = []
295
+ if enable_web_search:
296
+ try:
297
+ logger.info("Performing web search...")
298
+ websearch_json = web_search(question)
299
+ websearch_results = process_web_search_results(question, websearch_json)
300
+ logger.info(f"Retrieved {len(websearch_results)} web search results")
301
+ except Exception as e:
302
+ logger.error(f"Web search failed: {e!s}")
303
+ websearch_results = []
304
+
305
+ return {
306
+ "question": question,
307
+ "search_kwargs": search_kwargs,
308
+ "retrieved_docs": docs_dict,
309
+ "websearch_source": websearch_results,
310
+ }
311
+
312
+ # Define the document merging function
313
+ def merge_docs(state: GraphState) -> GraphState:
314
+ """Merge all retrieved documents into a single list, including web search results."""
315
+ docs_dict = state["retrieved_docs"]
316
+ websearch_results = state.get("websearch_source", [])
317
+ all_docs = []
318
+
319
+ # Add vectorstore documents
320
+ for i in range(len(retrievers)):
321
+ vectorstore_docs = docs_dict.get(f"docs_{i}", [])
322
+ all_docs.extend(vectorstore_docs)
323
+
324
+ # Convert web search results to Document objects
325
+ for web_result in websearch_results:
326
+ if "content" in web_result and "error" not in web_result:
327
+ web_doc = Document(
328
+ page_content=web_result["content"],
329
+ metadata={
330
+ "title": web_result.get("title", "Web Search Result"),
331
+ "url": web_result.get("url", "#"),
332
+ "type": web_result.get("type", "WebSearch"),
333
+ "similarity": web_result.get("similarity", 0.0),
334
+ "source_priority": web_result.get("source_priority", 2),
335
+ "citations": web_result.get("citations", []),
336
+ },
337
+ )
338
+ all_docs.append(web_doc)
339
+
340
+ logger.info(f"Merged {len(all_docs)} total documents (vectorstore + web search)")
341
+
342
+ return {
343
+ "question": state["question"],
344
+ "retrieved_docs": state["retrieved_docs"],
345
+ "websearch_source": state["websearch_source"],
346
+ "merged_docs": all_docs,
347
+ }
348
+
349
+ # Define the document processing function
350
+ def process_docs(state: GraphState) -> GraphState:
351
+ """Process the merged documents into a context string."""
352
+ docs = state["merged_docs"]
353
+
354
+ # Sort documents by source priority (1=internal, 2=web) and similarity
355
+ # Internal ATLAS sources should be prioritized
356
+ sorted_docs = sorted(
357
+ docs,
358
+ key=lambda x: (
359
+ x.metadata.get("source_priority", 1), # Internal sources first
360
+ -x.metadata.get("similarity", 0.0), # Then by similarity descending
361
+ ),
362
+ )
363
+
364
+ context = combine_documents(sorted_docs)
365
+ logger.info(f"Processed {len(docs)} documents into context string")
366
+
367
+ return {
368
+ "question": state["question"],
369
+ "retrieved_docs": state["retrieved_docs"],
370
+ "websearch_source": state["websearch_source"],
371
+ "merged_docs": state["merged_docs"],
372
+ "context": context,
373
+ }
374
+
375
+ # Define the answer generation function
376
+ def generate_answer(state: GraphState) -> GraphState:
377
+ """Generate an answer using the LLM with enhanced context including web search results."""
378
+ question = state["question"]
379
+ context = state["context"]
380
+
381
+ prompt_input = {"context": context, "question": question}
382
+
383
+ chain = prompt_template | model
384
+ response = chain.invoke(prompt_input)
385
+
386
+ # Extract the answer content
387
+ answer_content = response.content if hasattr(response, "content") else str(response)
388
+
389
+ logger.info("Generated answer using LLM")
390
+
391
+ return {
392
+ "question": state["question"],
393
+ "retrieved_docs": state["retrieved_docs"],
394
+ "websearch_source": state["websearch_source"],
395
+ "merged_docs": state["merged_docs"],
396
+ "context": state["context"],
397
+ "answer": answer_content,
398
+ }
399
+
400
+ # Build the graph with the defined state schema
401
+ graph = StateGraph(GraphState)
402
+
403
+ # Add nodes
404
+ graph.add_node("retrieve", retrieve_documents)
405
+ graph.add_node("merge", merge_docs)
406
+ graph.add_node("process", process_docs)
407
+ graph.add_node("generate", generate_answer)
408
+
409
+ # Define the edges
410
+ graph.add_edge("retrieve", "merge")
411
+ graph.add_edge("merge", "process")
412
+ graph.add_edge("process", "generate")
413
+ graph.add_edge("generate", END)
414
+
415
+ # Set the entry point
416
+ graph.set_entry_point("retrieve")
417
+
418
+ # Compile the graph
419
+ return graph.compile()
420
+
421
+
422
+ if __name__ == "__main__":
423
+ # Example of how to run the graph correctly
424
+ import os
425
+
426
+ os.environ["CHATLAS_EMBEDDING_MODEL_PATH"] = "<PATH TO YOUR EMBEDDING MODEL>"
427
+ os.environ["CHATLAS_OPENAI_KEY"] = "YOUR OPENAI API KEY"
428
+ os.environ["CHATLAS_DB_PASSWORD"] = "<>"
429
+
430
+ from ..prompt.starters import WEB_SEARCH_CHAT_PROMPT_TEMPLATE
431
+ from ..vectorstore import vectorstore
432
+
433
+ graph = websearch_retrieval_graph(vectorstore=vectorstore, model_name="gpt-4o-mini", enable_web_search=True)
434
+
435
+ ans = graph.invoke(
436
+ {
437
+ "question": "How many onions are in ATLAS?",
438
+ "search_kwargs": {
439
+ "k_text": 3,
440
+ "k": 10,
441
+ "date_filter": "01-01-2010",
442
+ "type": ["CDS", "twiki", "Indico"],
443
+ },
444
+ }
445
+ )
446
+
447
+ print(ans)
File without changes