chATLAS_Chains 0.1.4__tar.gz → 0.1.6__tar.gz

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.
Files changed (40) hide show
  1. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/PKG-INFO +11 -1
  2. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/README.md +10 -0
  3. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/advanced.py +41 -36
  4. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic.py +8 -1
  5. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rrf.py +29 -26
  6. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/groq.py +2 -1
  7. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/model_selection.py +11 -7
  8. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/PKG-INFO +11 -1
  9. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/SOURCES.txt +1 -0
  10. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/pyproject.toml +1 -1
  11. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_chains.py +58 -31
  12. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_llm.py +26 -4
  13. chatlas_chains-0.1.6/tests/test_rrf.py +126 -0
  14. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/LICENSE +0 -0
  15. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/__init__.py +0 -0
  16. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/__init__.py +0 -0
  17. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic_graph.py +0 -0
  18. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
  19. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
  20. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rerank.py +0 -0
  21. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/__init__.py +0 -0
  22. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/log.py +0 -0
  23. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/__init__.py +0 -0
  24. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  25. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/starters.py +0 -0
  26. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/query/query_rewriting.py +0 -0
  27. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/search/__init__.py +0 -0
  28. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/search/basic.py +0 -0
  29. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/__init__.py +0 -0
  30. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/doc_utils.py +0 -0
  31. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains/vectorstore.py +0 -0
  32. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  33. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/requires.txt +0 -0
  34. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/top_level.txt +0 -0
  35. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/setup.cfg +0 -0
  36. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/__init__.py +0 -0
  37. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/conftest.py +0 -0
  38. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_groq.py +0 -0
  39. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_search.py +0 -0
  40. {chatlas_chains-0.1.4 → chatlas_chains-0.1.6}/tests/test_utils.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: chATLAS_Chains
3
- Version: 0.1.4
3
+ Version: 0.1.6
4
4
  Summary: A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project.
5
5
  Author-email: Joe Egan <joseph.caimin.egan@cern.ch>
6
6
  License: Apache-2.0
@@ -91,6 +91,16 @@ CREATE EXTENSION IF NOT EXISTS vector;
91
91
  ```
92
92
  ## CHANGELOG
93
93
 
94
+ #### 0.1.6
95
+
96
+ Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
97
+
98
+ Add `fallback_models` optional argument to `advanced_rag`
99
+
100
+ #### 0.1.5
101
+
102
+ Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
103
+
94
104
  #### 0.1.4
95
105
 
96
106
  Support for Groq-hosted models
@@ -65,6 +65,16 @@ CREATE EXTENSION IF NOT EXISTS vector;
65
65
  ```
66
66
  ## CHANGELOG
67
67
 
68
+ #### 0.1.6
69
+
70
+ Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
71
+
72
+ Add `fallback_models` optional argument to `advanced_rag`
73
+
74
+ #### 0.1.5
75
+
76
+ Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
77
+
68
78
  #### 0.1.4
69
79
 
70
80
  Support for Groq-hosted models
@@ -11,6 +11,7 @@ upweighting results that appear in both
11
11
  """
12
12
 
13
13
  import argparse
14
+ import logging
14
15
  import os
15
16
  import sys
16
17
  from typing import TypedDict
@@ -22,6 +23,7 @@ from langgraph.graph.state import CompiledStateGraph
22
23
 
23
24
  from chATLAS_Chains.documents.rerank import rerank_documents
24
25
  from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
26
+ from chATLAS_Chains.llm.groq import RetryConfig
25
27
  from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS, get_chat_model
26
28
  from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
27
29
  from chATLAS_Chains.query.query_rewriting import rewrite_query
@@ -29,6 +31,8 @@ from chATLAS_Chains.search.basic import search_runnable
29
31
  from chATLAS_Chains.utils.doc_utils import combine_documents
30
32
  from chATLAS_Embed.Base import VectorStore
31
33
 
34
+ logger = logging.getLogger(__name__)
35
+
32
36
 
33
37
  # Define TypedDict for the simplified state
34
38
  class HybridGraphState(TypedDict, total=False):
@@ -55,6 +59,8 @@ def advanced_rag(
55
59
  pinecone_api_key: str | None = None,
56
60
  rrf_constant: float = 60.0,
57
61
  rrf_weights: dict[str, float] | None = None,
62
+ retry_config: RetryConfig | None = None,
63
+ fallback_models: list[str] | None = None,
58
64
  ) -> CompiledStateGraph:
59
65
  """
60
66
  Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
@@ -73,7 +79,9 @@ def advanced_rag(
73
79
  :param rerank_model: Name of the Pinecone reranker model to use.
74
80
  :param pinecone_api_key: Pinecone API key. If None, will use PINECONE_API_KEY environment variable.
75
81
  :param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
76
- :rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
82
+ :param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
83
+ :param retry_config: Optional RetryConfig for Groq API calls
84
+ :param fallback_models: Optional list of model names to fall back to if the Groq API fails.
77
85
 
78
86
  :return: A compiled LangGraph with the chosen features
79
87
 
@@ -88,13 +96,15 @@ def advanced_rag(
88
96
  searcher = search_runnable(vectorstore)
89
97
 
90
98
  # Initialise models
91
- model = get_chat_model(model_name, max_tokens, temperature, use_preview_models)
99
+ model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config, fallback_models)
92
100
 
93
101
  if enable_query_rewriting:
94
102
  query_rewriting_model_instance = get_chat_model(
95
103
  model_name=query_rewriting_model,
96
104
  temperature=query_rewriting_temperature,
97
105
  use_preview_models=use_preview_models,
106
+ retry_config=retry_config,
107
+ fallback_models=fallback_models,
98
108
  )
99
109
 
100
110
  # --------------- Define functions for the graph nodes ---------------
@@ -104,7 +114,7 @@ def advanced_rag(
104
114
  question = state["question"]
105
115
 
106
116
  if not enable_query_rewriting:
107
- print("[WARNING] Query rewriting disabled, but query_rewriting node was called.")
117
+ logger.warning("Query rewriting disabled, but query_rewriting node was called.")
108
118
  # return original state
109
119
  return {**state}
110
120
 
@@ -125,7 +135,7 @@ def advanced_rag(
125
135
 
126
136
  search_kwargs = state.get("search_kwargs", {})
127
137
  if not search_kwargs:
128
- print("[WARNING] No search_kwargs provided, using defaults.")
138
+ logger.warning("No search_kwargs provided, using defaults.")
129
139
  results = searcher.invoke(question)
130
140
  else:
131
141
  results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
@@ -133,7 +143,7 @@ def advanced_rag(
133
143
  if "docs" not in results:
134
144
  raise Exception("Search results missing 'docs' field")
135
145
 
136
- print(f"Retrieved {len(results['docs'])} documents")
146
+ logger.debug(f"Retrieved {len(results['docs'])} documents")
137
147
 
138
148
  return {
139
149
  **state,
@@ -147,7 +157,7 @@ def advanced_rag(
147
157
  """
148
158
 
149
159
  if not enable_reranking:
150
- print("[WARNING] Reranking disabled, but rerank node was called.")
160
+ logger.warning("Reranking disabled, but rerank node was called.")
151
161
  # return original state
152
162
  return {**state}
153
163
 
@@ -163,7 +173,7 @@ def advanced_rag(
163
173
  )
164
174
 
165
175
  except Exception as e:
166
- print(f"[WARNING] Reranking failed with exception {e}. Returning original documents.")
176
+ logger.warning(f"Reranking failed with exception {e}. Returning original documents.")
167
177
  reranked_docs = docs
168
178
 
169
179
  return {
@@ -174,13 +184,13 @@ def advanced_rag(
174
184
  def rrf(state: HybridGraphState) -> HybridGraphState:
175
185
  """Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
176
186
  if not enable_rrf:
177
- print("[WARNING] RRF disabled, but rrf node was called.")
187
+ logger.warning("RRF disabled, but rrf node was called.")
178
188
  # return original state
179
189
  return {**state}
180
190
 
181
191
  docs = state.get("docs", [])
182
192
  if not docs:
183
- print("[WARNING] No documents retrieved for RRF.")
193
+ logger.warning("No documents retrieved for RRF.")
184
194
  return {**state}
185
195
 
186
196
  try:
@@ -188,7 +198,7 @@ def advanced_rag(
188
198
  results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
189
199
  )
190
200
  except Exception as e:
191
- print(f"[WARNING] RRF failed with exception {e}. Returning original documents.")
201
+ logger.warning(f"RRF failed with exception {e}. Returning original documents.")
192
202
  rrf_docs = docs
193
203
 
194
204
  return {
@@ -198,7 +208,6 @@ def advanced_rag(
198
208
 
199
209
  def generate_answer(state: HybridGraphState) -> HybridGraphState:
200
210
  """Generate a response from the LLM"""
201
-
202
211
  # Format the prompt using LangChain template
203
212
  prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
204
213
  final_prompt = prompt_template.format_messages(**prompt_input)
@@ -263,37 +272,33 @@ if __name__ == "__main__":
263
272
  from chATLAS_Chains.vectorstore import get_vectorstore
264
273
 
265
274
  twiki = get_vectorstore("twiki_prod")
275
+ mkdocs = get_vectorstore("mkdocs_prod_v1")
276
+
277
+ retry_config = RetryConfig(
278
+ max_retries=1,
279
+ max_delay=120.0,
280
+ )
266
281
 
267
282
  # Create the hybrid graph
268
283
  graph = advanced_rag(
269
- vectorstore=[twiki],
284
+ vectorstore=[twiki, mkdocs],
270
285
  model_name=GROQ_PRODUCTION_MODELS[0],
271
286
  enable_query_rewriting=True,
272
287
  enable_rrf=True,
273
- enable_reranking=True,
288
+ enable_reranking=False,
274
289
  )
275
290
 
276
- # Test query
277
- try:
278
- ans = graph.invoke(
279
- {
280
- "question": "How can one check for and remove bad or corrupted events in the analysis?",
281
- "search_kwargs": {
282
- "k_text": 3,
283
- "k": 15,
284
- "date_filter": "01-01-2010",
285
- # "type": ["twiki"],
286
- },
287
- }
288
- )
289
-
290
- print(f"Number of docs is : {len(ans['docs'])}")
291
- print(f"Answer: {ans['answer']}")
292
-
293
- except Exception as e:
294
- print(f"❌ Graph execution failed with error: {e}")
295
- print(f"Error type: {type(e).__name__}")
296
- import traceback
291
+ ans = graph.invoke(
292
+ {
293
+ "question": "What is the crack veto in electron reconstruction?",
294
+ "search_kwargs": {
295
+ "k_text": 3,
296
+ "k": 15,
297
+ "date_filter": "01-01-2010",
298
+ # "type": ["twiki"],
299
+ },
300
+ }
301
+ )
297
302
 
298
- traceback.print_exc()
299
- sys.exit(1)
303
+ print(f"Number of docs is : {len(ans['docs'])}")
304
+ print(f"Answer: {ans['answer']}")
@@ -4,6 +4,7 @@ from typing import Optional
4
4
  from langchain_core.prompts import ChatPromptTemplate
5
5
  from langchain_core.runnables import RunnableSerializable
6
6
 
7
+ from chATLAS_Chains.llm.groq import RetryConfig
7
8
  from chATLAS_Chains.llm.model_selection import get_chat_model
8
9
  from chATLAS_Chains.search.basic import search_runnable
9
10
  from chATLAS_Chains.utils.doc_utils import combine_documents
@@ -16,6 +17,8 @@ def basic_retrieval_chain(
16
17
  model_name: str,
17
18
  max_tokens: int | None = None,
18
19
  temperature: float | None = None,
20
+ use_preview_models: bool = False,
21
+ retry_config: RetryConfig | None = None,
19
22
  ) -> RunnableSerializable:
20
23
  """
21
24
  Baseline RAG retrieval chain. Searches one or several vectorstores in parallel, passes retrieved documents to the model
@@ -30,13 +33,17 @@ def basic_retrieval_chain(
30
33
  :type max_tokens: int | None
31
34
  :param temperature: The temperature to use for the model's response generation. Defaults to None.
32
35
  :type temperature: float | None
36
+ :param use_preview_models: Whether to allow the use of preview models from Groq. Defaults to False.
37
+ :type use_preview_models: bool
38
+ :param retry_config: Configuration for retrying requests to the model in case of failures. Defaults to None, which will use the default retry configuration.
39
+ :type retry_config: RetryConfig | None
33
40
 
34
41
  :return: A LangChain RunnableSerializable chain that performs retrieval and response generation.
35
42
  :rtype: RunnableSerializable
36
43
 
37
44
  """
38
45
  prompt_template = ChatPromptTemplate.from_template(prompt)
39
- model = get_chat_model(model_name, max_tokens, temperature)
46
+ model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
40
47
 
41
48
  search = search_runnable(vectorstore)
42
49
 
@@ -5,7 +5,9 @@ Functions for performing Reciprocal Rank Fusion (RRF) on documents retrieved fro
5
5
  from langchain_core.documents import Document
6
6
 
7
7
 
8
- def reciprocal_rank_fusion(results: dict, k: float = 60, weights: dict[str, float] | None = None) -> list:
8
+ def reciprocal_rank_fusion(
9
+ results: dict[str, list[Document]], k: float = 60, weights: dict[str, float] | None = None
10
+ ) -> list:
9
11
  """
10
12
  Fuse search results using weighted Reciprocal Rank Fusion (RRF) algorithm.
11
13
 
@@ -39,24 +41,21 @@ def reciprocal_rank_fusion(results: dict, k: float = 60, weights: dict[str, floa
39
41
  if key not in weights:
40
42
  weights[key] = 1.0
41
43
 
42
- if len(results) != len(weights):
43
- raise ValueError("Number of results and weights are different")
44
-
45
44
  # Do this for each retriever in the dict
46
45
  for retriever, docs in results.items():
47
46
  for rank, doc in enumerate(docs, 1):
48
- id = doc.id
47
+ name = doc.metadata.get("name")
49
48
 
50
49
  # initialise score to zero
51
- if id not in rrf_scores:
52
- rrf_scores[id] = 0
53
- document_store[id] = doc
50
+ if name not in rrf_scores:
51
+ rrf_scores[name] = 0
52
+ document_store[name] = doc
54
53
 
55
54
  # weighted score for this retriever
56
- rrf_scores[id] += weights[retriever] * (1 / (k + rank))
55
+ rrf_scores[name] += weights[retriever] * (1 / (k + rank))
57
56
 
58
57
  # track which retrievers had this document
59
- metadata = document_store[id].metadata
58
+ metadata = document_store[name].metadata
60
59
  if "retrievers" not in metadata:
61
60
  metadata["retrievers"] = []
62
61
  if retriever not in metadata["retrievers"]:
@@ -64,9 +63,10 @@ def reciprocal_rank_fusion(results: dict, k: float = 60, weights: dict[str, floa
64
63
 
65
64
  # Sort and return results
66
65
  sorted_results = []
67
- for id in sorted(rrf_scores.keys(), key=lambda x: rrf_scores[x], reverse=True):
68
- doc = document_store[id]
69
- doc.metadata["rrf_score"] = rrf_scores[id]
66
+
67
+ for name in sorted(rrf_scores.keys(), key=lambda x: rrf_scores[x], reverse=True):
68
+ doc = document_store[name]
69
+ doc.metadata["rrf_score"] = rrf_scores[name]
70
70
 
71
71
  sorted_results.append(doc)
72
72
 
@@ -90,7 +90,6 @@ def split_docs_by_retriever(docs: list[Document]):
90
90
  retriever_types.append(search_type)
91
91
 
92
92
  split_docs = {retriever_type: [] for retriever_type in retriever_types}
93
-
94
93
  for doc in docs:
95
94
  search_type = doc.metadata["search_type"]
96
95
  if search_type not in split_docs:
@@ -102,36 +101,40 @@ def split_docs_by_retriever(docs: list[Document]):
102
101
 
103
102
  if __name__ == "__main__":
104
103
  # Example usage
105
- from langchain_core.documents import Document
106
104
 
107
105
  # Create some example documents
108
106
  doc1 = Document(
109
107
  page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
110
- id="higgs",
111
- metadata={"search_type": "vector"},
108
+ metadata={"search_type": "vector", "name": "higgs"},
109
+ source="twiki",
110
+ url="",
112
111
  )
113
112
  doc2 = Document(
114
113
  page_content="The top quark is the heaviest known elementary particle",
115
- id="top",
116
- metadata={"search_type": "vector"},
114
+ metadata={"search_type": "vector", "name": "top"},
115
+ source="twiki",
116
+ url="",
117
117
  )
118
118
 
119
119
  vector_results = [doc1, doc2]
120
120
 
121
121
  doc3 = Document(
122
122
  page_content="The gluon is the force carrier of Quantum Chromodynamics (QCD)",
123
- id="gluon",
124
- metadata={"search_type": "text"},
123
+ metadata={"search_type": "text", "name": "gluon"},
124
+ source="twiki",
125
+ url="",
125
126
  )
126
127
  doc4 = Document(
127
128
  page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
128
- id="higgs",
129
- metadata={"search_type": "text"},
129
+ metadata={"search_type": "text", "name": "higgs"},
130
+ source="twiki",
131
+ url="",
130
132
  )
131
133
  doc5 = Document(
132
134
  page_content="The Large Hadron Collider (LHC) is the world's largest and most powerful particle accelerator",
133
- id="lhc",
134
- metadata={"search_type": "text"},
135
+ metadata={"search_type": "text", "name": "lhc"},
136
+ source="twiki",
137
+ url="",
135
138
  )
136
139
 
137
140
  text_results = [doc3, doc4, doc5]
@@ -143,5 +146,5 @@ if __name__ == "__main__":
143
146
  print("RRF Results:")
144
147
  for doc in rrf_results:
145
148
  print(
146
- f"Document ID: {doc.id}, RRF Score: {doc.metadata.get('rrf_score', 0)}, Retrievers: {doc.metadata.get('retrievers', [])}"
149
+ f"Document name: {doc.metadata['name']}, RRF Score: {doc.metadata['rrf_score']}, Retrievers: {doc.metadata['retrievers']}"
147
150
  )
@@ -139,6 +139,7 @@ class AccGPTChatGroq(BaseChatModel, BaseModel):
139
139
  def __init__(self, **data):
140
140
  super().__init__(**data)
141
141
  if self.retry_config is None:
142
+ logger.debug("No retry_config provided, using default")
142
143
  self.retry_config = RetryConfig()
143
144
 
144
145
  if self.fallback_models is None:
@@ -565,7 +566,7 @@ if __name__ == "__main__":
565
566
  api_key = api_key.strip()
566
567
 
567
568
  base_url = os.getenv("CHATLAS_GROQ_BASE_URL")
568
- base_url = "http://localhost:3000"
569
+ # base_url = "http://localhost:3000"
569
570
  if not base_url or not base_url.strip():
570
571
  raise ValueError("CHATLAS_GROQ_BASE_URL not set in environment")
571
572
  base_url = base_url.strip()
@@ -26,8 +26,8 @@ GROQ_PREVIEW_MODELS = [
26
26
  "meta-llama/llama-prompt-guard-2-22m",
27
27
  "meta-llama/llama-prompt-guard-2-86m",
28
28
  "moonshotai/kimi-k2-instruct",
29
- # "playai-tts", # text to audio model
30
- # "playai-tts-arabic", # text to audio model
29
+ # "playai-tts", # text to audio model
30
+ # "playai-tts-arabic", # text to audio model
31
31
  "qwen/qwen3-32b",
32
32
  "gemma2-9b-it",
33
33
  ]
@@ -84,6 +84,7 @@ def get_chat_model(
84
84
  temperature: float | None = None,
85
85
  use_preview_models: bool = False,
86
86
  retry_config: RetryConfig | None = None,
87
+ fallback_models: list[str] | None = None,
87
88
  ):
88
89
  """
89
90
  Initialize chat model with the provided model name (if supported)
@@ -100,6 +101,12 @@ def get_chat_model(
100
101
  :param use_preview_models: If True, allows the use of preview models from Groq. Defaults to False.
101
102
  :type use_preview_models: bool
102
103
 
104
+ :param retry_config: Configuration for retrying requests to Groq models. If None, the default retry configuration is used.
105
+ :type retry_config: RetryConfig | None
106
+
107
+ :param fallback_models: A list of model names to fall back to if the Groq API fails.
108
+ :type fallback_models: list[str] | None
109
+
103
110
  :raises ValueError: If the environment variable `CHATLAS_OPENAI_KEY` is not set when using OpenAI models, or if `CHATLAS_GROQ_KEY` and `CHATLAS_GROQ_BASE_URL` are not set when using Groq models.
104
111
 
105
112
  :return: An instance of the specified chat model.
@@ -143,17 +150,14 @@ def get_chat_model(
143
150
  f"max_tokens ({max_tokens}) exceeds the model's maximum ({get_max_completion_tokens(model_name)})."
144
151
  )
145
152
 
146
- # default retry config
147
- if retry_config is None:
148
- retry_config = RetryConfig()
149
-
150
153
  llm = AccGPTChatGroq(
151
154
  model_name=model_name,
152
155
  api_key=api_key,
153
156
  base_url=base_url,
154
157
  max_tokens=max_tokens,
155
158
  temperature=temperature,
156
- retry_config=retry_config,
159
+ retry_config=retry_config, # if None, AccGPTChatGroq will use the default
160
+ fallback_models=fallback_models,
157
161
  )
158
162
 
159
163
  return llm
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: chATLAS_Chains
3
- Version: 0.1.4
3
+ Version: 0.1.6
4
4
  Summary: A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project.
5
5
  Author-email: Joe Egan <joseph.caimin.egan@cern.ch>
6
6
  License: Apache-2.0
@@ -91,6 +91,16 @@ CREATE EXTENSION IF NOT EXISTS vector;
91
91
  ```
92
92
  ## CHANGELOG
93
93
 
94
+ #### 0.1.6
95
+
96
+ Fix bug in `reciprocal_rank_fusion` which caused it to silently return only one document
97
+
98
+ Add `fallback_models` optional argument to `advanced_rag`
99
+
100
+ #### 0.1.5
101
+
102
+ Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
103
+
94
104
  #### 0.1.4
95
105
 
96
106
  Support for Groq-hosted models
@@ -33,5 +33,6 @@ tests/conftest.py
33
33
  tests/test_chains.py
34
34
  tests/test_groq.py
35
35
  tests/test_llm.py
36
+ tests/test_rrf.py
36
37
  tests/test_search.py
37
38
  tests/test_utils.py
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "chATLAS_Chains"
7
- version = "0.1.4"
7
+ version = "0.1.6"
8
8
  description = "A modular Python package for implementing Retrieval Augmented Generation chains for the chATLAS project."
9
9
  readme = "README.md"
10
10
  requires-python = ">=3.11"
@@ -13,16 +13,19 @@ from pydantic import ValidationError
13
13
 
14
14
  from chATLAS_Chains.chains.advanced import advanced_rag
15
15
  from chATLAS_Chains.chains.basic import basic_retrieval_chain
16
+ from chATLAS_Chains.llm.groq import RetryConfig
16
17
  from chATLAS_Chains.llm.model_selection import GROQ_PRODUCTION_MODELS
17
18
  from chATLAS_Chains.prompt.starters import CHAT_PROMPT_TEMPLATE
18
19
  from chATLAS_Chains.vectorstore import get_vectorstore
19
20
 
20
21
 
21
- def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
22
- """
23
- Test the search_runnable function with a populated vector store.
24
- """
22
+ @pytest.fixture
23
+ def retry_config():
24
+ """Retry config to use in tests"""
25
+ return RetryConfig(max_retries=5, max_delay=120.0)
26
+
25
27
 
28
+ def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
26
29
  chain = basic_retrieval_chain(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=twiki_vectorstore, model_name="gpt-4o-mini")
27
30
 
28
31
  assert isinstance(chain, RunnableSequence)
@@ -30,7 +33,7 @@ def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
30
33
 
31
34
  def test_search_runnable_returns_runnablesequence_with_list_of_vectorstores(three_vectorstores):
32
35
  chain = basic_retrieval_chain(prompt=CHAT_PROMPT_TEMPLATE, vectorstore=three_vectorstores, model_name="gpt-4o-mini")
33
- print(type(chain))
36
+
34
37
  assert isinstance(chain, RunnableSequence)
35
38
 
36
39
 
@@ -43,28 +46,36 @@ def test_search_runnable_error_on_invalid_input():
43
46
  )
44
47
 
45
48
 
46
- def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore):
49
+ def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore, retry_config):
47
50
  chain = basic_retrieval_chain(
48
- prompt=CHAT_PROMPT_TEMPLATE, vectorstore=twiki_vectorstore, model_name=GROQ_PRODUCTION_MODELS[0]
51
+ prompt=CHAT_PROMPT_TEMPLATE,
52
+ vectorstore=twiki_vectorstore,
53
+ model_name=GROQ_PRODUCTION_MODELS[0],
54
+ retry_config=retry_config,
49
55
  )
50
56
  output = chain.invoke("What is the Higgs boson?")
51
57
 
52
- assert "docs" in output
53
- assert "answer" in output
58
+ assert "docs" in output, "'docs' key not in output"
59
+ assert "answer" in output, "'answer' key not in output"
54
60
 
55
- assert isinstance(output["docs"], list)
61
+ assert isinstance(output["docs"], list), "'docs' is not a list"
62
+ assert all([isinstance(obj, Document) for obj in output["docs"]]), "'docs' should be a list of Documents"
56
63
 
57
64
 
58
- def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores):
65
+ def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores, retry_config):
59
66
  chain = basic_retrieval_chain(
60
- prompt=CHAT_PROMPT_TEMPLATE, vectorstore=three_vectorstores, model_name=GROQ_PRODUCTION_MODELS[0]
67
+ prompt=CHAT_PROMPT_TEMPLATE,
68
+ vectorstore=three_vectorstores,
69
+ model_name=GROQ_PRODUCTION_MODELS[0],
70
+ retry_config=retry_config,
61
71
  )
62
72
  output = chain.invoke("What is the Higgs boson?")
63
73
 
64
- assert "docs" in output
65
- assert "answer" in output
74
+ assert "docs" in output, "'docs' key not in output"
75
+ assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3, three vectorstores"
76
+ assert "answer" in output, "'answer' key not in output"
66
77
 
67
- assert isinstance(output["docs"], list)
78
+ assert isinstance(output["docs"], list), "'docs' is not a list"
68
79
 
69
80
  sources = set([doc.metadata.get("source") for doc in output["docs"]])
70
81
  expected_sources = {"twiki", "MkDocs", "CDS"}
@@ -73,7 +84,7 @@ def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores):
73
84
  )
74
85
 
75
86
 
76
- def test_advanced_rag_chain(twiki_vectorstore):
87
+ def test_advanced_rag_chain(twiki_vectorstore, retry_config):
77
88
  # TODO: this is required currently requried for search to work
78
89
  search_kwargs = {
79
90
  "k_text": 3,
@@ -83,16 +94,22 @@ def test_advanced_rag_chain(twiki_vectorstore):
83
94
  }
84
95
 
85
96
  chain = advanced_rag(
86
- prompt=CHAT_PROMPT_TEMPLATE, vectorstore=twiki_vectorstore, model_name=GROQ_PRODUCTION_MODELS[0]
97
+ prompt=CHAT_PROMPT_TEMPLATE,
98
+ vectorstore=twiki_vectorstore,
99
+ model_name=GROQ_PRODUCTION_MODELS[0],
100
+ retry_config=retry_config,
87
101
  )
88
102
  output = chain.invoke(
89
103
  {"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
90
104
  )
91
105
 
92
- assert "docs" in output
93
- assert "answer" in output
94
- assert isinstance(output["docs"], list)
95
- assert isinstance(output["docs"][0], Document)
106
+ assert "docs" in output, "'docs' key not in output"
107
+ assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3"
108
+ assert "answer" in output, "'answer' key not in output"
109
+ assert isinstance(output["docs"], list), "'docs' is not a list"
110
+ assert isinstance(output["docs"][0], Document), (
111
+ f"'docs' should contain instances of {Document}, but instead it has {type(output['docs'][0])}"
112
+ )
96
113
 
97
114
  # test with the additional options turned on (besides rerank)
98
115
  chain = advanced_rag(
@@ -107,13 +124,17 @@ def test_advanced_rag_chain(twiki_vectorstore):
107
124
  {"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
108
125
  )
109
126
 
110
- assert "docs" in output
111
- assert "answer" in output
112
- assert isinstance(output["docs"], list)
113
- assert isinstance(output["docs"][0], Document)
127
+ assert "docs" in output, "'docs' key not in output"
128
+ assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3"
129
+ assert "answer" in output, "'answer' key not in output"
130
+ assert isinstance(output["docs"], list), "'docs' is not a list"
131
+ assert all(isinstance(obj, Document) for obj in output["docs"]), (
132
+ f"'docs' should only contain instances of {Document}"
133
+ )
134
+ assert all("rrf_score" in doc.metadata for doc in output["docs"]), "'rrf_score' not in document metadata"
114
135
 
115
136
 
116
- def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores):
137
+ def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores, retry_config):
117
138
  search_kwargs = {
118
139
  "k_text": 3,
119
140
  "k": 15,
@@ -122,16 +143,22 @@ def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores):
122
143
  }
123
144
 
124
145
  chain = advanced_rag(
125
- prompt=CHAT_PROMPT_TEMPLATE, vectorstore=three_vectorstores, model_name=GROQ_PRODUCTION_MODELS[0]
146
+ prompt=CHAT_PROMPT_TEMPLATE,
147
+ vectorstore=three_vectorstores,
148
+ model_name=GROQ_PRODUCTION_MODELS[0],
149
+ retry_config=retry_config,
126
150
  )
127
151
  output = chain.invoke(
128
152
  {"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
129
153
  )
130
154
 
131
- assert "docs" in output
132
- assert "answer" in output
133
- assert isinstance(output["docs"], list)
134
- assert isinstance(output["docs"][0], Document)
155
+ assert "docs" in output, "'docs' key not in output"
156
+ assert len(output["docs"]) > 1, "Didn't return more than one doc with k=15, k_text=3, three vectorstores"
157
+ assert "answer" in output, "'answer' key not in output"
158
+ assert isinstance(output["docs"], list), "'docs' is not a list"
159
+ assert isinstance(output["docs"][0], Document), (
160
+ f"'docs' should contain instances of {Document}, but instead it has {type(output['docs'][0])}"
161
+ )
135
162
 
136
163
  sources = set([doc.metadata.get("source") for doc in output["docs"]])
137
164
  expected_sources = {"twiki", "MkDocs", "CDS"}
@@ -60,7 +60,9 @@ class TestModelSelection:
60
60
  """Test that Groq models are correctly initialized."""
61
61
  mock_class, mock_instance = mock_groq
62
62
 
63
- model = get_chat_model(model_name, use_preview_models=True)
63
+ retry_config = RetryConfig(max_retries=5, base_delay=5.0, max_delay=300.0)
64
+
65
+ model = get_chat_model(model_name, use_preview_models=True, retry_config=retry_config)
64
66
 
65
67
  mock_class.assert_called_once_with(
66
68
  model_name=model_name,
@@ -68,7 +70,8 @@ class TestModelSelection:
68
70
  base_url="http://fake-groq-url.com",
69
71
  max_tokens=get_max_completion_tokens(model_name),
70
72
  temperature=None,
71
- retry_config=RetryConfig(),
73
+ retry_config=retry_config,
74
+ fallback_models=None,
72
75
  )
73
76
  assert model is mock_instance
74
77
 
@@ -97,7 +100,8 @@ class TestModelSelection:
97
100
  base_url="http://fake-groq-url.com",
98
101
  max_tokens=get_max_completion_tokens(preview_model),
99
102
  temperature=None,
100
- retry_config=RetryConfig(),
103
+ retry_config=None,
104
+ fallback_models=None,
101
105
  )
102
106
  assert model is mock_instance
103
107
 
@@ -197,9 +201,27 @@ class TestModelSelection:
197
201
  base_url="http://whitespace-url.com",
198
202
  max_tokens=get_max_completion_tokens(GROQ_PRODUCTION_MODELS[0]),
199
203
  temperature=None,
200
- retry_config=RetryConfig(),
204
+ retry_config=None,
205
+ fallback_models=None,
201
206
  )
202
207
 
203
208
  def test_supported_models_lists(self):
204
209
  """Test that SUPPORTED_CHAT_MODELS is correctly composed of OPENAI_MODELS and GROQ_MODELS."""
205
210
  assert set(SUPPORTED_CHAT_MODELS) == set(OPENAI_MODELS + GROQ_MODELS)
211
+
212
+ def test_retry_config_set_correctly(self):
213
+ """Test that the retry_config is correctly passed to the Groq model."""
214
+ retry_config = RetryConfig(max_retries=3, base_delay=2.0, max_delay=100.0)
215
+
216
+ model = get_chat_model(GROQ_PRODUCTION_MODELS[0], use_preview_models=True, retry_config=retry_config)
217
+
218
+ assert isinstance(model, AccGPTChatGroq), "Returned model is not an instance of AccGPTChatGroq"
219
+ assert model.retry_config.max_retries == 3, (
220
+ f"max_retries not set correctly on model, it set to {model.retry_config.max_retries}"
221
+ )
222
+ assert model.retry_config.base_delay == 2.0, (
223
+ f"base_delay not set correctly on model, it set to {model.retry_config.base_delay}"
224
+ )
225
+ assert model.retry_config.max_delay == 100.0, (
226
+ f"max_delay not set correctly on model, it set to {model.retry_config.max_delay}"
227
+ )
@@ -0,0 +1,126 @@
1
+ import pytest
2
+ from langchain_core.documents import Document
3
+
4
+ from chATLAS_Chains.documents.rrf import reciprocal_rank_fusion, split_docs_by_retriever
5
+
6
+
7
+ @pytest.fixture
8
+ def vector_results():
9
+ doc1 = Document(
10
+ page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
11
+ metadata={"search_type": "vector", "name": "higgs"},
12
+ source="twiki",
13
+ url="",
14
+ )
15
+ doc2 = Document(
16
+ page_content="The top quark is the heaviest known elementary particle",
17
+ # name="top",
18
+ metadata={"search_type": "vector", "name": "top"},
19
+ source="twiki",
20
+ url="",
21
+ )
22
+
23
+ return [doc1, doc2]
24
+
25
+
26
+ @pytest.fixture
27
+ def text_results():
28
+ doc3 = Document(
29
+ page_content="The gluon is the force carrier of Quantum Chromodynamics (QCD)",
30
+ # name="gluon",
31
+ metadata={"search_type": "text", "name": "gluon"},
32
+ source="twiki",
33
+ url="",
34
+ )
35
+ doc4 = Document(
36
+ page_content="In the Standard Model, the Higgs potential is responsible for electroweak symmetry breaking",
37
+ # name="higgs",
38
+ metadata={"search_type": "text", "name": "higgs"},
39
+ source="twiki",
40
+ url="",
41
+ )
42
+ doc5 = Document(
43
+ page_content="The Large Hadron Collider (LHC) is the world's largest and most powerful particle accelerator",
44
+ # name="lhc",
45
+ metadata={"search_type": "text", "name": "lhc"},
46
+ source="twiki",
47
+ url="",
48
+ )
49
+
50
+ return [doc3, doc4, doc5]
51
+
52
+
53
+ @pytest.fixture
54
+ def docs_dict(vector_results, text_results):
55
+ return {
56
+ "vector": vector_results,
57
+ "text": text_results,
58
+ }
59
+
60
+
61
+ def test_split_docs_by_retriever(vector_results, text_results):
62
+ # pass it the combined list, it should re-split
63
+ combined = vector_results + text_results
64
+
65
+ split = split_docs_by_retriever(combined)
66
+
67
+ assert isinstance(split, dict), "return type should be dict"
68
+ assert "text" in split, "missing retriever type in returned dict"
69
+ assert "vector" in split, "missing retriever type in returned dict"
70
+ assert len(split["text"]) == len(text_results), "text results length mismatch"
71
+ assert len(split["vector"]) == len(vector_results), "vector results length mismatch"
72
+
73
+
74
+ def test_split_docs_by_retriever_missing_search_type(vector_results):
75
+ # remove search_type from one doc
76
+ vector_results[0].metadata.pop("search_type")
77
+
78
+ with pytest.raises(ValueError, match="search_type metadata field missing from Document.metadata"):
79
+ split_docs_by_retriever(vector_results)
80
+
81
+
82
+ def test_reciprocal_rank_fusion(docs_dict):
83
+ fused = reciprocal_rank_fusion(docs_dict)
84
+
85
+ assert isinstance(fused, list), "return type should be list"
86
+ assert len(fused) <= 5, "fused results length exceeds top_k"
87
+ assert all(isinstance(doc, Document) for doc in fused), f"all items in fused results should be {Document} instances"
88
+
89
+ assert all("rrf_score" in doc.metadata for doc in fused), "rrf_score missing from Document.metadata"
90
+
91
+ # Check that the documents are ordered by score (highest first)
92
+ scores = [doc.metadata["rrf_score"] for doc in fused]
93
+ assert scores == sorted(scores, reverse=True), "documents are not ordered by rrf_score"
94
+
95
+
96
+ def test_reciprocal_rank_fusion_with_docs_splitter(vector_results, text_results):
97
+ """
98
+ Same as above, but check it still works using the actual splitter function
99
+ """
100
+ fused = reciprocal_rank_fusion(split_docs_by_retriever(vector_results + text_results))
101
+
102
+ assert isinstance(fused, list), "return type should be list"
103
+ assert len(fused) <= 5, "fused results length exceeds top_k"
104
+ assert all(isinstance(doc, Document) for doc in fused), f"all items in fused results should be {Document} instances"
105
+
106
+ assert all("rrf_score" in doc.metadata for doc in fused), "rrf_score missing from Document.metadata"
107
+
108
+ # Check that the documents are ordered by score (highest first)
109
+ scores = [doc.metadata["rrf_score"] for doc in fused]
110
+ assert scores == sorted(scores, reverse=True), "documents are not ordered by rrf_score"
111
+
112
+
113
+ def test_reciprocal_rank_fusion_raise_for_negative_k(docs_dict):
114
+ with pytest.raises(ValueError, match="k must be positive"):
115
+ reciprocal_rank_fusion(docs_dict, k=-7)
116
+
117
+
118
+ def test_reciprocal_rank_fusion_raise_for_mismatched_weights(docs_dict):
119
+ weights = {
120
+ "vector": 1.0,
121
+ "text": 1.0,
122
+ "extra": 1.0,
123
+ }
124
+
125
+ with pytest.raises(ValueError, match="Weight 'extra' not found in provided results dictionary."):
126
+ reciprocal_rank_fusion(docs_dict, weights=weights)
File without changes
File without changes