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 +34 -0
- benchmark/conversational.py +0 -0
- chATLAS_Chains/__init__.py +0 -0
- chATLAS_Chains/chains/__init__.py +0 -0
- chATLAS_Chains/chains/basic.py +42 -0
- chATLAS_Chains/chains/basic_graph.py +163 -0
- chATLAS_Chains/chains/conversational.py +0 -0
- chATLAS_Chains/chains/websearch_retrieval_chain.py +447 -0
- chATLAS_Chains/llm/__init__.py +0 -0
- chATLAS_Chains/llm/model_selection.py +47 -0
- chATLAS_Chains/log.py +20 -0
- chATLAS_Chains/prompt/__init__.py +0 -0
- chATLAS_Chains/prompt/doc_joiners.py +5 -0
- chATLAS_Chains/prompt/starters.py +151 -0
- chATLAS_Chains/search/__init__.py +0 -0
- chATLAS_Chains/search/basic.py +45 -0
- chATLAS_Chains/utils/__init__.py +0 -0
- chATLAS_Chains/utils/doc_utils.py +28 -0
- chATLAS_Chains/vectorstore.py +94 -0
- chatlas_chains-0.1.3.dist-info/METADATA +113 -0
- chatlas_chains-0.1.3.dist-info/RECORD +30 -0
- chatlas_chains-0.1.3.dist-info/WHEEL +5 -0
- chatlas_chains-0.1.3.dist-info/licenses/LICENSE +201 -0
- chatlas_chains-0.1.3.dist-info/top_level.txt +3 -0
- tests/__init__.py +0 -0
- tests/conftest.py +274 -0
- tests/test_chains.py +54 -0
- tests/test_llm.py +81 -0
- tests/test_search.py +48 -0
- tests/test_utils.py +27 -0
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
|