chATLAS_Chains 0.1.5__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.5 → chatlas_chains-0.1.6}/PKG-INFO +7 -1
  2. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/README.md +6 -0
  3. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/advanced.py +32 -35
  4. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rrf.py +29 -26
  5. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/groq.py +2 -1
  6. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/model_selection.py +9 -5
  7. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/PKG-INFO +7 -1
  8. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/SOURCES.txt +1 -0
  9. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/pyproject.toml +1 -1
  10. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_chains.py +31 -19
  11. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_llm.py +26 -4
  12. chatlas_chains-0.1.6/tests/test_rrf.py +126 -0
  13. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/LICENSE +0 -0
  14. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/__init__.py +0 -0
  15. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/__init__.py +0 -0
  16. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic.py +0 -0
  17. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/basic_graph.py +0 -0
  18. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/enhanced_agentic_graph.py +0 -0
  19. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/chains/websearch_retrieval_chain.py +0 -0
  20. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/documents/rerank.py +0 -0
  21. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/llm/__init__.py +0 -0
  22. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/log.py +0 -0
  23. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/__init__.py +0 -0
  24. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/doc_joiners.py +0 -0
  25. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/prompt/starters.py +0 -0
  26. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/query/query_rewriting.py +0 -0
  27. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/search/__init__.py +0 -0
  28. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/search/basic.py +0 -0
  29. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/__init__.py +0 -0
  30. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/utils/doc_utils.py +0 -0
  31. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains/vectorstore.py +0 -0
  32. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/dependency_links.txt +0 -0
  33. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/requires.txt +0 -0
  34. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/chATLAS_Chains.egg-info/top_level.txt +0 -0
  35. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/setup.cfg +0 -0
  36. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/__init__.py +0 -0
  37. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/conftest.py +0 -0
  38. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_groq.py +0 -0
  39. {chatlas_chains-0.1.5 → chatlas_chains-0.1.6}/tests/test_search.py +0 -0
  40. {chatlas_chains-0.1.5 → 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.5
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,12 @@ 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
+
94
100
  #### 0.1.5
95
101
 
96
102
  Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
@@ -65,6 +65,12 @@ 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
+
68
74
  #### 0.1.5
69
75
 
70
76
  Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
@@ -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
@@ -30,6 +31,8 @@ from chATLAS_Chains.search.basic import search_runnable
30
31
  from chATLAS_Chains.utils.doc_utils import combine_documents
31
32
  from chATLAS_Embed.Base import VectorStore
32
33
 
34
+ logger = logging.getLogger(__name__)
35
+
33
36
 
34
37
  # Define TypedDict for the simplified state
35
38
  class HybridGraphState(TypedDict, total=False):
@@ -57,6 +60,7 @@ def advanced_rag(
57
60
  rrf_constant: float = 60.0,
58
61
  rrf_weights: dict[str, float] | None = None,
59
62
  retry_config: RetryConfig | None = None,
63
+ fallback_models: list[str] | None = None,
60
64
  ) -> CompiledStateGraph:
61
65
  """
62
66
  Advanced Agentic RAG graph with optional query rewriting, dual-stage reranking and self-evaluation.
@@ -77,6 +81,7 @@ def advanced_rag(
77
81
  :param rrf_constant: Constant to use for RRF (Reciprocal Rank Fusion) calculations.
78
82
  :param rrf_weights: How to weight the RRF score for each retriever. Default is 1.0 for all.
79
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.
80
85
 
81
86
  :return: A compiled LangGraph with the chosen features
82
87
 
@@ -91,13 +96,15 @@ def advanced_rag(
91
96
  searcher = search_runnable(vectorstore)
92
97
 
93
98
  # Initialise models
94
- model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config)
99
+ model = get_chat_model(model_name, max_tokens, temperature, use_preview_models, retry_config, fallback_models)
95
100
 
96
101
  if enable_query_rewriting:
97
102
  query_rewriting_model_instance = get_chat_model(
98
103
  model_name=query_rewriting_model,
99
104
  temperature=query_rewriting_temperature,
100
105
  use_preview_models=use_preview_models,
106
+ retry_config=retry_config,
107
+ fallback_models=fallback_models,
101
108
  )
102
109
 
103
110
  # --------------- Define functions for the graph nodes ---------------
@@ -107,7 +114,7 @@ def advanced_rag(
107
114
  question = state["question"]
108
115
 
109
116
  if not enable_query_rewriting:
110
- print("[WARNING] Query rewriting disabled, but query_rewriting node was called.")
117
+ logger.warning("Query rewriting disabled, but query_rewriting node was called.")
111
118
  # return original state
112
119
  return {**state}
113
120
 
@@ -128,7 +135,7 @@ def advanced_rag(
128
135
 
129
136
  search_kwargs = state.get("search_kwargs", {})
130
137
  if not search_kwargs:
131
- print("[WARNING] No search_kwargs provided, using defaults.")
138
+ logger.warning("No search_kwargs provided, using defaults.")
132
139
  results = searcher.invoke(question)
133
140
  else:
134
141
  results = searcher.invoke(question, config={"metadata": {"search_kwargs": search_kwargs}})
@@ -136,7 +143,7 @@ def advanced_rag(
136
143
  if "docs" not in results:
137
144
  raise Exception("Search results missing 'docs' field")
138
145
 
139
- print(f"Retrieved {len(results['docs'])} documents")
146
+ logger.debug(f"Retrieved {len(results['docs'])} documents")
140
147
 
141
148
  return {
142
149
  **state,
@@ -150,7 +157,7 @@ def advanced_rag(
150
157
  """
151
158
 
152
159
  if not enable_reranking:
153
- print("[WARNING] Reranking disabled, but rerank node was called.")
160
+ logger.warning("Reranking disabled, but rerank node was called.")
154
161
  # return original state
155
162
  return {**state}
156
163
 
@@ -166,7 +173,7 @@ def advanced_rag(
166
173
  )
167
174
 
168
175
  except Exception as e:
169
- print(f"[WARNING] Reranking failed with exception {e}. Returning original documents.")
176
+ logger.warning(f"Reranking failed with exception {e}. Returning original documents.")
170
177
  reranked_docs = docs
171
178
 
172
179
  return {
@@ -177,13 +184,13 @@ def advanced_rag(
177
184
  def rrf(state: HybridGraphState) -> HybridGraphState:
178
185
  """Perform Reciprocal Rank Fusion (RRF) on retrieved documents."""
179
186
  if not enable_rrf:
180
- print("[WARNING] RRF disabled, but rrf node was called.")
187
+ logger.warning("RRF disabled, but rrf node was called.")
181
188
  # return original state
182
189
  return {**state}
183
190
 
184
191
  docs = state.get("docs", [])
185
192
  if not docs:
186
- print("[WARNING] No documents retrieved for RRF.")
193
+ logger.warning("No documents retrieved for RRF.")
187
194
  return {**state}
188
195
 
189
196
  try:
@@ -191,7 +198,7 @@ def advanced_rag(
191
198
  results=split_docs_by_retriever(docs), k=rrf_constant, weights=rrf_weights
192
199
  )
193
200
  except Exception as e:
194
- print(f"[WARNING] RRF failed with exception {e}. Returning original documents.")
201
+ logger.warning(f"RRF failed with exception {e}. Returning original documents.")
195
202
  rrf_docs = docs
196
203
 
197
204
  return {
@@ -201,7 +208,6 @@ def advanced_rag(
201
208
 
202
209
  def generate_answer(state: HybridGraphState) -> HybridGraphState:
203
210
  """Generate a response from the LLM"""
204
-
205
211
  # Format the prompt using LangChain template
206
212
  prompt_input = {"context": combine_documents(state["docs"]), "question": state["question"]}
207
213
  final_prompt = prompt_template.format_messages(**prompt_input)
@@ -266,6 +272,7 @@ if __name__ == "__main__":
266
272
  from chATLAS_Chains.vectorstore import get_vectorstore
267
273
 
268
274
  twiki = get_vectorstore("twiki_prod")
275
+ mkdocs = get_vectorstore("mkdocs_prod_v1")
269
276
 
270
277
  retry_config = RetryConfig(
271
278
  max_retries=1,
@@ -274,34 +281,24 @@ if __name__ == "__main__":
274
281
 
275
282
  # Create the hybrid graph
276
283
  graph = advanced_rag(
277
- vectorstore=[twiki],
284
+ vectorstore=[twiki, mkdocs],
278
285
  model_name=GROQ_PRODUCTION_MODELS[0],
279
286
  enable_query_rewriting=True,
280
287
  enable_rrf=True,
281
- enable_reranking=True,
288
+ enable_reranking=False,
282
289
  )
283
290
 
284
- # Test query
285
- try:
286
- ans = graph.invoke(
287
- {
288
- "question": "How can one check for and remove bad or corrupted events in the analysis?",
289
- "search_kwargs": {
290
- "k_text": 3,
291
- "k": 15,
292
- "date_filter": "01-01-2010",
293
- # "type": ["twiki"],
294
- },
295
- }
296
- )
297
-
298
- print(f"Number of docs is : {len(ans['docs'])}")
299
- print(f"Answer: {ans['answer']}")
300
-
301
- except Exception as e:
302
- print(f"❌ Graph execution failed with error: {e}")
303
- print(f"Error type: {type(e).__name__}")
304
- 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
+ )
305
302
 
306
- traceback.print_exc()
307
- sys.exit(1)
303
+ print(f"Number of docs is : {len(ans['docs'])}")
304
+ print(f"Answer: {ans['answer']}")
@@ -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()
@@ -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.5
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,12 @@ 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
+
94
100
  #### 0.1.5
95
101
 
96
102
  Fix missing `retry_config` argument in `advanced_rag` caused by early PyPI upload
@@ -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.5"
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"
@@ -22,7 +22,7 @@ from chATLAS_Chains.vectorstore import get_vectorstore
22
22
  @pytest.fixture
23
23
  def retry_config():
24
24
  """Retry config to use in tests"""
25
- return RetryConfig(max_retries=2, max_delay=120.0)
25
+ return RetryConfig(max_retries=5, max_delay=120.0)
26
26
 
27
27
 
28
28
  def test_basic_retrieval_chain_returns_runnablesequence(twiki_vectorstore):
@@ -55,10 +55,11 @@ def test_basic_retrieval_chain_returns_docs_and_answer(twiki_vectorstore, retry_
55
55
  )
56
56
  output = chain.invoke("What is the Higgs boson?")
57
57
 
58
- assert "docs" in output
59
- assert "answer" in output
58
+ assert "docs" in output, "'docs' key not in output"
59
+ assert "answer" in output, "'answer' key not in output"
60
60
 
61
- 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"
62
63
 
63
64
 
64
65
  def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores, retry_config):
@@ -70,10 +71,11 @@ def test_basic_retrieval_chain_multiple_vectorstores(three_vectorstores, retry_c
70
71
  )
71
72
  output = chain.invoke("What is the Higgs boson?")
72
73
 
73
- assert "docs" in output
74
- 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"
75
77
 
76
- assert isinstance(output["docs"], list)
78
+ assert isinstance(output["docs"], list), "'docs' is not a list"
77
79
 
78
80
  sources = set([doc.metadata.get("source") for doc in output["docs"]])
79
81
  expected_sources = {"twiki", "MkDocs", "CDS"}
@@ -101,10 +103,13 @@ def test_advanced_rag_chain(twiki_vectorstore, retry_config):
101
103
  {"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
102
104
  )
103
105
 
104
- assert "docs" in output
105
- assert "answer" in output
106
- assert isinstance(output["docs"], list)
107
- 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
+ )
108
113
 
109
114
  # test with the additional options turned on (besides rerank)
110
115
  chain = advanced_rag(
@@ -119,10 +124,14 @@ def test_advanced_rag_chain(twiki_vectorstore, retry_config):
119
124
  {"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
120
125
  )
121
126
 
122
- assert "docs" in output
123
- assert "answer" in output
124
- assert isinstance(output["docs"], list)
125
- 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"
126
135
 
127
136
 
128
137
  def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores, retry_config):
@@ -143,10 +152,13 @@ def test_advanced_rag_chain_multiple_vectorstores(three_vectorstores, retry_conf
143
152
  {"question": "What is the Higgs boson?", "search_kwargs": search_kwargs},
144
153
  )
145
154
 
146
- assert "docs" in output
147
- assert "answer" in output
148
- assert isinstance(output["docs"], list)
149
- 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
+ )
150
162
 
151
163
  sources = set([doc.metadata.get("source") for doc in output["docs"]])
152
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